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()