from __future__ import annotations
from dataclasses import fields, is_dataclass
from typing import Any, Callable, ClassVar, Type, TypeVar, get_args, get_origin
from .cdr import XCDR1_MAX_ALIGNMENT, XCDR2_MAX_ALIGNMENT, CdrReader, CdrWriter
T = TypeVar("T")
class _IdlKind:
__slots__ = ("name", "write", "read", "is_primitive")
def __init__(
self,
name: str,
write: Callable[[CdrWriter, Any], None],
read: Callable[[CdrReader], Any],
is_primitive: bool = False,
) -> None:
self.name = name
self.write = write
self.read = read
self.is_primitive = is_primitive
def __repr__(self) -> str:
return f"IdlKind({self.name})"
Bool = _IdlKind("bool", CdrWriter.write_bool, CdrReader.read_bool, True)
Int8 = _IdlKind("int8", CdrWriter.write_i8, CdrReader.read_i8, True)
UInt8 = _IdlKind("uint8", CdrWriter.write_u8, CdrReader.read_u8, True)
Int16 = _IdlKind("int16", CdrWriter.write_i16, CdrReader.read_i16, True)
UInt16 = _IdlKind("uint16", CdrWriter.write_u16, CdrReader.read_u16, True)
Int32 = _IdlKind("int32", CdrWriter.write_i32, CdrReader.read_i32, True)
UInt32 = _IdlKind("uint32", CdrWriter.write_u32, CdrReader.read_u32, True)
Int64 = _IdlKind("int64", CdrWriter.write_i64, CdrReader.read_i64, True)
UInt64 = _IdlKind("uint64", CdrWriter.write_u64, CdrReader.read_u64, True)
Float32 = _IdlKind("float32", CdrWriter.write_f32, CdrReader.read_f32, True)
Float64 = _IdlKind("float64", CdrWriter.write_f64, CdrReader.read_f64, True)
String = _IdlKind("string", CdrWriter.write_string, CdrReader.read_string)
Bytes = _IdlKind("bytes", CdrWriter.write_bytes, CdrReader.read_bytes)
Octet = _IdlKind("octet", CdrWriter.write_u8, CdrReader.read_u8, True)
Char = _IdlKind("char", CdrWriter.write_char, CdrReader.read_char, True)
WChar = _IdlKind("wchar", CdrWriter.write_u16, CdrReader.read_u16, True)
WString = _IdlKind("wstring", CdrWriter.write_wstring, CdrReader.read_wstring)
class _IdlBoundedString(_IdlKind):
__slots__ = ("wide", "bound")
def __init__(self, bound: int, *, wide: bool = False) -> None:
self.wide = wide
self.bound = int(bound)
if self.bound < 0:
raise ValueError(f"string bound must be >= 0, got {self.bound}")
self.name = f"{'wstring' if wide else 'string'}<{self.bound}>"
self.write = self._write self.read = self._read
def _check(self, s: str, where: str) -> None:
if len(s) > self.bound:
raise ValueError(
f"bounded {self.name}: {where} length {len(s)} exceeds bound "
f"{self.bound}",
)
def _write(self, w: CdrWriter, s: str) -> None:
self._check(s, "value")
if self.wide:
w.write_wstring(s)
else:
w.write_string(s)
def _read(self, r: CdrReader) -> str:
s = r.read_wstring() if self.wide else r.read_string()
self._check(s, "wire")
return s
def __class_getitem__(cls, bound: Any) -> "_IdlBoundedString":
return cls(int(bound), wide=False)
class _IdlBoundedWString(_IdlBoundedString):
__slots__ = ()
def __class_getitem__(cls, bound: Any) -> "_IdlBoundedString":
return _IdlBoundedString(int(bound), wide=True)
class _IdlFixed(_IdlKind):
__slots__ = ("p", "s")
def __init__(self, p: int, s: int) -> None:
self.p = int(p)
self.s = int(s)
self.name = f"fixed<{self.p},{self.s}>"
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, v: Any) -> None:
w.write_fixed_bcd(str(v), self.p, self.s)
def _read(self, r: CdrReader) -> str:
return r.read_fixed_bcd(self.p, self.s)
def __class_getitem__(cls, params: Any) -> "_IdlFixed":
p, s = params
return cls(int(p), int(s))
class _IdlSequence(_IdlKind):
__slots__ = ("inner", "bound")
def __init__(self, inner: Any, bound: int | None = None) -> None:
self.inner = inner
self.bound = int(bound) if bound is not None else None
if self.bound is not None and self.bound < 0:
raise ValueError(f"sequence bound must be >= 0, got {self.bound}")
suffix = "" if self.bound is None else f", {self.bound}"
self.name = f"sequence<{_describe(inner)}{suffix}>"
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, values: Any) -> None:
values = list(values or [])
if self.bound is not None and len(values) > self.bound:
raise ValueError(
f"bounded sequence<{_describe(self.inner)}, {self.bound}>: "
f"got {len(values)} elements, exceeds bound {self.bound}",
)
def body(bw: CdrWriter) -> None:
bw.write_u32(len(values))
for v in values:
_write_any(bw, self.inner, v)
if _is_primitive_kind(self.inner):
body(w)
else:
_frame_dheader_write(w, body)
def _read(self, r: CdrReader) -> list:
if not _is_primitive_kind(self.inner):
r = r.read_dheader()
n = r.read_u32()
if self.bound is not None and n > self.bound:
raise ValueError(
f"bounded sequence<{_describe(self.inner)}, {self.bound}>: "
f"wire count {n} exceeds bound {self.bound}",
)
return [_read_any(r, self.inner) for _ in range(n)]
def __class_getitem__(cls, args: Any) -> "_IdlSequence":
if isinstance(args, tuple):
if len(args) != 2:
raise TypeError(
"Sequence[T] or Sequence[T, N] requires one or two parameters",
)
inner, bound = args
return cls(inner, int(bound))
return cls(args)
class _IdlArray(_IdlKind):
__slots__ = ("inner", "count")
def __init__(self, inner: Any, count: int) -> None:
if count <= 0:
raise ValueError(f"array count must be > 0, got {count}")
self.inner = inner
self.count = count
self.name = f"array<{_describe(inner)}, {count}>"
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, values: Any) -> None:
values = list(values or [])
if len(values) != self.count:
raise ValueError(
f"Array[{self.count}]: expected exactly {self.count} elements, "
f"got {len(values)}",
)
def body(bw: CdrWriter) -> None:
for v in values:
_write_any(bw, self.inner, v)
if _array_leaf_is_primitive(self):
body(w)
else:
_frame_dheader_write(w, body)
def _read(self, r: CdrReader) -> list:
if not _array_leaf_is_primitive(self):
r = r.read_dheader()
return [_read_any(r, self.inner) for _ in range(self.count)]
def __class_getitem__(cls, args: Any) -> "_IdlArray":
if not isinstance(args, tuple) or len(args) != 2:
raise TypeError("Array[T, N] requires exactly two parameters")
inner, count = args
return cls(inner, int(count))
class _IdlOptional(_IdlKind):
__slots__ = ("inner",)
def __init__(self, inner: Any) -> None:
self.inner = inner
self.name = f"optional<{_describe(inner)}>"
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, value: Any) -> None:
if value is None:
w.write_u8(0)
return
w.write_u8(1)
_write_any(w, self.inner, value)
def _read(self, r: CdrReader) -> Any:
flag = r.read_u8()
if flag == 0:
return None
return _read_any(r, self.inner)
def __class_getitem__(cls, inner: Any) -> "_IdlOptional":
return cls(inner)
class _IdlMap(_IdlKind):
__slots__ = ("key", "value", "bound")
def __init__(self, key: Any, value: Any, bound: int | None = None) -> None:
self.key = key
self.value = value
self.bound = int(bound) if bound is not None else None
if self.bound is not None and self.bound < 0:
raise ValueError(f"map bound must be >= 0, got {self.bound}")
suffix = "" if self.bound is None else f", {self.bound}"
self.name = f"map<{_describe(key)}, {_describe(value)}{suffix}>"
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, mapping: Any) -> None:
items = list((mapping or {}).items())
if self.bound is not None and len(items) > self.bound:
raise ValueError(
f"bounded map<..., {self.bound}>: got {len(items)} entries, "
f"exceeds bound {self.bound}",
)
items.sort(key=lambda kv: kv[0])
def body(bw: CdrWriter) -> None:
bw.write_u32(len(items))
for k, v in items:
_write_any(bw, self.key, k)
_write_any(bw, self.value, v)
if _is_primitive_kind(self.key) and _is_primitive_kind(self.value):
body(w)
else:
_frame_dheader_write(w, body)
def _read(self, r: CdrReader) -> dict:
if not (_is_primitive_kind(self.key) and _is_primitive_kind(self.value)):
r = r.read_dheader()
n = r.read_u32()
if self.bound is not None and n > self.bound:
raise ValueError(
f"bounded map<..., {self.bound}>: wire count {n} exceeds "
f"bound {self.bound}",
)
out: dict = {}
for _ in range(n):
k = _read_any(r, self.key)
v = _read_any(r, self.value)
out[k] = v
return out
def __class_getitem__(cls, args: Any) -> "_IdlMap":
if not isinstance(args, tuple) or len(args) not in (2, 3):
raise TypeError("Map[K, V] or Map[K, V, N] requires two or three parameters")
if len(args) == 3:
key, value, bound = args
return cls(key, value, int(bound))
key, value = args
return cls(key, value)
class _IdlEnum(_IdlKind):
__slots__ = ("enum_cls", "_w", "_r")
def __init__(self, enum_cls: type) -> None:
self.enum_cls = enum_cls
self.name = f"enum<{enum_cls.__name__}>"
bound = int(getattr(enum_cls, "_idl_bit_bound", 32))
if bound <= 8:
self._w, self._r = CdrWriter.write_i8, CdrReader.read_i8
elif bound <= 16:
self._w, self._r = CdrWriter.write_i16, CdrReader.read_i16
else:
self._w, self._r = CdrWriter.write_i32, CdrReader.read_i32
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, value: Any) -> None:
if value is None:
raise ValueError(f"Enum {self.enum_cls.__name__} must not be None")
self._w(w, int(value))
def _read(self, r: CdrReader) -> Any:
raw = self._r(r)
return self.enum_cls(raw)
def _holder_writer_reader(bits: int) -> tuple[Callable, Callable]:
if bits <= 8:
return (CdrWriter.write_u8, CdrReader.read_u8)
if bits <= 16:
return (CdrWriter.write_u16, CdrReader.read_u16)
if bits <= 32:
return (CdrWriter.write_u32, CdrReader.read_u32)
return (CdrWriter.write_u64, CdrReader.read_u64)
class _IdlBitmask(_IdlKind):
__slots__ = ("flag_cls", "_w", "_r")
def __init__(self, flag_cls: type, bit_count: int) -> None:
self.flag_cls = flag_cls
self.name = f"bitmask<{flag_cls.__name__}>"
self._w, self._r = _holder_writer_reader(bit_count)
self.is_primitive = True self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, value: Any) -> None:
if value is None:
raise ValueError(f"bitmask {self.flag_cls.__name__} must not be None")
self._w(w, int(value))
def _read(self, r: CdrReader) -> Any:
return self.flag_cls(self._r(r))
class _IdlBitset(_IdlKind):
__slots__ = ("_w", "_r")
def __init__(self, total_bits: int) -> None:
self.name = f"bitset<{total_bits}>"
self._w, self._r = _holder_writer_reader(total_bits)
self.is_primitive = True
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, value: Any) -> None:
self._w(w, int(value))
def _read(self, r: CdrReader) -> int:
return self._r(r)
class _IdlUnion(_IdlKind):
__slots__ = ("cases", "disc_kind", "default", "ext")
def __init__(
self,
disc_kind: Any,
cases: dict[int, tuple[str, Any]],
default: Any | None = None,
extensibility: str = "final",
) -> None:
self.disc_kind = _kind_from_annotation(disc_kind)
self.cases = {int(k): (v[0], v[1]) for k, v in cases.items()}
self.default = default
self.ext = extensibility
self.name = f"union<{self.disc_kind.name}>"
self.write = self._write self.read = self._read
def _resolve_case(self, disc: Any) -> tuple[str, Any] | None:
key = int(disc)
if key in self.cases:
return self.cases[key]
return self.default
def _write_body(self, w: CdrWriter, value: Any) -> None:
if value is None:
raise ValueError("union value must not be None")
disc = value.discriminator
self.disc_kind.write(w, disc)
case = self._resolve_case(disc)
if case is None:
raise ValueError(f"no case for discriminator {disc!r} and no default")
_fname, inner = case
_write_any(w, inner, value.value)
def _write(self, w: CdrWriter, value: Any) -> None:
if self.ext in ("appendable", "mutable"):
_frame_dheader_write(w, lambda bw: self._write_body(bw, value))
else:
self._write_body(w, value)
def _read_body(self, r: CdrReader) -> Any:
disc = self.disc_kind.read(r)
case = self._resolve_case(disc)
if case is None:
raise ValueError(f"no case for discriminator {disc!r} and no default")
_fname, inner = case
val = _read_any(r, inner)
return _UnionValue(discriminator=disc, value=val)
def _read(self, r: CdrReader) -> Any:
if self.ext in ("appendable", "mutable"):
return self._read_body(r.read_dheader())
return self._read_body(r)
class _UnionValue:
__slots__ = ("discriminator", "value")
def __init__(self, *, discriminator: Any, value: Any) -> None:
self.discriminator = discriminator
self.value = value
def __eq__(self, other: object) -> bool:
if not isinstance(other, _UnionValue):
return NotImplemented
return self.discriminator == other.discriminator and self.value == other.value
def __repr__(self) -> str:
return f"_UnionValue(discriminator={self.discriminator!r}, value={self.value!r})"
def idl_union(
*,
typename: str,
discriminator: Any,
cases: dict[int, tuple[str, Any]],
default: tuple[str, Any] | None = None,
extensibility: str = "final",
) -> _IdlKind:
kind = _IdlUnion(discriminator, cases, default, extensibility)
class _UnionFacade:
TYPE_NAME = typename
@staticmethod
def encode(v: Any, endian: str = "le") -> bytes:
w = CdrWriter(endian=endian)
kind.write(w, v)
return w.into_bytes()
@staticmethod
def decode(b: bytes, endian: str = "le") -> Any:
r = CdrReader(b, endian=endian)
return kind.read(r)
@staticmethod
def make(disc: Any, value: Any) -> _UnionValue:
return _UnionValue(discriminator=disc, value=value)
_idl_union_kind = kind
return _UnionFacade
def _struct_is_framed(struct_cls: type) -> bool:
ext = getattr(struct_cls, "_idl_extensibility", "final")
return ext in ("appendable", "mutable")
def _write_struct_body(w: CdrWriter, struct_cls: type, value: Any) -> None:
for fname, kind in struct_cls._idl_fields: kind.write(w, getattr(value, fname))
def _read_struct_body(r: CdrReader, struct_cls: type) -> Any:
values = {
fname: kind.read(r)
for fname, kind in struct_cls._idl_fields }
return struct_cls(**values)
class _IdlStruct(_IdlKind):
__slots__ = ("cls",)
def __init__(self, struct_cls: type) -> None:
self.cls = struct_cls
self.name = getattr(struct_cls, "TYPE_NAME", struct_cls.__name__)
self.write = self._write self.read = self._read
def _write(self, w: CdrWriter, value: Any) -> None:
if value is None:
raise ValueError(f"nested struct {self.name} must not be None")
ext = getattr(self.cls, "_idl_extensibility", "final")
if ext == "mutable":
kinds = self.cls._idl_fields mids = self.cls._idl_member_ids if w.max_alignment != XCDR2_MAX_ALIGNMENT:
_write_mutable_body_pl_cdr1(w, kinds, mids, value)
else:
_frame_dheader_write(
w, lambda bw: _write_mutable_body(bw, kinds, mids, value)
)
elif _struct_is_framed(self.cls):
_frame_dheader_write(w, lambda bw: _write_struct_body(bw, self.cls, value))
else:
_write_struct_body(w, self.cls, value)
def _read(self, r: CdrReader) -> Any:
ext = getattr(self.cls, "_idl_extensibility", "final")
if ext == "mutable":
kinds = self.cls._idl_fields mids = self.cls._idl_member_ids if r.max_alignment != XCDR2_MAX_ALIGNMENT:
return _read_mutable_body_pl_cdr1(r, kinds, mids, self.cls)
return _read_mutable_body(r.read_dheader(), kinds, mids, self.cls)
if _struct_is_framed(self.cls):
return _read_struct_body(r.read_dheader(), self.cls)
return _read_struct_body(r, self.cls)
Sequence = _IdlSequence
Array = _IdlArray
Optional = _IdlOptional
Map = _IdlMap
BoundedString = _IdlBoundedString
BoundedWString = _IdlBoundedWString
Fixed = _IdlFixed
class Bitset:
def __class_getitem__(cls, total_bits: Any) -> _IdlBitset:
return _IdlBitset(int(total_bits))
def _describe(t: Any) -> str:
if isinstance(t, _IdlKind):
return t.name
if isinstance(t, type) and is_dataclass(t):
return getattr(t, "TYPE_NAME", t.__name__)
return repr(t)
def _resolve_inner(kind: Any) -> _IdlKind | None:
if isinstance(kind, _IdlKind):
return kind
if isinstance(kind, type) and is_dataclass(kind):
return _IdlStruct(kind)
try:
return _kind_from_annotation(kind)
except TypeError:
return None
def _is_primitive_kind(kind: Any) -> bool:
resolved = _resolve_inner(kind)
return bool(resolved is not None and getattr(resolved, "is_primitive", False))
def _array_leaf_is_primitive(arr: "_IdlArray") -> bool:
inner = arr.inner
resolved = _resolve_inner(inner)
if isinstance(resolved, _IdlArray):
return _array_leaf_is_primitive(resolved)
return _is_primitive_kind(inner)
def _frame_dheader_write(w: CdrWriter, body_fn: Callable[[CdrWriter], None]) -> None:
if w.max_alignment != XCDR2_MAX_ALIGNMENT:
body_fn(w)
return
w._align(4) inner = CdrWriter(
max_alignment=w.max_alignment,
align_origin=w.position() + 4,
endian="be" if w._bo == ">" else "le",
)
body_fn(inner)
w.write_dheader(inner.into_bytes())
def _write_any(w: CdrWriter, kind: Any, value: Any) -> None:
resolved = _resolve_inner(kind)
if resolved is None:
raise TypeError(f"_write_any: unsupported kind {kind!r}")
resolved.write(w, value)
def _read_any(r: CdrReader, kind: Any) -> Any:
resolved = _resolve_inner(kind)
if resolved is None:
raise TypeError(f"_read_any: unsupported kind {kind!r}")
return resolved.read(r)
def _kind_from_annotation(annot: Any) -> _IdlKind:
import collections.abc as _collections_abc
import enum as _enum
import typing as _typing
if isinstance(annot, _IdlKind):
return annot
if isinstance(annot, _typing.ForwardRef):
import sys as _sys
name = annot.__forward_arg__
owner = getattr(annot, "__owner__", None) or getattr(annot, "owner", None)
ns: dict[str, Any] = {}
owner_mod = getattr(owner, "__module__", None)
mod = _sys.modules.get(owner_mod) if owner_mod else None
if mod is not None:
ns.update(vars(mod))
if owner is not None:
ns[getattr(owner, "__name__", name)] = owner
if name in ns:
return _kind_from_annotation(ns[name])
origin = get_origin(annot)
if origin is not None:
args = get_args(annot)
if origin in (list, _typing.List, _collections_abc.Sequence):
return _IdlSequence(_kind_from_annotation(args[0]))
if origin in (dict, _typing.Dict, _collections_abc.Mapping):
return _IdlMap(
_kind_from_annotation(args[0]),
_kind_from_annotation(args[1]),
)
if origin is _typing.Union:
non_none = [a for a in args if a is not type(None)]
if len(non_none) == 1:
return _IdlOptional(_kind_from_annotation(non_none[0]))
inner_union = getattr(annot, "_idl_union_kind", None)
if isinstance(inner_union, _IdlKind):
return inner_union
if isinstance(annot, type) and issubclass(annot, _enum.IntFlag):
bit_bound = getattr(annot, "_idl_bit_bound", 32)
return _IdlBitmask(annot, int(bit_bound))
if isinstance(annot, type) and issubclass(annot, _enum.IntEnum):
return _IdlEnum(annot)
if isinstance(annot, type) and is_dataclass(annot):
return _IdlStruct(annot)
if annot is int:
return Int32
if annot is bool:
return Bool
if annot is float:
return Float64
if annot is str:
return String
if annot is bytes:
return Bytes
raise TypeError(
f"@idl_struct: field type {annot!r} not supported. "
f"Use Bool/Int8/.../UInt64/Float32/Float64/String/Bytes, "
f"Sequence[T], Array[T, N], Optional[T], a nested @idl_struct "
f"dataclass or standard primitives (int/bool/float/str/bytes).",
)
_PRIMITIVE_WIRE_SIZE = {
"bool": 1, "int8": 1, "uint8": 1, "octet": 1, "char": 1,
"int16": 2, "uint16": 2, "wchar": 2,
"int32": 4, "uint32": 4, "float32": 4,
"int64": 8, "uint64": 8, "float64": 8,
}
_SIZE_TO_LC = {1: 0, 2: 1, 4: 2, 8: 3}
def _member_body_has_leading_dheader(kind: _IdlKind) -> bool:
name = getattr(kind, "name", "")
if name.startswith("string") or name.startswith("wstring"):
return True
if isinstance(kind, _IdlMap):
return not (_is_primitive_kind(kind.key) and _is_primitive_kind(kind.value))
if isinstance(kind, _IdlSequence):
return not _is_primitive_kind(kind.inner)
if isinstance(kind, _IdlStruct):
return _struct_is_framed(kind.cls)
return False
def _member_length_code(kind: _IdlKind) -> int:
name = getattr(kind, "name", "")
size = _PRIMITIVE_WIRE_SIZE.get(name)
if size is not None:
return _SIZE_TO_LC[size]
if _member_body_has_leading_dheader(kind):
return 5 if isinstance(kind, (_IdlBitmask, _IdlBitset)):
sub = CdrWriter()
kind.write(sub, 0)
return _SIZE_TO_LC.get(len(sub.into_bytes()), 4)
return 4
def _write_mutable_body(w: CdrWriter, kinds: list, member_ids: list, value: Any) -> None:
for (fname, kind), mid in zip(kinds, member_ids):
sub = CdrWriter(
max_alignment=w.max_alignment,
align_origin=0,
endian="be" if w._bo == ">" else "le",
)
kind.write(sub, getattr(value, fname))
w.write_emheader_lc(mid, _member_length_code(kind), sub.into_bytes())
def _write_mutable_body_pl_cdr1(
w: CdrWriter, kinds: list, member_ids: list, value: Any
) -> None:
for (fname, kind), mid in zip(kinds, member_ids):
sub = CdrWriter(
max_alignment=w.max_alignment,
align_origin=0,
endian="be" if w._bo == ">" else "le",
)
kind.write(sub, getattr(value, fname))
w.write_pl_cdr1_member(mid, sub.into_bytes())
w.write_pl_cdr1_sentinel()
def _read_mutable_body(r: CdrReader, kinds: list, member_ids: list, klass: type) -> Any:
by_id = {mid: (fname, kind) for (fname, kind), mid in zip(kinds, member_ids)}
values: dict = {}
while True:
entry = r.read_mutable_member()
if entry is None:
break
mid, must_understand, body = entry
slot = by_id.get(mid)
if slot is None:
if must_understand:
raise ValueError(f"unknown must-understand member id {mid}")
continue fname, kind = slot
values[fname] = kind.read(body)
for fname, _kind in kinds:
if fname not in values:
raise ValueError(f"missing non-optional member {fname!r}")
return klass(**values)
def _read_mutable_body_pl_cdr1(
r: CdrReader, kinds: list, member_ids: list, klass: type
) -> Any:
by_id = {mid: (fname, kind) for (fname, kind), mid in zip(kinds, member_ids)}
values: dict = {}
while True:
entry = r.read_pl_cdr1_member()
if entry is None:
break
mid, body = entry
slot = by_id.get(mid)
if slot is None:
continue fname, kind = slot
values[fname] = kind.read(body)
for fname, _kind in kinds:
if fname not in values:
raise ValueError(f"missing non-optional member {fname!r}")
return klass(**values)
def idl_struct(
*,
typename: str,
extensibility: str = "final",
member_ids: "list[int] | None" = None,
) -> Callable[[Type[T]], Type[T]]:
if extensibility not in ("final", "appendable", "mutable"):
raise ValueError(
f"@idl_struct: extensibility must be 'final', 'appendable' or "
f"'mutable', got {extensibility!r}",
)
def apply(cls: Type[T]) -> Type[T]:
if not is_dataclass(cls):
raise TypeError(
f"@idl_struct: {cls.__name__} is not a @dataclass — "
f"declaration order: @idl_struct(...) above @dataclass.",
)
import sys
import typing as _typing
module_globals: dict[str, Any] = {}
mod = sys.modules.get(cls.__module__)
if mod is not None:
module_globals.update(vars(mod))
module_globals.setdefault("Bool", Bool)
for _name, _kind in (
("Int8", Int8), ("UInt8", UInt8),
("Int16", Int16), ("UInt16", UInt16),
("Int32", Int32), ("UInt32", UInt32),
("Int64", Int64), ("UInt64", UInt64),
("Float32", Float32), ("Float64", Float64),
("String", String), ("Bytes", Bytes),
("Octet", Octet),
("Char", Char), ("WChar", WChar), ("WString", WString),
("Array", Array), ("Optional", Optional),
("Sequence", Sequence), ("Map", Map),
("BoundedString", BoundedString), ("BoundedWString", BoundedWString),
("Fixed", Fixed),
):
module_globals.setdefault(_name, _kind)
def _resolve(annot: Any) -> Any:
if isinstance(annot, _typing.ForwardRef):
annot = annot.__forward_arg__
if isinstance(annot, str):
try:
return eval(annot, module_globals) except NameError as exc:
raise TypeError(
f"@idl_struct: annotation string {annot!r} not "
f"resolvable in module {cls.__module__!r}. When using "
f"`from __future__ import annotations`, the "
f"kind constants must be imported in the module.",
) from exc
return annot
kinds: list[tuple[str, _IdlKind]] = []
for f in fields(cls):
kinds.append((f.name, _kind_from_annotation(_resolve(f.type))))
framed = extensibility in ("appendable", "mutable")
is_mutable = extensibility == "mutable"
mids = (
list(member_ids)
if member_ids is not None
else list(range(1, len(kinds) + 1))
)
if len(mids) != len(kinds):
raise ValueError(
f"@idl_struct {typename}: member_ids has {len(mids)} entries "
f"but the struct has {len(kinds)} members",
)
def _encode(self: Any, endian: str = "le", representation: int = 1) -> bytes:
max_align = XCDR1_MAX_ALIGNMENT if representation == 0 else XCDR2_MAX_ALIGNMENT
w = CdrWriter(endian=endian, max_alignment=max_align)
if is_mutable:
if representation == 0:
_write_mutable_body_pl_cdr1(w, kinds, mids, self)
else:
_frame_dheader_write(
w, lambda bw: _write_mutable_body(bw, kinds, mids, self)
)
elif framed:
_frame_dheader_write(
w, lambda bw: _write_struct_body(bw, cls, self)
)
else:
_write_struct_body(w, cls, self)
return w.into_bytes()
def _decode(
klass: Type[T],
data: bytes,
endian: str = "le",
representation: int = 1,
) -> T:
max_align = XCDR1_MAX_ALIGNMENT if representation == 0 else XCDR2_MAX_ALIGNMENT
r = CdrReader(data, endian=endian, max_alignment=max_align)
if is_mutable:
if representation == 0:
return _read_mutable_body_pl_cdr1(r, kinds, mids, klass)
return _read_mutable_body(r.read_dheader(), kinds, mids, klass)
body = r.read_dheader() if framed else r
values = {fname: kind.read(body) for fname, kind in kinds}
return klass(**values)
cls.TYPE_NAME = typename cls._idl_extensibility = extensibility cls._idl_fields = kinds cls._idl_member_ids = mids cls.encode = _encode cls.decode = classmethod(_decode) return cls
return apply
def is_idl_struct(obj: Any) -> bool:
cls: Any = obj if isinstance(obj, type) else type(obj)
return hasattr(cls, "TYPE_NAME") and hasattr(cls, "_idl_fields")
def type_name_of(cls_or_obj: Any) -> str:
cls: Any = cls_or_obj if isinstance(cls_or_obj, type) else type(cls_or_obj)
name: ClassVar[str] = getattr(cls, "TYPE_NAME", None) if name is None:
raise TypeError(f"{cls.__name__} hat keinen @idl_struct(typename=...)-Decorator")
return name