From 0c0ac5f75c2cd7f8b794c06b9aaac76cd5206be9 Mon Sep 17 00:00:00 2001 From: Justin Kim Date: Sat, 5 Sep 2026 20:22:29 +0900 Subject: [PATCH] Wrap wirelog's typed row API so float columns work in a Session wirelog 0.60.0 added float support, but a float could reach a program only through its source text. The untyped session entry points carry int64 lanes, and the engine refuses them outright where a float column is declared: wirelog_session_insert refuses the float-bearing relation, while wirelog_session_set_delta_cb and the untyped snapshot refuse whenever the program declares a float column anywhere. So insert(), step() and snapshot() do not work on such a program, and wrapping insert alone would have shipped a write with no way to read it back. Wrapped: insert_typed, remove_typed, snapshot_typed, step_typed and set_typed_delta_callback. Column types come from the relation's declaration, so callers pass plain Python rows and never build lane arrays. Left unwrapped: wirelog_session_make_compound_typed and the wirelog_extension_* family. Neither is needed for a float column, though a float inside a compound term does stay unreachable as a result, since the untyped make_compound refuses a FLOAT argument. TypedRowError subclasses ExecError so existing handlers keep working while gaining the row index, logical column and wirelog's own bounded diagnostic. The message buffer is caller-owned, which is why the struct field is POINTER(c_char) and not c_char_p - the latter surfaces as an immutable bytes and the engine's message would land where nothing can read it. TypedErrorCode is exported alongside it, since typed_code is public and naming it otherwise means importing from a module docs/api-stability.md declares private. The lane conversion lives in _core.lanes rather than on Session because the typed delta trampoline needs it too and cannot import the session module without a cycle. A FLOAT lane is host-order IEEE-754 bits, so every float crosses through an explicit reinterpret; UINT32/UINT64 decode unsigned and BOOL decodes to bool, where the untyped path returns everything signed. encode_lane accepts anything integer-like or float-like, so a NumPy scalar works as it does through insert_batch, and refuses a value that would not survive the conversion rather than storing it wrong - Decimal(2**63 - 1) names an integer no binary64 can hold. str is the deliberate exclusion: a STRING column carries an intern id and the advanced Session has no forward-intern entry point, so coercing would write a wrong id. insert_typed and remove_typed take TypedRow, not the EasySession Row alias, and the typed readers return LaneValue. Row admits str because EasySession.insert auto-interns it; declaring it here would let mypy wave through the one input that always raises and reject the NumPy scalars that work. The row count and the descriptor array size both come from one materialized list, so a Sequence whose __len__ changes between them cannot make wirelog read past the allocation. _delta_cb can now hold either callback kind, so set_delta_callback(), step() and close() dispatch on the armed kind. Reusing the other kind's handle hands ctypes a mismatched function pointer, which it rejects only after the stored callable has been overwritten. Clearing through the matching entry point is insurance rather than a fix: 0.60.0's set_delta_cb nulls typed_delta_cb unconditionally and gates its FLOAT refusal on a non-NULL callback, so either clear disarms both today - but nothing in the public C contract promises that stays true. Registration is guarded by has_typed_row_api(), mirroring wirelog_program_get_relation_ir in _parser. These symbols first ship in 0.60.0 while the loader floor still admits 0.52.0, and an unguarded attribute access would raise AttributeError at import time and make the package unimportable against an older engine. The typed methods raise WirelogVersionError there instead. Validated against both engines built from source, and mypy --strict clean over all 28 source files. Claude-Session: https://claude.ai/code/session_01MqHwvCimcvrg1osEY5MH7o --- docs/api-stability.md | 2 + src/pyrewire/__init__.py | 4 + src/pyrewire/_core/callbacks.py | 53 ++- src/pyrewire/_core/errors.py | 35 ++ src/pyrewire/_core/lanes.py | 199 ++++++++ src/pyrewire/_ffi/_advanced.py | 65 ++- src/pyrewire/_ffi/_enums.py | 14 + src/pyrewire/_ffi/_types.py | 84 ++++ src/pyrewire/session.py | 370 +++++++++++++- tests/test_typed_rows.py | 821 ++++++++++++++++++++++++++++++++ 10 files changed, 1623 insertions(+), 24 deletions(-) create mode 100644 src/pyrewire/_core/lanes.py create mode 100644 tests/test_typed_rows.py diff --git a/docs/api-stability.md b/docs/api-stability.md index 3388bd1..2d29ff0 100644 --- a/docs/api-stability.md +++ b/docs/api-stability.md @@ -66,6 +66,8 @@ The v1 stable public API includes these exported names: - `ParseError` - `InvalidIRError` - `ExecError` +- `TypedRowError` +- `TypedErrorCode` - `WirelogMemoryError` - `WirelogIOError` - `CompoundSaturatedError` diff --git a/src/pyrewire/__init__.py b/src/pyrewire/__init__.py index d7b02d3..0ae20d5 100644 --- a/src/pyrewire/__init__.py +++ b/src/pyrewire/__init__.py @@ -11,6 +11,7 @@ ExecError, InvalidIRError, ParseError, + TypedRowError, WirelogError, WirelogInternError, WirelogIOError, @@ -28,6 +29,7 @@ ErrorCode, IRNodeType, StrFn, + TypedErrorCode, ) from pyrewire._ffi._util import ( agg_fn_name, @@ -93,6 +95,8 @@ "WirelogError", "ParseError", "InvalidIRError", + "TypedRowError", + "TypedErrorCode", "ExecError", "WirelogMemoryError", "WirelogIOError", diff --git a/src/pyrewire/_core/callbacks.py b/src/pyrewire/_core/callbacks.py index a431c0a..76bf778 100644 --- a/src/pyrewire/_core/callbacks.py +++ b/src/pyrewire/_core/callbacks.py @@ -17,8 +17,8 @@ the registry slot. `CallbackHandle.drain()` then re-raises after wirelog has returned control to Python. -The two trampolines (`_delta_trampoline`, `_tuple_trampoline`) are -module-level singletons. Per-session instances would multiply ctypes +The three trampolines (`_delta_trampoline`, `_tuple_trampoline`, +`_typed_delta_trampoline`) are module-level singletons. Per-session instances would multiply ctypes overhead and complicate lifetime management. """ @@ -31,12 +31,15 @@ from dataclasses import dataclass, field from typing import Any -from .._ffi._types import OnDeltaFn, OnTupleFn +from .._ffi._types import OnDeltaFn, OnTupleFn, OnTypedTupleFn +from .lanes import decode_typed_row -# Event payloads buffered by the trampolines. The first element ("delta" -# or "tuple") is the trampoline kind; the rest is the decoded payload. +# Event payloads buffered by the trampolines. The first element is the +# trampoline kind; the rest is the decoded payload. DeltaEvent = tuple[str, str, tuple[int, ...], int] # ("delta", rel, row_ids, diff) TupleEvent = tuple[str, str, tuple[int, ...]] # ("tuple", rel, row_ids) +# ("typed_delta", rel, decoded_values, diff) - values are typed, not ids +TypedDeltaEvent = tuple[str, str, tuple[object, ...], int] Event = Any # union of the above @@ -62,7 +65,7 @@ def _next_token() -> int: # --- module-level trampolines ---------------------------------------------- # -# Both trampolines defensively swallow every exception. They store the +# Every trampoline defensively swallows all exceptions. They store the # exception on the registry slot (best effort) so `drain()` can re-raise # it. Never let one propagate into wirelog's C call. @@ -114,6 +117,34 @@ def _tuple_trampoline( state.last_error = exc +@OnTypedTupleFn # type: ignore[untyped-decorator] +def _typed_delta_trampoline( + relation: bytes, + row: Any, + diff: int, + user_data: Any, +) -> None: + """Typed counterpart of `_delta_trampoline`. + + wirelog refuses to install the UNTYPED delta callback on a program + whose schema carries a FLOAT column, so this is the only delta path + available to such a session. The descriptor is borrowed for this call + only; `decode_typed_row` copies every value out before returning. + """ + state: _TrampolineState | None = None + try: + raw = ctypes.cast(user_data, ctypes.c_void_p).value + token = int(raw) if raw else 0 + state = _REGISTRY.get(token) + if state is None or not row: + return + rel = relation.decode() if relation else "" + state.queue.append(("typed_delta", rel, decode_typed_row(row[0]), int(diff))) + except BaseException as exc: # never propagate to C + if state is not None: + state.last_error = exc + + # --- public API ------------------------------------------------------------ @@ -125,7 +156,7 @@ class CallbackHandle: __slots__ = ("kind", "token", "_state", "__weakref__") def __init__(self, kind: str, user_fn: Callable[..., Any] | None = None) -> None: - if kind not in ("delta", "tuple"): + if kind not in ("delta", "tuple", "typed_delta"): raise ValueError(f"unknown callback kind: {kind!r}") self.kind: str = kind self.token: int = _next_token() @@ -140,7 +171,11 @@ def user_data(self) -> ctypes.c_void_p: @property def fn(self) -> Any: """The module-level CFUNCTYPE instance matching `kind`.""" - return _delta_trampoline if self.kind == "delta" else _tuple_trampoline + if self.kind == "delta": + return _delta_trampoline + if self.kind == "typed_delta": + return _typed_delta_trampoline + return _tuple_trampoline def drain(self) -> list[Event]: """Pop and return all queued events. If a callback raised, the @@ -165,4 +200,4 @@ def __del__(self) -> None: pass -__all__ = ["CallbackHandle", "DeltaEvent", "TupleEvent", "Event"] +__all__ = ["CallbackHandle", "DeltaEvent", "TupleEvent", "TypedDeltaEvent", "Event"] diff --git a/src/pyrewire/_core/errors.py b/src/pyrewire/_core/errors.py index 01fe570..1de0bdc 100644 --- a/src/pyrewire/_core/errors.py +++ b/src/pyrewire/_core/errors.py @@ -110,6 +110,40 @@ class ExecError(WirelogError): code = int(ErrorCode.EXEC) +class TypedRowError(ExecError): + """Raised when `wirelog_session_insert_typed` / `remove_typed` rejects a row. + + A subclass of `ExecError` so existing `except ExecError` handlers keep + working, but it carries what the generic path throws away: which row + and column the engine objected to, and its own bounded diagnostic. + + Attributes: + typed_code: The `TypedErrorCode` wirelog reported - `DESCRIPTOR` + (the descriptor is malformed), `SCHEMA` (it disagrees with the + relation) or `VALUE` (a lane holds an unrepresentable value). + row_index: Index of the offending row within the batch, or `None` + when the engine reported no specific row. + column: Logical column index within that row, or `None` when the + engine reported no specific column. + engine_message: wirelog's own message, or `None` if it supplied none. + """ + + def __init__( + self, + message: str, + *, + typed_code: int, + row_index: int | None = None, + column: int | None = None, + engine_message: str | None = None, + ) -> None: + super().__init__(message) + self.typed_code = typed_code + self.row_index = row_index + self.column = column + self.engine_message = engine_message + + class WirelogMemoryError(WirelogError): """Raised when the wirelog engine cannot allocate memory (``WIRELOG_ERR_MEMORY``, code 4). @@ -282,6 +316,7 @@ def check(rc: int) -> None: "WirelogError", "ParseError", "InvalidIRError", + "TypedRowError", "ExecError", "WirelogMemoryError", "WirelogIOError", diff --git a/src/pyrewire/_core/lanes.py b/src/pyrewire/_core/lanes.py new file mode 100644 index 0000000..af491dc --- /dev/null +++ b/src/pyrewire/_core/lanes.py @@ -0,0 +1,199 @@ +# SPDX-License-Identifier: Apache-2.0 OR GPL-3.0-or-later +"""Conversion between Python values and wirelog's 64-bit typed lanes. + +wirelog's typed row ABI carries every column as a `uint64_t` lane plus a +`wirelog_column_type_t` saying how to read it. A FLOAT lane holds +host-order IEEE-754 binary64 bits rather than a C `double`, so a float +only survives the boundary through an explicit bit reinterpret - which is +exactly what the untyped `int64_t` entry points cannot do. + +These helpers live here rather than on `Session` because both the typed +insert path (`pyrewire.session`) and the typed delta trampoline +(`pyrewire._core.callbacks`) need them, and the trampoline cannot import +the session module without a cycle. +""" + +from __future__ import annotations + +import ctypes +import numbers +import operator +import struct +from typing import Any, SupportsFloat, SupportsIndex + +from .._ffi._enums import ColumnType + +# What a single lane can hold once decoded. BOOL decodes to `bool`, which +# is an `int` subclass, so the union is a documentation aid rather than a +# discriminator. +LaneValue = int | float | bool + +# What `encode_lane` accepts: anything integer-like or float-like. Wider +# than `LaneValue`, so a NumPy scalar or any `__index__` type goes in. +LaneInput = SupportsIndex | SupportsFloat + +_FLOAT = struct.Struct(" int: + """Pack one Python value into its 64-bit lane. + + Integer columns take the value's two's-complement bits. A float that + names an exact integer (`3.0`) is accepted there; a non-integral one + is rejected rather than silently truncated. + + Anything implementing `__index__` or `__float__` is accepted, so a + NumPy scalar works here as it does through `insert_batch`. `str` is + the deliberate exclusion. A float-like value heading for an integer + column is converted through `float`; if it is also a real numeric + type, the result is checked against the original and refused when the + conversion lost information - `Decimal(2**63 - 1)` names an integer a + binary64 cannot hold, and is refused rather than stored one off. + + Raises: + ValueError: a float-like value was given for an integer column and + it does not name an exactly representable integer. + TypeError: a `str`, or a value that is neither integer-like nor + float-like, was given. + OverflowError: the value is too large to convert to a float at + all, as `10**400` is. + """ + # Concrete types first. `isinstance` against a runtime-checkable + # Protocol probes for attributes and has no negative-result cache, so + # leaving it in front of the common case costs ~25% of a large typed + # insert. `type(...) is` deliberately excludes bool, which falls + # through to the __index__ path below and encodes as 0/1 either way. + is_int = type(value) is int + is_float = type(value) is float + if not is_int and not is_float: + if isinstance(value, str): + raise TypeError( + f"cannot encode str {value!r} into a lane: a STRING column carries an " + f"intern id, and the advanced Session has no forward-intern API. Pass " + f"the int64 id, seeding it with seed_intern(value, id) if you need to " + f"decode it back." + ) + if not isinstance(value, (SupportsIndex, SupportsFloat)): + raise TypeError(f"cannot encode {type(value).__name__} {value!r} into a lane") + + if column_type == ColumnType.FLOAT: + bits: int = _BITS.unpack(_FLOAT.pack(_as_float(value)))[0] + return bits + + if is_int: + return value & _UINT64_MASK # type: ignore[operator] + if not is_float and isinstance(value, SupportsIndex): + # `operator.index` enforces that __index__ returned an int, where + # calling the dunder directly would let a bogus return value reach + # the `&` below as some other type. + return operator.index(value) & _UINT64_MASK + + # Float-like heading for an integer column: accept it only when it + # names an exact integer, and only when converting through float did + # not already lose the value. + as_float = _as_float(value) + if not as_float.is_integer(): + raise ValueError(f"cannot store float {value!r} in a {ColumnType(column_type).name} column") + as_int = int(as_float) + if not is_float and _float_conversion_lost_value(value, as_int): + raise ValueError( + f"cannot store {type(value).__name__} {value!r} in a " + f"{ColumnType(column_type).name} column: it is not exactly representable " + f"as a 64-bit float, so the stored value would be {as_int}" + ) + return as_int & _UINT64_MASK + + +def _float_conversion_lost_value(value: object, as_int: int) -> bool: + """Whether routing `value` through `float` lost information. + + The check runs only for a type registered with `numbers.Number`. + `Decimal` and `Fraction` carry an exact value that `float` may round + - `Decimal(2**63 - 1)` becomes `2**63` - and they compare numerically + against `int`, so for them the question is meaningful and the answer + is trustworthy. + + Everything else is accepted without asking, because for most + `__float__` providers the protocol is the only thing exposed: + `float(value)` IS the value and nothing can have been lost. Asking + anyway would compare an `int` against an object with no numeric + `__eq__`, fall back to identity, and reject every such value while + claiming it was unrepresentable - which is how this guard read + before, and it was wrong. + + `numbers.Number` is a proxy for "can answer the question", not the + thing itself, and it is imprecise in both directions: an exact + numeric type that declines to register is not checked, and a + non-numeric type that registers without defining `__eq__` is checked + and wrongly refused. Gating on `type(value).__eq__` instead trades + one of those for a worse one - a wrapper whose `__eq__` returns + `False` rather than `NotImplemented` for an `int` goes back to being + falsely rejected. The documented contract is therefore scoped to real + numeric types rather than to this mechanism. + + A comparison that raises is treated as "cannot tell, accept", so a + user-defined `__eq__` cannot throw out of `encode_lane` past its + documented contract. + """ + if not isinstance(value, numbers.Number): + return False + try: + # `!=` through operator, not the literal, so mypy does not reject + # the intentional cross-type comparison the ABCs make valid. + return bool(operator.ne(as_int, value)) + except Exception: + return False + + +def _as_float(value: LaneInput) -> float: + """`float(value)`, with a non-numeric `__float__` reported as a TypeError. + + `ArithmeticError` passes through: `float(10**400)` overflows, and that + is an honest `OverflowError` about the value rather than a fault in + some `__float__` the type does not even have. + """ + try: + return float(value) + except (TypeError, ValueError, ArithmeticError): + raise + except Exception as exc: # a __float__ that raises something else + raise TypeError( + f"cannot encode {type(value).__name__} {value!r} into a lane: " + f"__float__ raised {exc!r}" + ) from exc + + +def decode_lane(lane: int, column_type: ColumnType) -> LaneValue: + """Unpack one 64-bit lane back into a Python value. + + FLOAT becomes `float`, BOOL becomes `bool`, UINT32 / UINT64 stay + unsigned, and everything else is read as a signed 64-bit integer. + """ + if column_type == ColumnType.FLOAT: + value: float = _FLOAT.unpack(_BITS.pack(lane))[0] + return value + if column_type == ColumnType.BOOL: + return bool(lane) + if column_type in (ColumnType.UINT32, ColumnType.UINT64): + return int(lane) + return int(ctypes.c_int64(lane).value) + + +def decode_typed_row(row: Any) -> tuple[LaneValue, ...]: + """Decode a borrowed `TypedRowStruct` into a tuple of Python values. + + The descriptor and its lane storage belong to wirelog only for the + duration of the callback that received them, so this copies every + value out before returning. + """ + values: list[LaneValue] = [] + for c in range(row.logical_ncols): + lane = int(row.lanes[row.lane_offsets[c]]) + values.append(decode_lane(lane, ColumnType(row.types[c]))) + return tuple(values) + + +__all__ = ["LaneValue", "LaneInput", "encode_lane", "decode_lane", "decode_typed_row"] diff --git a/src/pyrewire/_ffi/_advanced.py b/src/pyrewire/_ffi/_advanced.py index 903e05c..6dd45a7 100644 --- a/src/pyrewire/_ffi/_advanced.py +++ b/src/pyrewire/_ffi/_advanced.py @@ -1,7 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 OR GPL-3.0-or-later """Raw ctypes bindings for the wirelog advanced session API (#20). -Covers 8 entry points from `wirelog/wirelog-advanced.h`: +Covers 12 entry points from `wirelog/wirelog-advanced.h`: - `wirelog_session_create(program, backend, num_workers, &out) -> wirelog_error_t` - `wirelog_session_destroy(session) -> void` @@ -11,6 +11,10 @@ - `wirelog_session_snapshot(session, cb, user_data) -> wirelog_error_t` - `wirelog_session_set_delta_cb(session, cb, user_data) -> wirelog_error_t` - `wirelog_session_make_compound(session, functor, arity, args, &handle_out) -> wirelog_error_t` +- `wirelog_session_insert_typed(session, rel, rows, num_rows, &err) -> wirelog_error_t` +- `wirelog_session_remove_typed(session, rel, rows, num_rows, &err) -> wirelog_error_t` +- `wirelog_session_snapshot_typed(session, cb, user_data) -> wirelog_error_t` +- `wirelog_session_set_typed_delta_cb(session, cb, user_data) -> wirelog_error_t` The advanced session BORROWS its `wirelog_program_t`; the high-level `Session` class (#21) is responsible for keeping the program alive @@ -19,6 +23,13 @@ `insert` / `remove` take BATCHED rows (`num_rows * num_cols`), unlike the easy facade which takes one row per call. This is the documented advanced-API shape. + +The four `*_typed` entry points are the only way a FLOAT column +crosses the FFI boundary with its value intact. The untyped entry +points carry `int64_t` lanes, and on a relation that declares a FLOAT +column wirelog refuses them outright rather than distorting the value. +The typed descriptors are borrowed for the duration of the call and +never retained. """ from __future__ import annotations @@ -30,11 +41,31 @@ CompoundArgStruct, OnDeltaFn, OnTupleFn, + OnTypedTupleFn, ProgramHandle, SessionHandle, + TypedErrorStruct, + TypedRowStruct, +) + +TYPED_ROW_ENTRY_POINTS = ( + "wirelog_session_insert_typed", + "wirelog_session_remove_typed", + "wirelog_session_snapshot_typed", + "wirelog_session_set_typed_delta_cb", ) +def has_typed_row_api() -> bool: + """Whether the loaded libwirelog exports the typed row entry points. + + False on any engine older than 0.60.0. Callers should raise + `WirelogVersionError` rather than let ctypes fail with a bare + `AttributeError`. + """ + return all(hasattr(LIB, name) for name in TYPED_ROW_ENTRY_POINTS) + + def _register() -> None: LIB.wirelog_session_create.restype = ctypes.c_int LIB.wirelog_session_create.argtypes = [ @@ -85,5 +116,37 @@ def _register() -> None: ctypes.POINTER(ctypes.c_uint64), ] + # The typed row entry points first ship in wirelog 0.60.0, while the + # loader floor still admits 0.52.0. Register them guardedly, the same + # way `_parser` handles `wirelog_program_get_relation_ir`: an + # unguarded attribute access here would raise AttributeError at import + # time and take the whole package down on an older engine. + if has_typed_row_api(): + LIB.wirelog_session_insert_typed.restype = ctypes.c_int + LIB.wirelog_session_insert_typed.argtypes = [ + SessionHandle, + ctypes.c_char_p, + ctypes.POINTER(TypedRowStruct), + ctypes.c_uint32, + ctypes.POINTER(TypedErrorStruct), + ] + + LIB.wirelog_session_remove_typed.restype = ctypes.c_int + # `list(...)`: sharing one mutable list between two ctypes + # function objects would let a later in-place edit retype both. + LIB.wirelog_session_remove_typed.argtypes = list(LIB.wirelog_session_insert_typed.argtypes) + + LIB.wirelog_session_snapshot_typed.restype = ctypes.c_int + LIB.wirelog_session_snapshot_typed.argtypes = [ + SessionHandle, + OnTypedTupleFn, + ctypes.c_void_p, + ] + + LIB.wirelog_session_set_typed_delta_cb.restype = ctypes.c_int + LIB.wirelog_session_set_typed_delta_cb.argtypes = list( + LIB.wirelog_session_snapshot_typed.argtypes + ) + _register() diff --git a/src/pyrewire/_ffi/_enums.py b/src/pyrewire/_ffi/_enums.py index 5834287..315d3a6 100644 --- a/src/pyrewire/_ffi/_enums.py +++ b/src/pyrewire/_ffi/_enums.py @@ -37,6 +37,19 @@ class ColumnType(IntEnum): BOOL = 6 +class TypedErrorCode(IntEnum): + """Mirrors `wirelog_typed_error_code_t`. + + Reported through `wirelog_typed_error_v1_t.code` by the typed + insert / remove entry points. + """ + + NONE = 0 + DESCRIPTOR = 1 # the row descriptor itself is malformed + SCHEMA = 2 # descriptor is well-formed but disagrees with the relation + VALUE = 3 # a lane holds a value the column cannot represent + + class CompoundKind(IntEnum): """Mirrors `wirelog_compound_kind_t`.""" @@ -135,6 +148,7 @@ class IRNodeType(IntEnum): __all__ = [ "ErrorCode", + "TypedErrorCode", "ColumnType", "CompoundKind", "CmpOp", diff --git a/src/pyrewire/_ffi/_types.py b/src/pyrewire/_ffi/_types.py index 3c4b67d..18a9dc2 100644 --- a/src/pyrewire/_ffi/_types.py +++ b/src/pyrewire/_ffi/_types.py @@ -86,6 +86,71 @@ class CompoundArgStruct(ctypes.Structure): ] +# --- wirelog_typed_row_v1_t ------------------------------------------------- +# +# The versioned typed-row descriptor from wirelog-types.h. `struct_size` +# and `abi_version` are the forward-compatibility handshake: the C side +# rejects a descriptor whose `struct_size` it does not recognise, so both +# fields MUST be filled from the constants below rather than left zero. +# +# A row carries two schemas. The LOGICAL one (`types`, `logical_ncols`) is +# the relation's declared columns. The PHYSICAL one (`physical_types`, +# `physical_nlanes`, `physical_stride`) is the lane layout, which differs +# from the logical one only for a relation declaring an `inline` compound +# column. `lane_offsets[c]` maps logical column `c` to its first lane. +# +# FLOAT lanes hold host-order IEEE-754 binary64 bits, not a C double, so +# every float crosses this boundary through an explicit bit reinterpret. + + +class TypedRowStruct(ctypes.Structure): + """Mirrors `wirelog_typed_row_v1_t` from wirelog-types.h.""" + + _fields_ = [ + ("struct_size", ctypes.c_uint32), + ("abi_version", ctypes.c_uint16), + ("reserved", ctypes.c_uint16), + ("logical_ncols", ctypes.c_uint32), + ("physical_nlanes", ctypes.c_uint32), + ("physical_stride", ctypes.c_uint32), + ("types", ctypes.POINTER(ctypes.c_uint32)), + ("lane_offsets", ctypes.POINTER(ctypes.c_uint32)), + ("physical_types", ctypes.POINTER(ctypes.c_uint32)), + ("lanes", ctypes.POINTER(ctypes.c_uint64)), + ] + + +class TypedErrorStruct(ctypes.Structure): + """Mirrors `wirelog_typed_error_v1_t`. + + `message` is a CALLER-owned buffer; wirelog writes into it only when + `message_capacity` is non-zero. Pass `POINTER(c_char)` rather than + `c_char_p` so ctypes hands the C side a writeable pointer instead of + an immutable `bytes`. + """ + + _fields_ = [ + ("struct_size", ctypes.c_uint32), + ("code", ctypes.c_uint32), + ("row_index", ctypes.c_uint32), + ("logical_col", ctypes.c_uint32), + ("message", ctypes.POINTER(ctypes.c_char)), + ("message_capacity", ctypes.c_uint32), + ] + + +TYPED_ROW_STRUCT_SIZE = ctypes.sizeof(TypedRowStruct) +TYPED_ERROR_STRUCT_SIZE = ctypes.sizeof(TypedErrorStruct) +TYPED_ROW_ABI_VERSION = 1 + +# wirelog reports "not applicable" for the row / column fields of a typed +# error as UINT32_MAX rather than 0, which is a valid index. +TYPED_ERROR_NO_INDEX = 0xFFFFFFFF + +# Bytes reserved for the engine's bounded diagnostic message. +TYPED_ERROR_MESSAGE_CAPACITY = 256 + + # --- wirelog_easy_open_opts_t ----------------------------------------------- class EasyOpenOptsStruct(ctypes.Structure): """Mirrors `wirelog_easy_open_opts_t`. `size` MUST be set to @@ -125,6 +190,17 @@ class EasyOpenOptsStruct(ctypes.Structure): ctypes.c_void_p, ) +# `wirelog_on_typed_tuple_fn`. The descriptor and its lane storage are +# owned by the library only for the duration of the call, so a trampoline +# must copy out anything it wants to keep. +OnTypedTupleFn = ctypes.CFUNCTYPE( + None, + ctypes.c_char_p, # const char *relation + ctypes.POINTER(TypedRowStruct), # const wirelog_typed_row_v1_t *row + ctypes.c_int32, # int32_t diff (+1 / -1) + ctypes.c_void_p, # void *user_data +) + # --- wirelog_io_adapter_t (ABI version 2) ---------------------------------- # @@ -187,12 +263,20 @@ class IOAdapterStruct(ctypes.Structure): "SchemaStruct", "StratumStruct", "CompoundArgStruct", + "TypedRowStruct", + "TypedErrorStruct", "EasyOpenOptsStruct", "IOAdapterStruct", # Constants "EASY_OPEN_OPTS_SIZE", + "TYPED_ROW_STRUCT_SIZE", + "TYPED_ERROR_STRUCT_SIZE", + "TYPED_ROW_ABI_VERSION", + "TYPED_ERROR_NO_INDEX", + "TYPED_ERROR_MESSAGE_CAPACITY", "WIRELOG_IO_ABI_VERSION", # Callback types "OnTupleFn", "OnDeltaFn", + "OnTypedTupleFn", ] diff --git a/src/pyrewire/session.py b/src/pyrewire/session.py index 3565d35..9f52cc4 100644 --- a/src/pyrewire/session.py +++ b/src/pyrewire/session.py @@ -21,19 +21,36 @@ from typing import Any from ._core.callbacks import CallbackHandle -from ._core.errors import ExecError, WirelogInternError, WirelogModeError, check +from ._core.errors import ( + ExecError, + TypedRowError, + WirelogInternError, + WirelogModeError, + WirelogVersionError, + check, +) from ._core.intern import InternTable +from ._core.lanes import LaneInput, LaneValue, encode_lane from ._ffi import LIB from ._ffi import _advanced as _advanced_ffi # noqa: F401 -- registers argtypes from ._ffi import _easy as _easy_ffi # noqa: F401 -- registers argtypes -from ._ffi._enums import BackendKind, ColumnType +from ._ffi._advanced import has_typed_row_api +from ._ffi._enums import BackendKind, ColumnType, CompoundKind, TypedErrorCode from ._ffi._types import ( EASY_OPEN_OPTS_SIZE, + TYPED_ERROR_MESSAGE_CAPACITY, + TYPED_ERROR_NO_INDEX, + TYPED_ERROR_STRUCT_SIZE, + TYPED_ROW_ABI_VERSION, + TYPED_ROW_STRUCT_SIZE, CompoundArgStruct, EasyOpenOptsStruct, EasySessionHandle, OnDeltaFn, + OnTypedTupleFn, SessionHandle, + TypedErrorStruct, + TypedRowStruct, ) from .compound import Compound, CompoundArg from .program import Program, Schema @@ -46,6 +63,14 @@ Value = int | str | bool | float Row = Sequence[Value] +# The typed path's own vocabulary. `Row` admits `str` because +# `EasySession.insert` auto-interns it; the typed path has no +# forward-intern entry point and rejects `str` by design, while accepting +# anything integer-like or float-like that `Row` does not name. Declaring +# `Row` there would let mypy wave through the one input that always +# raises and reject the NumPy scalars that work. +TypedRow = Sequence[LaneInput] + class _Mode(Enum): """Tracks which evaluation mode the session has entered. The @@ -680,6 +705,289 @@ def remove(self, relation: str, rows: Sequence[Sequence[int]]) -> None: ) check(rc) + # --- typed rows: the FLOAT-capable path -------------------------------- + # + # `insert` / `remove` above carry `int64_t` lanes. What happens to a + # Python float there is NOT a bit reinterpret: `_flatten` calls + # `int(v)`, so on an integer relation 2.9 is stored truncated as 2. + # On a relation that declares a FLOAT column the engine does not get + # that far: `wirelog_session_insert` and `wirelog_session_remove` + # refuse that relation, while `set_delta_cb` and the untyped snapshot + # refuse whenever the program carries a float anywhere - including in + # a compound slot, not only as a declared column. So `insert()` into + # such a relation, and `step()` / `snapshot()` on such a program, all + # raise `ExecError`. + # The typed entry points carry an explicit per-column type alongside + # the bits, which is what makes a FLOAT column usable from a session + # at all. Requires wirelog >= 0.60.0. + + @staticmethod + def _require_typed_row_api(method: str) -> None: + if not has_typed_row_api(): + raise WirelogVersionError( + f"Session.{method} requires libwirelog with the typed row entry " + f"points (wirelog >= 0.60.0); the loaded engine does not export " + f"them. The loader floor still admits 0.52.0, so a source install " + f"against an older system libwirelog reaches this." + ) + + def _column_types(self, relation: str) -> tuple[ColumnType, ...]: + """Logical column types for `relation`, read from the program schema. + + Raises `ExecError` if the relation is undeclared, or if it declares + an `inline` compound column - those spread one logical column over + several physical lanes, and this wrapper only builds the flat + one-lane-per-column descriptor. + """ + schema = self._program.schema(relation) + if schema is None: + raise ExecError(f"no schema for relation: {relation!r}") + for index, column in enumerate(schema.columns): + if column.compound_kind == CompoundKind.INLINE: + raise ExecError( + f"relation {relation!r} column {index} ({column.name!r}) is an " + "inline compound; typed insert/remove supports only relations " + "whose logical columns map one-to-one onto physical lanes" + ) + return tuple(column.type for column in schema.columns) + + def _build_typed_rows( + self, rows: Sequence[TypedRow], types: tuple[ColumnType, ...] + ) -> tuple[Any, int, list[Any]]: + """Build the `TypedRowStruct` array and everything it points at. + + ctypes does retain the backing arrays: assigning the struct into + an array element bumps their refcount through `_objects`, which + is observable as `sys.getrefcount` going 1 -> 2. They are also + returned in an explicit keep-alive list so the guarantee rests on + a reference this code holds rather than on that implementation + detail of ctypes. + """ + ncols = len(types) + type_codes = (ctypes.c_uint32 * ncols)(*(int(t) for t in types)) + lane_offsets = (ctypes.c_uint32 * ncols)(*range(ncols)) + keepalive: list[Any] = [type_codes, lane_offsets] + + # Materialize once. A `Sequence` whose `__len__` disagrees between + # the sizing here and the count handed to C - a hostile `__len__`, + # or a list another thread appends to mid-call - would make wirelog + # read past the end of this array. + materialized = list(rows) + array = (TypedRowStruct * len(materialized))() + for i, row in enumerate(materialized): + if len(row) != ncols: + raise ValueError( + f"row {i} has {len(row)} cols, expected {ncols} from the " f"relation schema" + ) + lanes = (ctypes.c_uint64 * ncols)() + for j, value in enumerate(row): + lanes[j] = encode_lane(value, types[j]) + keepalive.append(lanes) + array[i] = TypedRowStruct( + struct_size=TYPED_ROW_STRUCT_SIZE, + abi_version=TYPED_ROW_ABI_VERSION, + reserved=0, + logical_ncols=ncols, + physical_nlanes=ncols, + physical_stride=ncols, + types=type_codes, + lane_offsets=lane_offsets, + physical_types=type_codes, + lanes=lanes, + ) + return array, len(materialized), keepalive + + @staticmethod + def _raise_typed_error(relation: str, err: TypedErrorStruct, buf: Any) -> None: + engine_message = buf.value.decode("utf-8", "replace") if buf.value else None + row_index = None if err.row_index == TYPED_ERROR_NO_INDEX else int(err.row_index) + column = None if err.logical_col == TYPED_ERROR_NO_INDEX else int(err.logical_col) + try: + code_name = TypedErrorCode(err.code).name + except ValueError: # a code this build of PyreWire does not know + code_name = f"UNKNOWN({err.code})" + + where = f"relation {relation!r}" + if row_index is not None: + where += f", row {row_index}" + if column is not None: + where += f", column {column}" + detail = f": {engine_message}" if engine_message else "" + raise TypedRowError( + f"typed row rejected ({code_name}) - {where}{detail}", + typed_code=int(err.code), + row_index=row_index, + column=column, + engine_message=engine_message, + ) + + def _typed_iud(self, relation: str, rows: Sequence[TypedRow], fn: Any) -> None: + if not rows: + return + types = self._column_types(relation) + # `nrows` comes from what was actually built, never from a second + # `len(rows)`: the two can disagree, and C would read the gap. + array, nrows, keepalive = self._build_typed_rows(rows, types) + if not nrows: + return + buf = ctypes.create_string_buffer(TYPED_ERROR_MESSAGE_CAPACITY) + err = TypedErrorStruct( + struct_size=TYPED_ERROR_STRUCT_SIZE, + code=int(TypedErrorCode.NONE), + row_index=TYPED_ERROR_NO_INDEX, + logical_col=TYPED_ERROR_NO_INDEX, + message=ctypes.cast(buf, ctypes.POINTER(ctypes.c_char)), + message_capacity=TYPED_ERROR_MESSAGE_CAPACITY, + ) + with self._serialize(): + rc = fn( + self._handle, + relation.encode("utf-8"), + array, + ctypes.c_uint32(nrows), + ctypes.byref(err), + ) + del keepalive # the arrays had to stay referenced across the call above + if rc != 0: + if err.code != int(TypedErrorCode.NONE): + self._raise_typed_error(relation, err, buf) + check(rc) + + def insert_typed(self, relation: str, rows: Sequence[TypedRow]) -> None: + """Batched insert that preserves FLOAT columns. + + Column types come from the relation's declaration in the borrowed + program, so rows are plain Python values in declaration order: + + session.insert_typed("sample", [(1, 2.5), (2, 3.5)]) + + `-0.0` and `+0.0` canonicalize to the same `+0.0` on ingress, so + rows differing only in a float's sign of zero collapse to one. + NaN and the infinities are rejected by the engine, not stored: + they raise `TypedRowError` with code `VALUE` and the message + `non-finite float`. The batch is validated as a whole before any + of it is applied, so a rejected row leaves nothing behind. + + Raises: + TypedRowError: wirelog rejected a row; the exception names the + row and column. + ExecError: the relation is undeclared, or declares an `inline` + compound column. + ValueError: a row's width disagrees with the schema, or a + non-integral float was given for an integer column. + TypeError: a `str`, or a value that is neither integer-like + nor float-like, appeared in a row. + OverflowError: a value is too large to convert to a float. + """ + self._require_typed_row_api("insert_typed()") + self._require_mode(_Mode.INCREMENTAL) + self._typed_iud(relation, rows, LIB.wirelog_session_insert_typed) + + def remove_typed(self, relation: str, rows: Sequence[TypedRow]) -> None: + """Like `insert_typed` but emits z-set decrements. + + A float retracts the row it was inserted with only if it encodes to + the same lane bits, which is why `insert_typed`'s zero + canonicalization matters here: removing `-0.0` retracts a row + inserted as `+0.0`. + """ + self._require_typed_row_api("remove_typed()") + self._require_mode(_Mode.INCREMENTAL) + self._typed_iud(relation, rows, LIB.wirelog_session_remove_typed) + + def snapshot_typed(self) -> list[tuple[str, tuple[LaneValue, ...]]]: + """Materialize every derived relation with FLOAT columns decoded. + + On a program that declares a FLOAT column the untyped `snapshot()` + does not return a distorted value - it raises `ExecError`, because + the engine refuses to install an untyped tuple callback there at + all. This decodes each column by its reported type: FLOAT to + `float`, BOOL to `bool`, UINT32 / UINT64 unsigned, everything else + signed. + + Commits the session to QUERY mode, exactly as `snapshot()` does. + """ + self._require_typed_row_api("snapshot_typed()") + self._require_mode(_Mode.QUERY) + # Through `CallbackHandle`, not a bare closure: a decode that + # raises inside a raw ctypes callback is swallowed at the C + # boundary, and this method would then return a short row set with + # a WIRELOG_OK return code. The handle stashes the exception and + # `drain()` re-raises it, which is the contract `_core.callbacks` + # documents and the untyped `snapshot()` already follows. + cb = CallbackHandle("typed_delta") + try: + with self._serialize(): + rc = LIB.wirelog_session_snapshot_typed(self._handle, cb.fn, cb.user_data) + check(rc) + events = cb.drain() + finally: + cb.close() + return [(rel, vals) for _kind, rel, vals, _diff in events] + + def set_typed_delta_callback( + self, fn: Callable[[str, tuple[LaneValue, ...], int], None] | None + ) -> None: + """Register a delta callback that decodes columns by type. + + wirelog REFUSES to install the untyped delta callback on a program + whose schema carries a FLOAT column - `wirelog_session_set_delta_cb` + returns `WIRELOG_ERR_EXEC` - so `set_delta_callback()` and `step()` + are unusable there. This is the path such a session must take. + + `fn` is NOT invoked per row. Registering it arms the typed + trampoline, and the events it buffers are returned by + `step_typed()`; `fn` itself is only stored. That mirrors the + untyped `set_delta_callback`, which does not dispatch to its + callable either - `EasySession.step` is the one place a + user callable is called per event. Read the return value of + `step_typed()` rather than expecting `fn` to fire. + + Passing `None` clears the callback. + """ + self._require_typed_row_api("set_typed_delta_callback()") + self._require_mode(_Mode.INCREMENTAL) + if fn is None: + if self._delta_cb is not None: + self._clear_delta_cb() + return + if self._delta_cb is None or self._delta_cb.kind != "typed_delta": + if self._delta_cb is not None: + self._clear_delta_cb() + self._delta_cb = CallbackHandle("typed_delta") + self._delta_cb._state.user_fn = fn + with self._serialize(): + rc = LIB.wirelog_session_set_typed_delta_cb( + self._handle, self._delta_cb.fn, self._delta_cb.user_data + ) + check(rc) + + def step_typed(self) -> list[tuple[str, tuple[LaneValue, ...], int]]: + """Drive one fixpoint step with typed delta decoding. + + The typed counterpart of `step()`: same `(relation, row, diff)` + events, with FLOAT columns decoded as `float`. On a program that + declares a FLOAT column this is the only one of the two that + works - `step()` raises `ExecError`, because wirelog refuses to + install the untyped delta callback there. + """ + self._require_typed_row_api("step_typed()") + self._require_mode(_Mode.INCREMENTAL) + if self._delta_cb is None or self._delta_cb.kind != "typed_delta": + if self._delta_cb is not None: + self._clear_delta_cb() + self._delta_cb = CallbackHandle("typed_delta") + with self._serialize(): + rc = LIB.wirelog_session_set_typed_delta_cb( + self._handle, self._delta_cb.fn, self._delta_cb.user_data + ) + check(rc) + with self._serialize(): + rc = LIB.wirelog_session_step(self._handle) + check(rc) + events = self._delta_cb.drain() + return [(rel, vals, diff) for _kind, rel, vals, diff in events] + # --- zero-copy NumPy path (#22) ---------------------------------------- def insert_batch(self, relation: str, rows: Any) -> None: @@ -739,19 +1047,48 @@ def _batch_iud(self, relation: str, rows: Any, fn: Any) -> None: # --- step / snapshot / callbacks -------------------------------------- + def _clear_delta_cb(self) -> None: + """Clear wirelog's pointer for whichever callback kind is armed. + + Either entry point would in fact disarm both: 0.60.0's + `wirelog_session_set_delta_cb` nulls `typed_delta_cb` + unconditionally, and its FLOAT-relation refusal is gated on a + non-NULL callback, so a NULL clear is never rejected. Dispatching + on the armed kind does not fix a live bug; it keeps this side from + depending on that, since nothing in the public C contract promises + one setter will keep clearing the other's slot. + """ + cb = self._delta_cb + if cb is None: + return + try: + with self._serialize(): + if cb.kind == "typed_delta": + rc = LIB.wirelog_session_set_typed_delta_cb( + self._handle, OnTypedTupleFn(), None + ) + else: + rc = LIB.wirelog_session_set_delta_cb(self._handle, OnDeltaFn(), None) + check(rc) + finally: + cb.close() + self._delta_cb = None + def set_delta_callback(self, fn: Callable[[str, tuple[int, ...], int], None] | None) -> None: """Register or clear the delta callback. The session enters INCREMENTAL mode on the first call.""" self._require_mode(_Mode.INCREMENTAL) if fn is None: if self._delta_cb is not None: - with self._serialize(): - rc = LIB.wirelog_session_set_delta_cb(self._handle, OnDeltaFn(), None) - check(rc) - self._delta_cb.close() - self._delta_cb = None + self._clear_delta_cb() return - if self._delta_cb is None: + # `kind` matters since the typed path landed: reusing a + # "typed_delta" handle here would hand `OnTypedTupleFn` to + # `wirelog_session_set_delta_cb`, which ctypes rejects with an + # opaque ArgumentError after `user_fn` has already been clobbered. + if self._delta_cb is None or self._delta_cb.kind != "delta": + if self._delta_cb is not None: + self._clear_delta_cb() self._delta_cb = CallbackHandle("delta") self._delta_cb._state.user_fn = fn with self._serialize(): @@ -764,7 +1101,9 @@ def step(self) -> list[tuple[str, tuple[int, ...], int]]: """Drive one fixpoint step. Returns `(relation, row, diff)` events that wirelog emitted for the delta callback during this step.""" self._require_mode(_Mode.INCREMENTAL) - if self._delta_cb is None: + if self._delta_cb is None or self._delta_cb.kind != "delta": + if self._delta_cb is not None: + self._clear_delta_cb() self._delta_cb = CallbackHandle("delta") with self._serialize(): rc = LIB.wirelog_session_set_delta_cb( @@ -858,13 +1197,16 @@ def close(self) -> None: c.invalidate() self._compounds.clear() if self._delta_cb is not None: - # Clear wirelog's pointer before tearing the slot down. + # Clear wirelog's pointer before tearing the slot down, + # through the entry point matching the armed kind. See + # `_clear_delta_cb` for why that is insurance rather than + # a fix. try: - LIB.wirelog_session_set_delta_cb(self._handle, OnDeltaFn(), None) + self._clear_delta_cb() except Exception: - pass - self._delta_cb.close() - self._delta_cb = None + if self._delta_cb is not None: + self._delta_cb.close() + self._delta_cb = None if self._handle.value: LIB.wirelog_session_destroy(self._handle) self._handle = SessionHandle() diff --git a/tests/test_typed_rows.py b/tests/test_typed_rows.py new file mode 100644 index 0000000..3b65253 --- /dev/null +++ b/tests/test_typed_rows.py @@ -0,0 +1,821 @@ +# SPDX-License-Identifier: Apache-2.0 OR GPL-3.0-or-later +"""Tests for the typed session path: `Session.insert_typed` / `remove_typed` +/ `snapshot_typed`. + +These wrap `wirelog_session_insert_typed` and friends, which arrived in +wirelog 0.60.0. They are the only way a FLOAT column reaches a session: +the untyped `insert()` carries `int64_t` lanes and truncates via +`int(v)`, and on a relation that declares a FLOAT column the engine +refuses the untyped entry points outright. +""" + +from __future__ import annotations + +import ctypes +import math +import struct +from dataclasses import replace + +import pytest + +from pyrewire import Program, Session +from pyrewire._core.errors import ExecError, TypedRowError, WirelogVersionError +from pyrewire._core.lanes import decode_lane, encode_lane +from pyrewire._ffi._advanced import has_typed_row_api +from pyrewire._ffi._enums import ColumnType, CompoundKind, TypedErrorCode +from pyrewire._ffi._loader import _parse_version, _pep440_base +from pyrewire._ffi._types import ( + TYPED_ERROR_STRUCT_SIZE, + TYPED_ROW_ABI_VERSION, + TYPED_ROW_STRUCT_SIZE, + OnTypedTupleFn, + TypedErrorStruct, + TypedRowStruct, +) +from pyrewire._ffi._util import wirelog_version + +FLOAT_SRC = """ +.decl sample(a: int64, value: float) +.decl seen(a: int64, value: float) +seen(A, V) :- sample(A, V). +""" + +INT_SRC = """ +.decl edge(x: int64, y: int64) +.decl reach(x: int64, y: int64) +reach(X, Y) :- edge(X, Y). +""" + +# Inline facts, so a QUERY-mode session has something to snapshot without +# an insert first. `Session`'s mode machine commits to INCREMENTAL on the +# first insert, which then rejects `snapshot_typed()`. +FLOAT_INLINE_SRC = FLOAT_SRC + """ +sample(1, 2.5). +sample(2, 3.5). +""" + + +def _wirelog_older_than(minimum: tuple[int, int, int]) -> bool: + return _parse_version(_pep440_base(wirelog_version())) < minimum + + +typed_api = pytest.mark.skipif( + _wirelog_older_than((0, 60, 0)), + reason="the typed row entry points first ship in wirelog 0.60.0", +) + + +# ---------------------------------------------------------------------- +# ABI shape — these hold on any engine, so they are not version-gated. +# ---------------------------------------------------------------------- + + +def test_typed_row_struct_matches_documented_abi(): + """Field order and offsets mirror `wirelog_typed_row_v1_t`. + + A silently reordered field would not fail to compile here; it would + hand wirelog a descriptor whose pointers land in the wrong slots. + """ + assert TypedRowStruct.struct_size.offset == 0 + assert TypedRowStruct.abi_version.offset == 4 + assert TypedRowStruct.logical_ncols.offset == 8 + assert TypedRowStruct.physical_nlanes.offset == 12 + assert TypedRowStruct.physical_stride.offset == 16 + # Four pointers, naturally aligned after the header. + assert TypedRowStruct.types.offset == 24 + assert TypedRowStruct.lane_offsets.offset == 32 + assert TypedRowStruct.physical_types.offset == 40 + assert TypedRowStruct.lanes.offset == 48 + assert TYPED_ROW_STRUCT_SIZE == ctypes.sizeof(TypedRowStruct) + + +def test_typed_error_struct_matches_documented_abi(): + assert TypedErrorStruct.struct_size.offset == 0 + assert TypedErrorStruct.code.offset == 4 + assert TypedErrorStruct.row_index.offset == 8 + assert TypedErrorStruct.logical_col.offset == 12 + assert TypedErrorStruct.message.offset == 16 + assert TypedErrorStruct.message_capacity.offset == 24 + assert TYPED_ERROR_STRUCT_SIZE == ctypes.sizeof(TypedErrorStruct) + + +def test_typed_error_message_is_a_writeable_buffer(): + """`message` must be `POINTER(c_char)`, not `c_char_p`. + + `c_char_p` surfaces in Python as an immutable `bytes`, so wirelog's + bounded diagnostic would be written into a buffer nothing can read + back. + """ + assert TypedErrorStruct.message.__get__ is not None + field_type = dict(TypedErrorStruct._fields_)["message"] + assert field_type is not ctypes.c_char_p + assert field_type == ctypes.POINTER(ctypes.c_char) + + +def test_typed_tuple_callback_signature(): + assert OnTypedTupleFn._restype_ is None + assert OnTypedTupleFn._argtypes_ == ( + ctypes.c_char_p, + ctypes.POINTER(TypedRowStruct), + ctypes.c_int32, + ctypes.c_void_p, + ) + + +def test_typed_error_codes_match_the_c_enum(): + assert int(TypedErrorCode.NONE) == 0 + assert int(TypedErrorCode.DESCRIPTOR) == 1 + assert int(TypedErrorCode.SCHEMA) == 2 + assert int(TypedErrorCode.VALUE) == 3 + + +def test_abi_version_is_one(): + assert TYPED_ROW_ABI_VERSION == 1 + + +# ---------------------------------------------------------------------- +# Lane encoding — pure, so no engine needed. +# ---------------------------------------------------------------------- + + +def test_float_lane_roundtrips_through_ieee754_bits(): + for value in (2.5, -0.0, 0.0, 1e308, -1e-308, math.pi): + lane = encode_lane(value, ColumnType.FLOAT) + assert lane == struct.unpack(" 0 + + +@typed_api +def test_remove_typed_retracts_a_float_row(): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + s.insert_typed("sample", [(1, 2.5), (2, 3.5)]) + assert len(s.step_typed()) == 2 + s.remove_typed("sample", [(1, 2.5)]) + assert s.step_typed() == [("seen", (1, 2.5), -1)] + + +@typed_api +def test_remove_typed_with_negative_zero_retracts_a_positive_zero_row(): + """Both spellings canonicalize on ingress, so the retraction lands.""" + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + s.insert_typed("sample", [(1, 0.0)]) + assert s.step_typed() == [("seen", (1, 0.0), 1)] + s.remove_typed("sample", [(1, -0.0)]) + assert s.step_typed() == [("seen", (1, 0.0), -1)] + + +@typed_api +def test_typed_path_works_for_a_plain_int_relation(): + """Nothing about the typed path is float-specific.""" + with Program.from_string(INT_SRC) as prog, Session(prog) as s: + s.insert_typed("edge", [(1, 2), (2, 3)]) + events = sorted(s.step_typed()) + assert events == [("reach", (1, 2), 1), ("reach", (2, 3), 1)] + + +@typed_api +def test_empty_rows_is_a_no_op(): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + s.insert_typed("sample", []) + s.remove_typed("sample", []) + assert s.step_typed() == [] + + +@typed_api +def test_registered_typed_callback_is_stored_but_never_invoked(): + """`fn` arms the trampoline; the events come back from `step_typed()`. + + The registered callable is deliberately NOT called per row -- the + untyped `set_delta_callback` behaves the same way, and only + `EasySession.step` dispatches to a user callable. Pinning it here so + the docstring and the behavior cannot drift apart again: the previous + version of this test was named "...delivers_decoded_rows" and asserted + nothing about `fn`, which is exactly the claim that was false. + """ + calls: list[tuple] = [] + + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + s.set_typed_delta_callback(lambda rel, row, diff: calls.append((rel, row, diff))) + s.insert_typed("sample", [(1, 2.5)]) + events = s.step_typed() + + assert events == [("seen", (1, 2.5), 1)] + assert calls == [] + + +@typed_api +def test_clear_typed_delta_callback(): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + s.set_typed_delta_callback(lambda rel, row, diff: None) + s.set_typed_delta_callback(None) + + +# ---------------------------------------------------------------------- +# Error surface. +# ---------------------------------------------------------------------- + + +@typed_api +def test_wrong_row_width_names_the_row_before_reaching_the_engine(): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + with pytest.raises(ValueError, match=r"row 1 has 1 cols, expected 2"): + s.insert_typed("sample", [(1, 2.5), (2,)]) + + +@typed_api +def test_undeclared_relation_raises_exec_error(): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + with pytest.raises(ExecError, match="no schema for relation"): + s.insert_typed("nope", [(1, 2.5)]) + + +@typed_api +def test_column_types_come_from_the_program_schema(): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + assert s._column_types("sample") == (ColumnType.INT64, ColumnType.FLOAT) + + +@typed_api +def test_typed_row_error_carries_the_engine_diagnostic(): + """A row wirelog rejects must arrive with its detail attached. + + The generic path raises a bare `ExecError` reading "execution error"; + the point of threading `wirelog_typed_error_v1_t` through is that the + caller learns the code, the row, and the engine's own message. + """ + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + # A descriptor whose arity disagrees with the relation. Built by + # hand because `insert_typed` derives the width from the schema and + # would reject this before the FFI call ever happens. + lanes = (ctypes.c_uint64 * 1)(7) + codes = (ctypes.c_uint32 * 1)(int(ColumnType.INT64)) + offsets = (ctypes.c_uint32 * 1)(0) + row = TypedRowStruct( + struct_size=TYPED_ROW_STRUCT_SIZE, + abi_version=TYPED_ROW_ABI_VERSION, + reserved=0, + logical_ncols=1, + physical_nlanes=1, + physical_stride=1, + types=codes, + lane_offsets=offsets, + physical_types=codes, + lanes=lanes, + ) + s._build_typed_rows = lambda rows, types: ( + (TypedRowStruct * 1)(row), + 1, + [lanes, codes, offsets], + ) + + with pytest.raises(TypedRowError) as excinfo: + s.insert_typed("sample", [(1, 2.5)]) + + err = excinfo.value + assert isinstance(err, ExecError) # existing handlers still catch it + assert err.typed_code != int(TypedErrorCode.NONE) + assert "relation 'sample'" in str(err) + assert err.engine_message # wirelog filled the caller-owned buffer + + +@typed_api +def test_non_finite_float_is_rejected_by_the_engine_naming_the_row(): + """The reachable `VALUE` error, through the real descriptor builder. + + The monkeypatched test above forges a descriptor `insert_typed` can + never produce, so it pins neither `row_index` accuracy nor which code + is reported. This one goes through `_build_typed_rows` unmodified. + """ + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + rows = [(1, 1.5), (2, 2.5), (3, float("nan")), (4, 4.5)] + with pytest.raises(TypedRowError) as excinfo: + s.insert_typed("sample", rows) + + err = excinfo.value + assert err.typed_code == int(TypedErrorCode.VALUE) + assert err.row_index == 2 # the NaN row, not the first or the last + assert err.column == 1 + assert "non-finite" in (err.engine_message or "") + + # The batch is validated as a whole: nothing before the bad row + # was applied. + assert s.step_typed() == [] + + +@typed_api +@pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf")]) +def test_every_non_finite_spelling_is_rejected(value): + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + with pytest.raises(TypedRowError): + s.insert_typed("sample", [(1, value)]) + + +@typed_api +def test_string_in_a_typed_row_is_rejected_rather_than_coerced(): + """`int("5")` is 5, which is a valid intern id for some other symbol. + + The advanced Session has no forward-intern entry point, so there is + nothing here that could turn text into the right id. Refuse instead of + writing a wrong one. + """ + with Program.from_string(INT_SRC) as prog, Session(prog) as s: + with pytest.raises(TypeError, match="intern id"): + s.insert_typed("edge", [("5", 2)]) + + +@typed_api +def test_snapshot_typed_reraises_a_decode_error_instead_of_truncating(): + """A raise inside the callback must not become a short row set. + + Through a raw ctypes callback the exception is swallowed at the C + boundary and the method returns the rows decoded so far with a + WIRELOG_OK return code -- a silently wrong answer. Routing through + `CallbackHandle` stashes it and `drain()` re-raises. + """ + # The decode happens inside the trampoline, so patch it there. + import pyrewire._core.callbacks as callbacks_mod + + calls = {"n": 0} + + def exploding_decode(row): + calls["n"] += 1 + raise RuntimeError("decode boom") + + with Program.from_string(FLOAT_INLINE_SRC) as prog, Session(prog) as s: + original = callbacks_mod.decode_typed_row + callbacks_mod.decode_typed_row = exploding_decode + try: + with pytest.raises(RuntimeError, match="decode boom"): + s.snapshot_typed() + finally: + callbacks_mod.decode_typed_row = original + + assert calls["n"] > 0 # the callback really ran + + +@typed_api +def test_typed_and_untyped_callbacks_do_not_share_a_handle(): + """Switching kinds must re-arm, not reuse the other kind's handle. + + Reusing it hands `OnTypedTupleFn` to `wirelog_session_set_delta_cb`, + which ctypes rejects with an opaque ArgumentError after `user_fn` has + already been overwritten. + """ + with Program.from_string(INT_SRC) as prog, Session(prog) as s: + s.set_typed_delta_callback(lambda *a: None) + assert s._delta_cb is not None and s._delta_cb.kind == "typed_delta" + + s.set_delta_callback(lambda *a: None) + assert s._delta_cb is not None and s._delta_cb.kind == "delta" + + s.insert("edge", [(1, 2)]) + assert s.step() == [("reach", (1, 2), 1)] + + +@typed_api +def test_closing_a_session_clears_through_the_armed_kinds_entry_point(): + """`close()` must call the clear matching the armed kind. + + Asserting only `_delta_cb is None` afterwards pins nothing: the old + code set that too, so the test stayed green with the kind dispatch + reverted. Record which C entry point was actually called. + """ + from pyrewire._ffi import LIB + + for src, kind, expected in ( + (FLOAT_SRC, "typed_delta", "wirelog_session_set_typed_delta_cb"), + (INT_SRC, "delta", "wirelog_session_set_delta_cb"), + ): + called: list[str] = [] + real_typed = LIB.wirelog_session_set_typed_delta_cb + real_untyped = LIB.wirelog_session_set_delta_cb + + with Program.from_string(src) as prog: + s = Session(prog) + if kind == "typed_delta": + s.set_typed_delta_callback(lambda *a: None) + else: + s.set_delta_callback(lambda *a: None) + assert s._delta_cb is not None and s._delta_cb.kind == kind + + def _typed(*a, _r=real_typed): + called.append("wirelog_session_set_typed_delta_cb") + return _r(*a) + + def _untyped(*a, _r=real_untyped): + called.append("wirelog_session_set_delta_cb") + return _r(*a) + + LIB.wirelog_session_set_typed_delta_cb = _typed + LIB.wirelog_session_set_delta_cb = _untyped + try: + s.close() + finally: + LIB.wirelog_session_set_typed_delta_cb = real_typed + LIB.wirelog_session_set_delta_cb = real_untyped + + assert called == [expected], f"{kind}: {called}" + assert s._delta_cb is None + + +@typed_api +def test_snapshot_typed_releases_its_registry_slot_on_the_raising_path(): + """Pin the `finally: cb.close()`, not refcounting. + + Counting `_REGISTRY` after a successful call proves nothing: + `CallbackHandle.__del__` reclaims the slot by refcount whether or not + `close()` ran, so such a test passes with the `finally` deleted. The + explicit close is load-bearing exactly when the call raises - the + propagating traceback pins the frame holding `cb`, so `__del__` does + not fire promptly. Drive that path. + """ + import pyrewire._core.callbacks as callbacks_mod + + with Program.from_string(FLOAT_INLINE_SRC) as prog, Session(prog) as s: + before = len(callbacks_mod._REGISTRY) + + original = callbacks_mod.decode_typed_row + + def boom(row): + raise RuntimeError("decode boom") + + callbacks_mod.decode_typed_row = boom + try: + with pytest.raises(RuntimeError, match="decode boom"): + s.snapshot_typed() + finally: + callbacks_mod.decode_typed_row = original + + # Still inside the frame that caught the exception, so the + # traceback may still reference the failed call. Only the explicit + # close() can have released the slot by now. + assert len(callbacks_mod._REGISTRY) == before + + +@typed_api +def test_step_typed_rearms_after_an_untyped_callback_was_installed(): + """Switching kinds into `step_typed()` must re-arm, not reuse.""" + with Program.from_string(INT_SRC) as prog, Session(prog) as s: + s.set_delta_callback(lambda *a: None) + assert s._delta_cb is not None and s._delta_cb.kind == "delta" + + s.insert_typed("edge", [(1, 2)]) + assert s.step_typed() == [("reach", (1, 2), 1)] + assert s._delta_cb is not None and s._delta_cb.kind == "typed_delta" + + +@typed_api +def test_row_count_comes_from_what_was_built_not_a_second_len(): + """A `Sequence` whose `__len__` changes must not make C read past the end. + + Sizing the descriptor array from one `len(rows)` and handing C a + second one lets wirelog walk off the allocation. Both now come from + the materialized list. + """ + + class ShiftyLen(list): + def __init__(self, items): + super().__init__(items) + self._calls = 0 + + def __len__(self): + self._calls += 1 + return super().__len__() if self._calls <= 2 else 64 + + with Program.from_string(FLOAT_SRC) as prog, Session(prog) as s: + s.insert_typed("sample", ShiftyLen([(1, 1.5)])) + # Exactly the one real row was inserted; no descriptor beyond it + # was ever handed to the engine. + assert s.step_typed() == [("seen", (1, 1.5), 1)] + + +def test_encode_lane_accepts_index_and_float_protocols(): + """NumPy scalars go through `insert_batch`; they must work here too.""" + + class Indexy: + def __index__(self): + return 7 + + class Floaty: + def __float__(self): + return 2.5 + + assert encode_lane(Indexy(), ColumnType.INT64) == 7 + assert decode_lane(encode_lane(Floaty(), ColumnType.FLOAT), ColumnType.FLOAT) == 2.5 + # Float-like into an integer column still has to name an exact integer. + with pytest.raises(ValueError, match="cannot store float"): + encode_lane(Floaty(), ColumnType.INT64) + with pytest.raises(TypeError, match="cannot encode"): + encode_lane(object(), ColumnType.INT64) + + +def test_integral_float_only_value_is_accepted_not_falsely_rejected(): + """A bare `__float__` provider has nothing to lose. + + The lossiness guard compares the float result against the original, + which is only meaningful for a real numeric type. A plain + `__float__` wrapper has no numeric `__eq__`, so that comparison falls + back to identity and rejects every such value while claiming it was + unrepresentable. The previous guard did exactly that. + + The existing protocol test uses 2.5, which short-circuits at the + integrality check and never reaches the comparison - which is why + nothing caught it. + """ + + class Meters: + def __init__(self, v): + self.v = v + + def __float__(self): + return self.v + + assert encode_lane(Meters(4.0), ColumnType.INT64) == 4 + assert encode_lane(Meters(-7.0), ColumnType.INT64) == (-7) & 0xFFFFFFFFFFFFFFFF + assert decode_lane(encode_lane(Meters(2.5), ColumnType.FLOAT), ColumnType.FLOAT) == 2.5 + # Non-integral into an integer column is still refused. + with pytest.raises(ValueError, match="cannot store float"): + encode_lane(Meters(2.5), ColumnType.INT64) + + +def test_a_raising_eq_cannot_escape_encode_lane(): + """The guard must not let a user `__eq__` throw past the contract. + + `EqBoom` has to be registered as a `numbers.Number`, or it returns at + the isinstance gate and never reaches the comparison this test is + about - which is what the first version of it did. + """ + import numbers + + class EqBoom: + def __float__(self): + return 4.0 + + def __eq__(self, other): + raise KeyError("boom") + + numbers.Number.register(EqBoom) + + assert encode_lane(EqBoom(), ColumnType.INT64) == 4 + + +def test_int_too_large_for_float_keeps_its_overflow_error(): + """`float(10**400)` overflowing is about the value, not a `__float__`. + + Reporting it as a TypeError blaming `__float__` would be doubly + wrong: `int` has no `__float__`, and the value is integer-like. + """ + with pytest.raises(OverflowError): + encode_lane(10**400, ColumnType.FLOAT) + + +def test_lossy_integer_like_is_refused_rather_than_stored_wrong(): + """`Decimal(2**63-1)` names an integer no binary64 can hold. + + Routing it through `float` would store 2**63, one off from the value + the caller named, with no error. Refuse instead. + """ + from decimal import Decimal + + exact = Decimal(2**53) # representable + assert encode_lane(exact, ColumnType.INT64) == 2**53 + + with pytest.raises(ValueError, match="not exactly representable"): + encode_lane(Decimal(2**63 - 1), ColumnType.INT64) + + from fractions import Fraction + + assert encode_lane(Fraction(3, 1), ColumnType.INT64) == 3 + with pytest.raises(ValueError, match="not exactly representable"): + encode_lane(Fraction(2**63 - 1, 1), ColumnType.INT64) + + +def test_index_returning_a_non_int_is_a_clean_type_error(): + """`operator.index` enforces the protocol's own contract. + + Calling `__index__()` directly would let a bogus return value reach + the mask and fail as an unrelated operand error. + """ + + class Liar: + def __index__(self): + return "not an int" + + with pytest.raises(TypeError, match="__index__ returned non-int"): + encode_lane(Liar(), ColumnType.INT64) + + +def test_a_raising_float_conversion_is_reported_as_a_type_error(): + class Exploding: + def __float__(self): + raise RuntimeError("boom") + + with pytest.raises(TypeError, match="__float__ raised"): + encode_lane(Exploding(), ColumnType.FLOAT) + + +def test_encode_lane_accepts_numpy_scalars_when_numpy_is_present(): + np = pytest.importorskip("numpy") + assert encode_lane(np.int64(5), ColumnType.INT64) == 5 + assert encode_lane(np.int32(-3), ColumnType.INT64) == (-3) & 0xFFFFFFFFFFFFFFFF + assert decode_lane(encode_lane(np.float64(2.5), ColumnType.FLOAT), ColumnType.FLOAT) == 2.5 + + +@typed_api +def test_inline_compound_relation_is_refused_with_a_clear_message(): + """`_column_types` builds a flat one-lane-per-column descriptor, which + an inline compound column does not fit. Refuse loudly rather than hand + wirelog a descriptor whose lanes mean something else.""" + with Program.from_string(INT_SRC) as prog, Session(prog) as s: + real = prog.schema("edge") + assert real is not None + inline_col = replace(real.columns[1], compound_kind=CompoundKind.INLINE) + patched = replace(real, columns=(real.columns[0], inline_col)) + s._program.schema = lambda relation: patched # type: ignore[method-assign] + + with pytest.raises(ExecError, match="inline compound"): + s.insert_typed("edge", [(1, 2)]) + + +# ---------------------------------------------------------------------- +# Older engines. +# ---------------------------------------------------------------------- + + +def test_typed_api_detection_agrees_with_the_engine_version(): + """`has_typed_row_api()` must track the 0.60.0 boundary. + + The loader floor still admits wirelog 0.52.0, so PyreWire has to keep + importing there - registering these argtypes unguarded would raise + `AttributeError` at import time and take the whole package down. + """ + assert has_typed_row_api() is (not _wirelog_older_than((0, 60, 0))) + + +@pytest.mark.skipif( + not _wirelog_older_than((0, 60, 0)), + reason="covers the pre-0.60.0 engine path", +) +@pytest.mark.parametrize( + "call", + [ + pytest.param(lambda s: s.insert_typed("edge", [(1, 2)]), id="insert_typed"), + pytest.param(lambda s: s.remove_typed("edge", [(1, 2)]), id="remove_typed"), + pytest.param(lambda s: s.insert_typed("edge", []), id="insert_typed-empty"), + pytest.param(lambda s: s.snapshot_typed(), id="snapshot_typed"), + pytest.param(lambda s: s.step_typed(), id="step_typed"), + pytest.param(lambda s: s.set_typed_delta_callback(lambda *a: None), id="set-cb"), + pytest.param(lambda s: s.set_typed_delta_callback(None), id="clear-cb"), + ], +) +def test_every_typed_method_raises_version_error_on_an_older_engine(call): + """All five guards, not just `insert_typed`'s. + + A mutation dropping the guard from any of the other four left the + suite green when only `insert_typed` was covered; the symbol is simply + absent on an older engine, so the failure would surface as a raw + ctypes `AttributeError` instead of this typed error. + """ + with Program.from_string(INT_SRC) as prog, Session(prog) as s: + with pytest.raises(WirelogVersionError, match="0.60.0"): + call(s) + + +def test_package_imports_without_the_typed_entry_points(monkeypatch): + """Simulate an older engine: the guard, not ctypes, must decide.""" + from pyrewire._ffi import _advanced + + monkeypatch.setattr(_advanced, "TYPED_ROW_ENTRY_POINTS", ("wirelog_no_such_symbol",)) + assert _advanced.has_typed_row_api() is False + # Registration is a no-op in that state rather than an AttributeError. + _advanced._register()