"""The declarative frontend: a metaclass that reads class-body annotations
and field descriptors (rustruct.fields), lazily compiles a rustruct.Codec on
first use, and converts between typed Struct instances and the plain dict
Codec.pack()/unpack() deal with.
byteorder is an inherited class kwarg. An open dispatch registry is a class
kwarg pair: `registry=True` on the base, one arbitrary-named kwarg (its
value is the tag) on each concrete subclass. Registration is eager, but a
registry only freezes into a fixed cases tuple on first schema resolution,
so subclasses imported after the base -- but before first use -- are still
visible.
Resolution (building fields_tuple/Shapes, the Codec, and specialized
__init__/to_mapping/from_mapping methods compiled as AST, never
source-string templating) is per-class and lazy, triggered by first use.
`StructMeta.__new__` installs a small trampoline into every class's own
`__dict__` so resolving an ancestor class never leaks its compiled methods
onto an unresolved subclass.
"""
import ast
import inspect
import typing
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Protocol, cast
from .core import compile as compile_codec
from .fields import MISSING, FieldSpec, Registry
from .scalars import ScalarType
class FromMapping(Protocol):
"""The compiled/trampoline from_mapping signature: `mapping` always,
`ctx` only when a nested Shape needs it. A plain `Callable[[dict, tuple],
Struct]` can't express that second parameter's default, hence this
Protocol instead."""
def __call__(self, mapping: dict[str, Any], ctx: tuple = ()) -> "Struct": ...
class Shape:
"""Converts a decoded wire value to/from its Python-facing value for one
field. `ctx` is a tuple of enclosing frames, outermost first: a raw wire
dict per frame during unpack, the Struct instance per frame during pack.
Only SwitchShape looks past the innermost frame, since a switch's `on=`
can live in an enclosing scope (e.g. IPv4's `protocol` driving dispatch
inside its `body` window)."""
def to_python(self, value, ctx):
return value
def to_wire(self, value, ctx):
return value
SCALAR_SHAPE = Shape()
class StructShape(Shape):
__slots__ = ("cls",)
def __init__(self, cls):
self.cls = cls
def to_python(self, value, ctx):
return self.cls.from_mapping(value, ctx)
def to_wire(self, value, ctx):
return value.to_mapping(ctx) if isinstance(value, Struct) else value
class ArrayShape(Shape):
__slots__ = ("elem",)
def __init__(self, elem):
self.elem = elem
def to_python(self, value, ctx):
return [self.elem.to_python(v, ctx) for v in value]
def to_wire(self, value, ctx):
return [self.elem.to_wire(v, ctx) for v in value]
class ConvertShape(Shape):
"""Wraps a plain scalar/bytes shape with a value-level transcoder, e.g.
u32 <-> ipaddress.IPv4Address."""
__slots__ = ("decode", "encode")
def __init__(self, decode, encode):
self.decode = decode
self.encode = encode
def to_python(self, value, ctx):
return self.decode(value)
def to_wire(self, value, ctx):
return self.encode(value)
def lookup_in_chain(ctx, name):
for frame in reversed(ctx):
if isinstance(frame, dict):
if name in frame:
return frame[name]
elif hasattr(frame, name):
return getattr(frame, name)
raise KeyError(name)
class SwitchShape(Shape):
"""`on_field` is the sibling/ancestor field name driving dispatch, when
`on=` is a bare ref. A non-ref `on=` expression can't be re-evaluated
here, so values pass through unconverted (SCALAR_SHAPE)."""
__slots__ = ("on_field", "cases", "default")
def __init__(self, on_field, cases, default):
self.on_field = on_field
# Read-only: this dispatch table is fixed once the class resolves,
# so nothing can add/remove a case through this Shape afterward.
self.cases = MappingProxyType(dict(cases))
self.default = default
def pick(self, ctx):
if self.on_field is None:
return self.default or SCALAR_SHAPE
try:
tag = lookup_in_chain(ctx, self.on_field)
except KeyError:
return self.default or SCALAR_SHAPE
return self.cases.get(tag, self.default or SCALAR_SHAPE)
def to_python(self, value, ctx):
return self.pick(ctx).to_python(value, ctx)
def to_wire(self, value, ctx):
return self.pick(ctx).to_wire(value, ctx)
def type_spec(value):
"""Resolve any of {ScalarType subclass, Struct subclass, FieldSpec, raw
(kind, opts) tuple} to a (kind, opts, Shape) triple. Struct subclasses
resolve through the target class's own lazy, cached resolution, so
reused nested types are only ever resolved once."""
if isinstance(value, FieldSpec):
return resolve_field_spec(value)
if isinstance(value, type) and issubclass(value, ScalarType):
return value.kind, {}, SCALAR_SHAPE
if isinstance(value, type) and issubclass(value, Struct):
resolved = ensure_resolved(value)
opts = {"fields": resolved.fields_tuple, "byteorder": value.byteorder}
return "struct", opts, StructShape(value)
if isinstance(value, tuple) and len(value) == 2 and isinstance(value[0], str):
return value[0], dict(value[1]), SCALAR_SHAPE
raise TypeError(f"cannot resolve a wire type for {value!r}")
def resolve_cases(cases_src):
if isinstance(cases_src, Registry):
pairs = cases_src.items()
elif isinstance(cases_src, dict):
pairs = cases_src.items()
else:
pairs = list(cases_src)
cases_opts = []
cases_shapes = {}
for tag, val in pairs:
kind, opts, shape = type_spec(val)
cases_opts.append((tag, (kind, opts)))
cases_shapes[tag] = shape
return tuple(cases_opts), cases_shapes
def resolve_field_spec(spec):
if spec.kind is None:
raise TypeError(
"a kind-less FieldSpec (described()) can't be resolved on its own; it must pair with a typed annotation"
)
match spec.kind:
case "array":
elem_kind, elem_opts, elem_shape = type_spec(spec.opts["elem"])
opts = dict(spec.opts)
opts["elem"] = (elem_kind, elem_opts)
# An array of plain scalars needs no per-element conversion;
# collapsing to SCALAR_SHAPE lets compiled methods pass the
# core's list straight through instead of copying element-wise.
shape = SCALAR_SHAPE if elem_shape is SCALAR_SHAPE else ArrayShape(elem_shape)
return "array", opts, shape
case "switch":
cases_opts, cases_shapes = resolve_cases(spec.opts["cases"])
opts = {"on": spec.opts["on"], "cases": cases_opts}
default_shape = None
if spec.opts.get("default") is not None:
default_kind, default_opts, default_shape = type_spec(spec.opts["default"])
opts["default"] = (default_kind, default_opts)
on = spec.opts["on"]
on_field = on[1] if isinstance(on, tuple) and on[0] == "ref" else None
return "switch", opts, SwitchShape(on_field, cases_shapes, default_shape)
case "struct":
struct_cls = spec.opts["struct"]
resolved = ensure_resolved(struct_cls)
opts = {k: v for k, v in spec.opts.items() if k != "struct"}
opts["fields"] = resolved.fields_tuple
opts["byteorder"] = struct_cls.byteorder
return "struct", opts, StructShape(struct_cls)
case "convert":
base_kind, base_opts, _base_shape = type_spec(spec.opts["base"])
return base_kind, base_opts, ConvertShape(spec.opts["decode"], spec.opts["encode"])
case "cond":
# then_shape is reused as-is: to_python/to_wire for this field
# only run when the key is present, guarded by from_mapping's
# "if name in mapping" check and to_mapping's own None-check.
then_kind, then_opts, then_shape = type_spec(spec.opts["then"])
opts = {"pred": spec.opts["pred"], "then": (then_kind, then_opts)}
return "cond", opts, then_shape
case _:
return spec.kind, dict(spec.opts), SCALAR_SHAPE
class FieldDecl:
"""One class-body declaration -- name plus its raw annotation/FieldSpec,
default, and help -- before `resolve_field_spec` turns it into a wire
kind/opts/Shape. Not to be confused with `Field`, the (name, kind, opts)
NamedTuple this eventually compiles down to."""
__slots__ = ("name", "value_spec", "default", "default_factory", "help")
def __init__(self, name, value_spec, default, default_factory, help): # noqa: A002
self.name = name
self.value_spec = value_spec
self.default = default
self.default_factory = default_factory
self.help = help
Kind = typing.Literal[
"u8",
"i8",
"u16",
"i16",
"u32",
"i32",
"u64",
"i64",
"f32",
"f64",
"bool",
"raw",
"bytes",
"str",
"cstr",
"bits",
"flags",
"struct",
"array",
"switch",
"cond",
"digest",
]
"""Every kind string rustruct.compile()'s Rust side actually recognizes
(crates/rustruct-py/src/lib.rs's parse_type match) -- a closed set, so
`Field.kind` can be a Literal instead of a bare `str` and a typo shows up
at type-check time instead of a runtime SchemaError."""
class Field(typing.NamedTuple):
"""One entry of `resolved.fields_tuple`, matching rustruct.compile()'s
expected shape (core.pyi's `Field`) -- a real NamedTuple so a call site
reads as `Field(name=..., kind=..., opts=...)`. Still a plain tuple at
the C level (pyo3's downcast is structural), so the core side needs no
changes to accept it."""
name: str | None
kind: Kind
opts: dict
def build_wire_fields(cls, namespace, bases):
# inspect.get_annotations(), not namespace["__annotations__"]: PEP 649
# (default since 3.14) computes annotations lazily via a hidden
# __annotate__, absent from the class-body namespace at __new__ time.
# get_annotations() handles both eager (<3.14) and lazy (3.14+) cases
# uniformly, and returns only this class's own annotations.
own = {}
for name, annotation in inspect.get_annotations(cls).items():
if typing.get_origin(annotation) is typing.ClassVar:
continue
if name.startswith("__") and name.endswith("__"):
continue
raw = namespace.get(name, MISSING)
if isinstance(raw, FieldSpec):
spec = raw
value_spec = annotation if spec.kind is None else spec
default, default_factory, help_ = spec.default, spec.default_factory, spec.help
else:
value_spec = annotation
default = raw
default_factory = MISSING
help_ = None
own[name] = FieldDecl(name, value_spec, default, default_factory, help_)
# Each base already carries its own fully-accumulated wire_fields (this
# same merge ran when it was built), so a flat walk of direct bases --
# not the whole MRO -- is enough. A name inherited from a base and then
# redeclared here overrides that base's FieldDecl but keeps its
# original wire position, exactly like dataclass field inheritance.
fields = {}
for base in bases:
for f in getattr(base, "wire_fields", ()):
fields[f.name] = f
fields.update(own)
return tuple(fields.values())
class Resolved:
"""A class's own compiled schema, cached once on `cls.resolved_cache`
(see `ensure_resolved`) and never touched again -- every container
here is a read-only snapshot (`MappingProxyType`/`frozenset`/`tuple`),
not a live mutable structure some other code could reach through the
class and corrupt out from under every future pack()/unpack() call."""
__slots__ = ("fields_tuple", "shapes", "no_input_names", "needs_ctx")
def __init__(self, fields_tuple, shapes, no_input_names, needs_ctx):
self.fields_tuple = fields_tuple
self.shapes = MappingProxyType(dict(shapes))
self.no_input_names = frozenset(no_input_names)
self.needs_ctx = needs_ctx
def shape_needs_ctx(shape):
"""Whether converting a value through `shape` could ever need to look
past its immediate frame (only a bare-ref SwitchShape does, or something
nested that does). Lets the common switch-free case skip building the
ctx tuple, avoiding a per-call allocation."""
if isinstance(shape, SwitchShape):
return shape.on_field is not None
if isinstance(shape, StructShape):
return ensure_resolved(shape.cls).needs_ctx
if isinstance(shape, ArrayShape):
return shape_needs_ctx(shape.elem)
return False
def collect_refs(value, into):
"""Walk an Expr-shaped tuple ("*"/int/("ref", name)/("add", a, b)/...)
and collect every ("ref", name) target it mentions."""
if isinstance(value, tuple):
if len(value) == 2 and value[0] == "ref" and isinstance(value[1], str):
into.add(value[1])
else:
for item in value:
collect_refs(item, into)
def ast_load(name):
return ast.Name(id=name, ctx=ast.Load())
def ast_call(func, args):
return ast.Call(func=func, args=args, keywords=[])
def ast_bind(namespace, prefix, value):
"""Put `value` into the compiled function's globals under a fresh name
and return a Name node loading it -- how generated code references
runtime objects (shapes, defaults, other classes' compiled methods)."""
name = f"{prefix}{len(namespace)}"
namespace[name] = value
return ast_load(name)
def ast_subscript(owner, key):
return ast.Subscript(value=ast_load(owner), slice=ast.Constant(value=key), ctx=ast.Load())
def ast_subscript_store(owner, key):
return ast.Subscript(value=ast_load(owner), slice=ast.Constant(value=key), ctx=ast.Store())
def ast_self_dict(ctx):
return ast.Attribute(value=ast_load("self"), attr="__dict__", ctx=ctx)
def ast_ctx_push(frame):
"""ctx = (*ctx, <frame>)"""
return ast.Assign(
targets=[ast.Name(id="ctx", ctx=ast.Store())],
value=ast.Tuple(
elts=[ast.Starred(value=ast_load("ctx"), ctx=ast.Load()), ast_load(frame)],
ctx=ast.Load(),
),
)
def ast_signature(posargs, kwonlyargs=(), kw_defaults=(), defaults=()):
return ast.arguments(
posonlyargs=[],
args=[ast.arg(arg=a) for a in posargs],
vararg=None,
kwonlyargs=[ast.arg(arg=a) for a in kwonlyargs],
kw_defaults=list(kw_defaults),
kwarg=None,
defaults=list(defaults),
)
def ast_function(name, args, body, namespace):
"""Compile a single FunctionDef built as an AST (never via source-string
templating) and return the resulting function, with `namespace` as its
globals."""
fn = ast.FunctionDef(name=name, args=args, body=body, decorator_list=[])
# type_params: a FunctionDef field from PEP 695 (3.12+), required by
# compile() on 3.12+ and simply unused before that. setattr(), since
# typeshed's FunctionDef stub only declares this field for a 3.12+
# target and this project's floor is 3.11.
setattr(fn, "type_params", []) # noqa: B010
module = ast.Module(body=[fn], type_ignores=[])
ast.fix_missing_locations(module)
exec(compile(module, f"<rustruct:{name}>", "exec"), namespace)
return namespace[name]
def compile_init(fields, no_input_names):
"""Generate a real __init__ with a genuine keyword-only signature, so
CPython's own call-binding machinery (fast, C-level) handles required-
vs-optional and rejects unknown keywords, instead of a Python-level loop
popping from **kwargs and branching on default/default_factory/
no_input_names per field per call."""
if not fields:
return ast_function("__init__", ast_signature(["self"]), [ast.Pass()], {})
namespace = {"_MISSING": MISSING}
kwonly, kw_defaults, body = [], [], []
for f in fields:
kwonly.append(f.name)
if f.default_factory is not MISSING:
factory = ast_bind(namespace, "_factory", f.default_factory)
kw_defaults.append(ast_load("_MISSING"))
value = ast.IfExp(
test=ast.Compare(left=ast_load(f.name), ops=[ast.Is()], comparators=[ast_load("_MISSING")]),
body=ast_call(factory, []),
orelse=ast_load(f.name),
)
elif f.default is not MISSING:
kw_defaults.append(ast_bind(namespace, "_default", f.default))
value = ast_load(f.name)
elif f.name in no_input_names:
# derived/const: pack() always recomputes/overwrites this, so
# there's nothing meaningful to require here.
kw_defaults.append(ast.Constant(value=None))
value = ast_load(f.name)
else:
kw_defaults.append(None) # keyword-only with no default: required
value = ast_load(f.name)
body.append(
ast.Assign(
targets=[ast.Attribute(value=ast_load("self"), attr=f.name, ctx=ast.Store())],
value=value,
)
)
args = ast_signature(["self"], kwonlyargs=kwonly, kw_defaults=kw_defaults)
return ast_function("__init__", args, body, namespace)
def emit_to_python(shape, expr, namespace, depth=0):
"""Return an AST expression converting wire-value `expr` to its Python
value, inlining the common shape combinations so the per-element hot
path (arrays of nested structs, arrays of converts) runs with zero
Shape-object dispatch. Falls back to a generic `.to_python()` call for
anything else (switches). `expr` is a zero-arg factory producing a
fresh node per use -- AST nodes must not be shared between positions."""
match shape:
case _ if shape is SCALAR_SHAPE:
return expr()
case StructShape(cls=cls):
# The target class is fixed by the schema (no polymorphism on
# unpack) and is already resolved -- ensure_resolved() ran on it
# while this class's own fields were being resolved -- so this
# binds its *compiled* from_mapping directly.
fm = ast_bind(namespace, "_fm", cls.from_mapping)
return ast_call(fm, [expr(), ast_load("ctx")])
case ConvertShape(decode=decode):
return ast_call(ast_bind(namespace, "_dec", decode), [expr()])
case ArrayShape(elem=elem):
var = f"v{depth}"
inner = emit_to_python(elem, lambda: ast_load(var), namespace, depth + 1)
return ast.ListComp(
elt=inner,
generators=[
ast.comprehension(target=ast.Name(id=var, ctx=ast.Store()), iter=expr(), ifs=[], is_async=0)
],
)
case _:
sh = ast_bind(namespace, "_sh", shape)
return ast_call(ast.Attribute(value=sh, attr="to_python", ctx=ast.Load()), [expr(), ast_load("ctx")])
def emit_to_wire(shape, expr, namespace, depth=0):
"""The pack-direction twin of emit_to_python. Nested struct values
dispatch through the value's own to_mapping (not a schema-bound
function): the runtime value may be a *subclass* of the declared elem
class, and raw dicts are passed through unconverted, exactly as
StructShape.to_wire always did."""
match shape:
case _ if shape is SCALAR_SHAPE:
return expr()
case StructShape():
# <expr>.to_mapping(ctx) if isinstance(<expr>, Struct) else <expr>
return ast.IfExp(
test=ast_call(ast_load("isinstance"), [expr(), ast_load("_Struct")]),
body=ast_call(ast.Attribute(value=expr(), attr="to_mapping", ctx=ast.Load()), [ast_load("ctx")]),
orelse=expr(),
)
case ConvertShape(encode=encode):
return ast_call(ast_bind(namespace, "_enc", encode), [expr()])
case ArrayShape(elem=elem):
var = f"v{depth}"
inner = emit_to_wire(elem, lambda: ast_load(var), namespace, depth + 1)
return ast.ListComp(
elt=inner,
generators=[
ast.comprehension(target=ast.Name(id=var, ctx=ast.Store()), iter=expr(), ifs=[], is_async=0)
],
)
case _:
sh = ast_bind(namespace, "_sh", shape)
return ast_call(ast.Attribute(value=sh, attr="to_wire", ctx=ast.Load()), [expr(), ast_load("ctx")])
def compile_to_mapping(fields, no_input_names, shapes, needs_ctx, cond_names):
namespace = {"_Struct": Struct}
args = ast_signature(["self", "ctx"], defaults=[ast.Tuple(elts=[], ctx=ast.Load())])
if not no_input_names and not cond_names and all(shapes[f.name] is SCALAR_SHAPE for f in fields):
# All-scalar, nothing omitted: one C-level dict copy of the
# instance dict beats a per-key literal (the core ignores any
# stray extra keys on pack, verified).
body = [ast.Return(value=ast_call(ast_load("dict"), [ast_self_dict(ast.Load())]))]
return ast_function("to_mapping", args, body, namespace)
body = []
if needs_ctx:
# Skipping the ctx append when nothing below could ever look past
# this frame avoids a tuple allocation on every single call -- the
# common case for e.g. an array of plain-scalar structs (see
# shape_needs_ctx).
body.append(ast_ctx_push("self"))
body.append(ast.Assign(targets=[ast.Name(id="d", ctx=ast.Store())], value=ast_self_dict(ast.Load())))
if not cond_names:
# No conditionally-present field: still one dict literal, just
# built from statement-free key/value lists (the common case).
keys, values = [], []
for f in fields:
if f.name in no_input_names:
# Derived/const/digest: pack() always recomputes/overwrites
# it, so it is omitted.
continue
keys.append(ast.Constant(value=f.name))
values.append(emit_to_wire(shapes[f.name], lambda name=f.name: ast_subscript("d", name), namespace))
body.append(ast.Return(value=ast.Dict(keys=keys, values=values)))
return ast_function("to_mapping", args, body, namespace)
# At least one `when()` field: its Python value is None exactly when
# wire-absent (a Shape like ArrayShape/StructShape would crash trying to
# convert a bare None), so its key is only added when not None.
body.append(ast.Assign(targets=[ast.Name(id="out", ctx=ast.Store())], value=ast.Dict(keys=[], values=[])))
for f in fields:
if f.name in no_input_names:
continue
value_expr = emit_to_wire(shapes[f.name], lambda name=f.name: ast_subscript("d", name), namespace)
assign = ast.Assign(targets=[ast_subscript_store("out", f.name)], value=value_expr)
match f.name in cond_names:
case True:
body.append(
ast.If(
test=ast.Compare(
left=ast_subscript("d", f.name), ops=[ast.IsNot()], comparators=[ast.Constant(value=None)]
),
body=[assign],
orelse=[],
)
)
case False:
body.append(assign)
body.append(ast.Return(value=ast_load("out")))
return ast_function("to_mapping", args, body, namespace)
def compile_from_mapping(cls, fields, shapes, needs_ctx, cond_names):
"""The mapping is expected to be complete for every non-`when()` field;
a missing key raises KeyError rather than falling back to a default.
Built via object.__new__ plus a direct __dict__ assignment, since
__init__'s parameter binding is pure overhead when every value is
already known by name. A `when()` field is the one exception: the core
omits its key when the predicate was false, so it gets its own presence
check and default."""
namespace = {"_cls": cls, "_new": object.__new__}
args = ast_signature(["mapping", "ctx"], defaults=[ast.Tuple(elts=[], ctx=ast.Load())])
new_self = ast.Assign(
targets=[ast.Name(id="self", ctx=ast.Store())],
value=ast_call(ast_load("_new"), [ast_load("_cls")]),
)
if not cond_names and all(shapes[f.name] is SCALAR_SHAPE for f in fields):
# All-scalar, nothing conditionally absent: adopt the dict the core
# just built as the instance dict outright. Safe because the core
# hands back a fresh dict per struct; unsound once a field can be
# missing on purpose, hence the cond_names branch below.
body = [
new_self,
ast.Assign(targets=[ast_self_dict(ast.Store())], value=ast_load("mapping")),
ast.Return(value=ast_load("self")),
]
return ast_function("from_mapping", args, body, namespace)
body = []
if needs_ctx:
body.append(ast_ctx_push("mapping"))
body.append(new_self)
if not cond_names:
keys, values = [], []
for f in fields:
keys.append(ast.Constant(value=f.name))
values.append(
emit_to_python(shapes[f.name], lambda name=f.name: ast_subscript("mapping", name), namespace)
)
body.append(ast.Assign(targets=[ast_self_dict(ast.Store())], value=ast.Dict(keys=keys, values=values)))
body.append(ast.Return(value=ast_load("self")))
return ast_function("from_mapping", args, body, namespace)
# At least one `when()` field: build the instance dict incrementally so
# each one can fall back to its own default when the core's mapping
# doesn't have that key at all (predicate was false at pack time).
body.append(ast.Assign(targets=[ast.Name(id="d", ctx=ast.Store())], value=ast.Dict(keys=[], values=[])))
for f in fields:
value_expr = emit_to_python(shapes[f.name], lambda name=f.name: ast_subscript("mapping", name), namespace)
present = ast.Assign(targets=[ast_subscript_store("d", f.name)], value=value_expr)
if f.name in cond_names:
default_node = ast_bind(namespace, "_default", f.default)
absent = ast.Assign(targets=[ast_subscript_store("d", f.name)], value=default_node)
body.append(
ast.If(
test=ast.Compare(
left=ast.Constant(value=f.name), ops=[ast.In()], comparators=[ast_load("mapping")]
),
body=[present],
orelse=[absent],
)
)
else:
body.append(present)
body.append(ast.Assign(targets=[ast_self_dict(ast.Store())], value=ast_load("d")))
body.append(ast.Return(value=ast_load("self")))
return ast_function("from_mapping", args, body, namespace)
def ensure_resolved(cls):
resolved = cls.__dict__.get("resolved_cache")
if resolved is not None:
return resolved
fields = cls.wire_fields
fields_tuple = []
shapes = {}
referenced = set()
const_names = set()
digest_names = set()
cond_names = set()
for f in fields:
kind, opts, shape = type_spec(f.value_spec)
fields_tuple.append(Field(f.name, kind, opts))
shapes[f.name] = shape
if "const" in opts:
const_names.add(f.name)
match kind:
case "digest":
digest_names.add(f.name)
case "cond":
cond_names.add(f.name)
case "bytes" | "str":
collect_refs(opts.get("len"), referenced)
case "array":
collect_refs(opts.get("count"), referenced)
case "struct":
collect_refs(opts.get("size"), referenced)
# switch's "on" is deliberately NOT collected here: unlike len/count/
# size, referencing a field from `on=` does not make it derived --
# the discriminant stays an ordinary, caller-supplied field.
own_names = {f.name for f in fields}
# Fields the caller never has to supply: derived (a sibling's len/
# count/size references them), const (pack() always writes the literal
# value), or digest (pack() always recomputes it) -- pack() ignores
# whatever's given for any of these.
no_input_names = (referenced & own_names) | const_names | digest_names
# Whether this class's own fields ever need the ancestor-scope ctx chain
# -- computed after the shapes above are all built, so any nested
# struct's own needs_ctx (already resolved, recursively) is available.
needs_ctx = any(shape_needs_ctx(s) for s in shapes.values())
resolved = Resolved(tuple(fields_tuple), shapes, no_input_names, needs_ctx)
cls.resolved_cache = resolved
# Compile specialized methods for this *exact* class and install them
# directly into cls.__dict__, replacing the per-class trampolines
# StructMeta.__new__ installed at creation time. Every call after this
# one, for this class, skips the wire_fields/Shape interpretive loop
# entirely -- see the module docstring for why this can't just be
# inherited from an ancestor.
cls.__init__ = compile_init(fields, no_input_names)
cls.to_mapping = compile_to_mapping(fields, no_input_names, shapes, needs_ctx, cond_names)
cls.from_mapping = staticmethod(compile_from_mapping(cls, fields, shapes, needs_ctx, cond_names))
return resolved
def get_codec(cls):
codec = cls.__dict__.get("codec_cache")
if codec is None:
resolved = ensure_resolved(cls)
codec = compile_codec(resolved.fields_tuple, byteorder=cls.byteorder)
cls.codec_cache = codec
return codec
def make_trampolines(cls):
"""A per-class __init__/to_mapping/from_mapping that resolves+compiles
(idempotent, cached) on first call and immediately delegates this same
call to the freshly-compiled version. Installed directly into every
class's OWN __dict__ at creation time -- see the module docstring."""
def trampoline_init(self, **kwargs):
ensure_resolved(cls)
return cls.__dict__["__init__"](self, **kwargs)
def trampoline_to_mapping(self, ctx=()):
ensure_resolved(cls)
return cls.__dict__["to_mapping"](self, ctx)
def trampoline_from_mapping(mapping, ctx=()):
ensure_resolved(cls)
return cls.__dict__["from_mapping"](mapping, ctx)
return trampoline_init, trampoline_to_mapping, trampoline_from_mapping
class StructMeta(type):
def __new__(mcls, name, bases, namespace, *, byteorder=None, registry=False, **kwargs):
# A metaclass's __new__ can't express "every class this builds is a
# Struct subclass" in its own signature -- cast once so every
# attribute set below resolves against Struct's shape, not `type`.
cls = cast("type[Struct]", super().__new__(mcls, name, bases, namespace))
if byteorder is not None:
cls.byteorder = byteorder
if registry:
cls.dispatch_registry = Registry()
cls.registry_key = None
elif kwargs:
if len(kwargs) != 1:
raise TypeError(f"{name}: expected exactly one registration keyword, got {sorted(kwargs)}")
((_, key),) = kwargs.items()
target = getattr(cls, "dispatch_registry", None)
if target is None:
raise TypeError(f"{name}: no registry found on any base class (mark the base with registry=True)")
target.add(key, cls)
cls.registry_key = key
cls.wire_fields = build_wire_fields(cls, namespace, bases)
cls.resolved_cache = None
cls.codec_cache = None
init, to_mapping, from_mapping = make_trampolines(cls)
cls.__init__ = init
cls.to_mapping = to_mapping
cls.from_mapping = staticmethod(from_mapping)
return cls
[docs]
class Struct(metaclass=StructMeta):
byteorder: ClassVar[str] = "big"
dispatch_registry: ClassVar[Registry | None] = None
registry_key: ClassVar[Any] = None
if TYPE_CHECKING:
# StructMeta.__new__ sets these three unconditionally on every class
# it builds. Declaring them here never actually executes -- it just
# gives type checkers a real shape to check subclass constructor
# calls and pack()/unpack() usage against.
wire_fields: ClassVar[tuple[FieldDecl, ...]] = ()
resolved_cache: ClassVar["Resolved | None"] = None
codec_cache: ClassVar[Any] = None
def __init__(self, **kwargs: Any) -> None: ...
def to_mapping(self, ctx: tuple = ()) -> dict[str, Any]: ...
# ClassVar[FromMapping], not `@staticmethod def`: it's assigned a
# real `staticmethod(...)` object at runtime (see StructMeta.__new__
# and ensure_resolved), so its declared type must match that value
# rather than the auto-unwrapped shape a literal def would get.
from_mapping: ClassVar[FromMapping]
def __repr__(self):
parts = ", ".join(f"{f.name}={getattr(self, f.name)!r}" for f in type(self).wire_fields)
return f"{type(self).__name__}({parts})"
def __eq__(self, other):
if type(self) is not type(other):
return NotImplemented
return all(getattr(self, f.name) == getattr(other, f.name) for f in type(self).wire_fields)
__hash__ = None
[docs]
def pack(self):
return get_codec(type(self)).pack(self.to_mapping())
[docs]
def pack_into(self, buf, offset=0):
return get_codec(type(self)).pack_into(buf, offset, self.to_mapping())
[docs]
@classmethod
def unpack(cls, buf):
return cls.from_mapping(get_codec(cls).unpack(buf))
[docs]
@classmethod
def unpack_from(cls, buf, offset=0):
mapping, pos = get_codec(cls).unpack_from(buf, offset)
return cls.from_mapping(mapping), pos
[docs]
@classmethod
def parse(cls, buf, offset=0):
result = get_codec(cls).parse(buf, offset)
if not result:
return result
mapping, pos = result
return cls.from_mapping(mapping), pos