diff --git a/CHANGELOG.md b/CHANGELOG.md index 0ba192a7c..0b9020258 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,6 +15,7 @@ v2.16.0 is a feature release with the following features, fixes and enhancements `dlq.auto.flush=true` or give the serde its own `RuleRegistry` (closable on shutdown) for durability. - Add support for inline validation rules (#2326) +- Add Variant, Decimal, and Timestamp CEL functions (#2332) ### Fixes diff --git a/pyproject.toml b/pyproject.toml index d5e80f1d9..5224b873e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,6 +26,11 @@ Homepage = "https://github.com/confluentinc/confluent-kafka-python" [tool.mypy] ignore_missing_imports = true +# The generated protobuf modules are not statically analysable: _builder injects the message +# classes into globals() at import time. Excluding them keeps them out of the build set, which +# is what lets follow_imports = "skip" below apply - it is ignored for files mypy was asked to +# check directly. +exclude = '_pb2\.py$' [[tool.mypy.overrides]] module = [ @@ -35,12 +40,18 @@ module = [ ] disable_error_code = ["assignment", "no-redef"] +# Generated protobuf modules build their message classes at import time, so the classes are +# not statically visible - neither in the generated file itself nor to anything annotating +# against them. follow_imports = "skip" makes the module Any, which covers both. [[tool.mypy.overrides]] module = [ "confluent_kafka.schema_registry.confluent.meta_pb2", + "confluent_kafka.schema_registry.confluent.type.decimal_pb2", + "confluent_kafka.schema_registry.confluent.type.variant_pb2", "confluent_kafka.schema_registry.confluent.types.decimal_pb2", ] ignore_errors = true +follow_imports = "skip" [tool.black] line-length = 120 diff --git a/src/confluent_kafka/admin/__init__.py b/src/confluent_kafka/admin/__init__.py index ec53c9520..f6e581964 100644 --- a/src/confluent_kafka/admin/__init__.py +++ b/src/confluent_kafka/admin/__init__.py @@ -33,6 +33,7 @@ from .._model import ElectionType as _ElectionType from .._model import TopicCollection as _TopicCollection from ..cimpl import KafkaException # noqa: F401 +from ..cimpl import _AdminClientImpl # noqa: F401 from ..cimpl import ( # noqa: F401 CONFIG_SOURCE_DEFAULT_CONFIG, CONFIG_SOURCE_DYNAMIC_BROKER_CONFIG, @@ -53,7 +54,6 @@ NewTopic, ) from ..cimpl import TopicPartition as _TopicPartition -from ..cimpl import _AdminClientImpl from ._acl import AclOperation # noqa: F401 from ._acl import AclBinding, AclBindingFilter, AclPermissionType # noqa: F401 from ._cluster import DescribeClusterResult # noqa: F401 diff --git a/src/confluent_kafka/schema_registry/common/avro.py b/src/confluent_kafka/schema_registry/common/avro.py index cf3c3c843..cabbcd583 100644 --- a/src/confluent_kafka/schema_registry/common/avro.py +++ b/src/confluent_kafka/schema_registry/common/avro.py @@ -10,6 +10,7 @@ from fastavro import repository, validate from fastavro.schema import load_schema +from confluent_kafka.schema_registry.confluent.type.variant_utils import Variant from confluent_kafka.schema_registry.serde import ( VALIDATION_RULES_PROP, FieldTransform, @@ -25,6 +26,33 @@ from .schema_registry_client import RuleKind, Schema + +# The Avro `variant` logical type: a record {metadata: bytes, value: bytes} carrying a +# Spark/Parquet Variant. A field with this logical type decodes to / encodes from a Variant, +# so serde consumers and CEL rules see a first-class Variant rather than raw bytes. fastavro +# keys logical handlers by "-" = "record-variant"; this is the Python +# counterpart of Java's io.confluent.avro.type.VariantConversion. +def _variant_from_avro(data, writer_schema, reader_schema=None): # noqa: ARG001 + return Variant(bytes(data["value"]), bytes(data["metadata"])) + + +def _variant_to_avro(data, schema): # noqa: ARG001 + if isinstance(data, Variant): + # standalone_value_bytes, not .value: a navigated sub-variant's own value starts at + # its position, and .value is the whole shared buffer. + return {"metadata": data.metadata, "value": data.standalone_value_bytes()} + return data + + +try: + from fastavro import read as _fastavro_read + from fastavro import write as _fastavro_write + + _fastavro_read.LOGICAL_READERS["record-variant"] = _variant_from_avro + _fastavro_write.LOGICAL_WRITERS["record-variant"] = _variant_to_avro +except Exception: # pragma: no cover - guards against a fastavro API shape change + pass + __all__ = [ 'AvroMessage', 'AvroSchema', @@ -116,7 +144,11 @@ def parse_schema_with_repo(schema_str: str, named_schemas: Dict[str, AvroSchema] def transform( ctx: RuleContext, schema: AvroSchema, message: AvroMessage, field_transform: FieldTransform ) -> AvroMessage: - if message is None or schema is None: + # Only the schema being absent stops the walk. A `None` *value* is the null branch of a + # `["null", T]` union and has to reach the rule: the reference binds it as CEL null so a + # rule can guard with `value == null`, and returning early here skipped the rule entirely - + # indistinguishable, to the caller, from a rule that ran and passed. + if schema is None: return message field_ctx = ctx.current_field() if field_ctx is not None: @@ -126,7 +158,12 @@ def transform( if subschema is None: return message submessage = transform(ctx, subschema, submessage, field_transform) - if isinstance(message, tuple) and len(message) == 2: + # The branch a transformed value belongs to follows from the value, not from the branch + # it arrived on - the reference keeps no branch at all and resolves it from the datum. + # Keep the notation while its branch still accepts the result, so two same-shaped + # records are never swapped; drop it otherwise and let fastavro resolve, since + # ("null", x) would be written as null with x silently dropped. + if isinstance(message, tuple) and len(message) == 2 and _branch_accepts(subschema, submessage): return (message[0], submessage) return submessage elif isinstance(schema, dict): @@ -142,6 +179,11 @@ def transform( return message return {key: transform(ctx, schema["values"], value, field_transform) for key, value in message.items()} elif schema_type == 'record': + # A null record has no fields to walk. Guarded before the isinstance check below + # so a legitimate null does not log an "incompatible message type" warning; the + # reference guards the record case, and only the record case, the same way. + if message is None: + return message if not isinstance(message, dict): log.warning("Incompatible message type for record schema") return message @@ -364,6 +406,15 @@ def _union_branch_matches(subschema: AvroSchema, branch_name: str, exact: bool) return '.' not in name and not subschema.get("namespace") and branch_name.rsplit('.', 1)[-1] == name +def _branch_accepts(subschema: AvroSchema, message: AvroMessage) -> bool: + """Whether ``subschema`` can hold ``message``, by the same test ``_resolve_union`` uses.""" + try: + validate(message, _collapse_schema(deepcopy(subschema))) + return True + except: # noqa: E722 + return False + + def _resolve_union(schema: AvroSchema, message: AvroMessage) -> Tuple[Optional[AvroSchema], AvroMessage]: is_wrapped_union = isinstance(message, tuple) and len(message) == 2 is_typed_union = isinstance(message, dict) and '-type' in message diff --git a/src/confluent_kafka/schema_registry/common/protobuf.py b/src/confluent_kafka/schema_registry/common/protobuf.py index b6782ecc6..198437432 100644 --- a/src/confluent_kafka/schema_registry/common/protobuf.py +++ b/src/confluent_kafka/schema_registry/common/protobuf.py @@ -1,8 +1,10 @@ import base64 +import datetime +import decimal import io import sys from collections import deque -from decimal import MAX_PREC, Context, Decimal +from decimal import MAX_EMAX, MAX_PREC, MIN_EMIN, ROUND_HALF_UP, Context, Decimal from typing import Any, Deque, Dict, List, Optional, Set, Tuple from google.protobuf import __version__ as _protobuf_version @@ -41,12 +43,18 @@ import confluent_kafka.schema_registry.confluent.meta_pb2 as meta_pb2 from confluent_kafka.schema_registry import RuleKind -from confluent_kafka.schema_registry.confluent.types import decimal_pb2 +from confluent_kafka.schema_registry.confluent.type import decimal_pb2, variant_pb2 +from confluent_kafka.schema_registry.confluent.type.decimal_utils import ( + unscaled_to_bytes, +) +from confluent_kafka.schema_registry.confluent.type.variant_utils import Variant +from confluent_kafka.schema_registry.confluent.types import decimal_pb2 as legacy_decimal_pb2 from confluent_kafka.schema_registry.serde import ( FieldTransform, FieldType, RuleConditionError, RuleContext, + RuleError, ValidationRule, ValidationRuleError, ValidationRuleExecutor, @@ -73,6 +81,8 @@ '_is_builtin', 'decimal_to_protobuf', 'protobuf_to_decimal', + 'variant_to_protobuf', + 'protobuf_to_variant', ] # Convert an int to bytes (inverse of ord()) @@ -253,6 +263,137 @@ def _init_pool(pool: DescriptorPool): pool.AddSerializedFile(meta_pb2.DESCRIPTOR.serialized_pb) pool.AddSerializedFile(decimal_pb2.DESCRIPTOR.serialized_pb) + pool.AddSerializedFile(variant_pb2.DESCRIPTOR.serialized_pb) + # The path confluent.type.Decimal used to occupy. A schema importing it is never sent with a + # reference - _is_builtin matches the whole confluent/ prefix - so the pool is the only place + # a reader can resolve it from. The stub declares nothing and publicly imports the canonical + # file, so it re-exports confluent.type.Decimal without a second declaration of the symbol, + # which a pool refuses. Added after the canonical file, which it depends on. Variant needs no + # such stub: it had not shipped under the old path. + pool.AddSerializedFile(legacy_decimal_pb2.DESCRIPTOR.serialized_pb) + + +# Message types a CEL rule works with as a single value rather than as a record. +# +# Avro carries the same concepts as logical types on a primitive, so the field is a leaf there +# and a CEL_FIELD rule reaches it. In protobuf they are messages, and without this the walk +# descends into their internals and transforms `value`/`scale` or `seconds`/`nanos` one at a +# time instead - which is not what the rule asked for, and which an untagged rule would do +# silently. Ported from the JVM client's ProtobufSchema.isCelLeafMessage (#4538). +# +# Variant is deliberately *not* a leaf: it is a record in Avro too, so skipping it is the +# behaviour that matches, and a variant is reached with a message-level CEL rule instead. +DECIMAL_TYPE_NAME = "confluent.type.Decimal" +TIMESTAMP_TYPE_NAME = "google.protobuf.Timestamp" + + +def is_cel_leaf_message(desc: Optional[Descriptor]) -> bool: + """Whether *desc* is a message type bound to CEL as a single value.""" + return desc is not None and desc.full_name in (DECIMAL_TYPE_NAME, TIMESTAMP_TYPE_NAME) + + +# The widest coefficient this client can *encode*, as opposed to compute with. The wire form is +# the unscaled integer in base 256, and decimal <-> binary radix conversion is quadratic; 4300 is +# CPython's own ``int_max_str_digits``, the cap it puts on str <-> int for exactly that reason, and +# the number every client in the family adopts so they agree on which decimals can be written. +# +# Defined here, in the lower layer, and imported by the CEL writer - there were two constants +# named ``_MAX_COEFFICIENT_DIGITS`` in this client with different values, which is a trap for the +# next reader. The other one is now ``_MAX_BIGINTEGER_DIGITS``, which is what it always meant. +MAX_ENCODABLE_COEFFICIENT_DIGITS = 4300 + + +def set_decimal_message(target: Message, value: decimal.Decimal) -> None: + """Writes a Python Decimal into a confluent.type.Decimal message. + + Precision and scale describe the value itself rather than a declared column width, which is + the same mapping the JVM client uses (DecimalUtils.fromBigDecimal) and the reverse of how a + decimal is read back. + """ + sign, digits, exponent = value.as_tuple() + if not isinstance(exponent, int): + raise ValueError("cannot write a non-finite decimal to " + DECIMAL_TYPE_NAME) + # Checked here rather than left to CPython: the ``int(...)`` below is a str -> int + # conversion, which raises "Exceeds the limit (4300 digits) for integer string conversion" + # naming neither the decimal nor the field. This is the *field-level* write-back, the twin + # of the message-level ``_set_decimal``, and it had the unhelpful version. + if len(digits) > MAX_ENCODABLE_COEFFICIENT_DIGITS: + raise ValueError( + f"decimal coefficient has {len(digits)} digits, past the " + f"{MAX_ENCODABLE_COEFFICIENT_DIGITS} this client can encode into " + DECIMAL_TYPE_NAME + ) + unscaled = int("".join(str(d) for d in digits) or "0") + if sign: + unscaled = -unscaled + # The scale is the negated exponent, negative included: BigDecimal("1E+3") reports + # unscaled 1 with scale -3, and the proto field is a signed int32, so normalising a + # positive exponent into the digits would write a different value than the JVM does. + scale = -exponent + target.value = unscaled_to_bytes(unscaled) + target.precision = len(digits) + target.scale = scale + + +def set_timestamp_message(target: Message, value: datetime.datetime) -> None: + """Writes a datetime into a google.protobuf.Timestamp message.""" + if value.tzinfo is None: + value = value.replace(tzinfo=datetime.timezone.utc) + delta = value - _EPOCH + target.seconds = delta.days * 86400 + delta.seconds + target.nanos = delta.microseconds * 1000 + + +_EPOCH = datetime.datetime(1970, 1, 1, tzinfo=datetime.timezone.utc) + + +def rebuild_value_type(ctx, fd: FieldDescriptor, value: Any) -> Message: + """Rebuilds a leaf value-type message from what a CEL_FIELD rule returned. + + An identity rule hands back the message it was given; a computed rule hands back a Python + Decimal or datetime, which has to be encoded. Anything else is a rule-authoring mistake and + is reported as one rather than written as a default. + """ + desc = fd.message_type + if value is None: + raise _value_type_error(ctx, fd, "null", "a decimal or timestamp") + if isinstance(value, Message) and value.DESCRIPTOR.full_name == desc.full_name: + # Already the right message, which is what an identity rule produces. + return value + # A timestamp is bound as a datetime, which cannot hold the nanos it was read with, so an + # echoed one is copied from its source message instead of re-encoded. Same mechanism as the + # branch above; the difference is only that this binding converts rather than wrapping. + source = getattr(value, "msg", None) + if isinstance(source, Message) and source.DESCRIPTOR.full_name == desc.full_name: + return source + out = _message_factory(desc) + if desc.full_name == DECIMAL_TYPE_NAME: + if not isinstance(value, decimal.Decimal): + raise _value_type_error(ctx, fd, type(value).__name__, "a decimal") + set_decimal_message(out, value) + return out + if not isinstance(value, datetime.datetime): + raise _value_type_error(ctx, fd, type(value).__name__, "a timestamp") + set_timestamp_message(out, value) + return out + + +def _message_factory(desc: Descriptor) -> Message: + """A new message of *desc*'s type, built from the descriptor so that a message parsed + dynamically from a registered schema is written back in kind.""" + return message_factory.GetMessageClass(desc)() + + +def _value_type_error(ctx, fd: FieldDescriptor, actual: str, expected: str) -> Exception: + return RuleError( + "Rule returned " + + actual + + " for field '" + + fd.full_name + + "', which is a " + + fd.message_type.full_name + + "; expected " + + expected + ) def transform(ctx: RuleContext, descriptor: Descriptor, message: Any, field_transform: FieldTransform) -> Any: @@ -262,7 +403,7 @@ def transform(ctx: RuleContext, descriptor: Descriptor, message: Any, field_tran return [transform(ctx, descriptor, item, field_transform) for item in message] if isinstance(message, dict): return {key: transform(ctx, descriptor, value, field_transform) for key, value in message.items()} - if isinstance(message, Message): + if isinstance(message, Message) and not is_cel_leaf_message(message.DESCRIPTOR): # Driven by the runtime message's fields, each matched by name to the # schema-side descriptor, which is the one carrying the inline tags. The two # can differ under use.latest.version, and only the runtime field can be read @@ -315,6 +456,19 @@ def _transform_field( if new_value is False: raise RuleConditionError(ctx.rule) else: + if fd.type == FieldDescriptor.TYPE_MESSAGE and is_cel_leaf_message(fd.message_type): + # The rule saw this field as a single value, so it hands back a decimal or a + # datetime rather than the message; encode it before writing. + # + # A repeated leaf field needs the same treatment per element. The walk applies + # the rule to each element, so what comes back is a *list* of decimals - and + # writing those raw failed with "Expected a message object, but got + # Decimal(...)". Only the singular case was rebuilt before, so a field rule + # over a repeated value type could not be written back at all. + if _is_repeated(fd): + new_value = [rebuild_value_type(ctx, fd, item) for item in new_value] + else: + new_value = rebuild_value_type(ctx, fd, new_value) _set_field(fd, message, new_value) finally: ctx.exit_field() @@ -642,6 +796,10 @@ def get_type(fd: FieldDescriptor) -> FieldType: if is_map_field(fd): return FieldType.MAP if fd.type == FieldDescriptor.TYPE_MESSAGE: + # Report the same primitive type the Avro counterpart does, so that CEL_FIELD applies + # to the field and a rule written against one format ports to the other. + if is_cel_leaf_message(fd.message_type): + return FieldType.BYTES if fd.message_type.full_name == DECIMAL_TYPE_NAME else FieldType.LONG return FieldType.RECORD if fd.type == FieldDescriptor.TYPE_ENUM: return FieldType.ENUM @@ -700,6 +858,23 @@ def _is_builtin(name: str) -> bool: return name.startswith('confluent/') or name.startswith('google/protobuf/') or name.startswith('google/type/') +# Exact, with the exponent range widened: the default +/-999999 is narrower than the int32 +# scale a confluent.type.Decimal field permits, and the 28-digit default precision would +# silently round a wide unscaled value. +_EXACT_CONTEXT = Context(prec=MAX_PREC, rounding=ROUND_HALF_UP, Emax=MAX_EMAX, Emin=MIN_EMIN) + + +# The widest coefficient a BigDecimal can hold: BigInteger tops out at Integer.MAX_VALUE bits, +# which is 646456993 decimal digits, and setScale reports anything wider as "BigInteger would +# overflow supported range". Bisected against the JDK on BigDecimal("1.23"): setScale(1e8) and +# setScale(-1e8) succeed, setScale(646456993) and setScale(-1e9) do not. +# +# `rules/cel/decimal_funcs._quantize` bounds its own rescale by the same JDK limit for the same +# reason. The two cannot share one constant: this module needs the protobuf runtime, which is +# an optional extra, and that one has to import without it. +_MAX_BIGINTEGER_DIGITS = 646456993 + + def decimal_to_protobuf(value: Decimal, scale: int) -> decimal_pb2.Decimal: # type: ignore[name-defined] """ Converts a Decimal to a Protobuf value. @@ -715,25 +890,83 @@ def decimal_to_protobuf(value: Decimal, scale: int) -> decimal_pb2.Decimal: # t delta = exp + scale # type: ignore[operator] - if delta < 0: - raise ValueError("Scale provided does not match the decimal") - unscaled_datum = 0 for digit in digits: unscaled_datum = (unscaled_datum * 10) + digit - unscaled_datum = 10**delta * unscaled_datum - - bytes_req = (unscaled_datum.bit_length() + 8) // 8 + if delta >= 0: + # Widening: the coefficient grows by `delta` digits, and the JVM refuses a result + # wider than BigInteger can hold - instantly, where `10**delta` grinds first and then + # *succeeds*. Measured against the JDK and this function: setScale(1, 1e7) is accepted + # by both (1.4s there, 5s and a 4 MB field here), setScale(1, 1e9) throws + # "BigInteger would overflow supported range" there while here it ran past a 240s + # timeout still working towards a several-hundred-megabyte value. + if delta + len(digits) > _MAX_BIGINTEGER_DIGITS: + raise ValueError("Scale provided is too wide for the decimal") + unscaled_datum = 10**delta * unscaled_datum + else: + # Narrowing the scale, which BigDecimal.setScale(scale) allows whenever no rounding is + # needed - only the digits being dropped have to be zeros. Refusing every reduction + # rejected exact conversions: Decimal("1.50") at scale 1, or Decimal("1000") at the + # negative scale -3 that protobuf_to_decimal itself can produce. + # + # Whether the dropped digits are zeros is read off the digit tuple rather than + # discovered by dividing. Building a 10**-delta divisor just to find a non-zero + # remainder cost 178s at a scale of -1e8 and would run for hours at -1e9, to reach a + # rejection the trailing digits already prove. The JVM answers the same, reaching it + # through the division ("Rounding necessary"), so only the path changes. + drop = -delta + if unscaled_datum != 0: + trailing_zeros = 0 + for digit in reversed(digits): + if digit != 0: + break + trailing_zeros += 1 + if drop > trailing_zeros: + raise ValueError("Scale provided does not match the decimal") + # drop <= trailing_zeros <= len(digits) now, so the divisor is no wider than the + # coefficient already in hand. + unscaled_datum //= 10**drop + # A zero is the exception: it has no digits to lose, so it narrows to any scale. The + # JVM agrees and gets there without the division - setScale(-1e9) on BigDecimal("0") + # is exact and instant, while this function would have spent hours on the divisor. if sign: unscaled_datum = -unscaled_datum - bytes = unscaled_datum.to_bytes(bytes_req, byteorder="big", signed=True) + bytes = unscaled_to_bytes(unscaled_datum) result = decimal_pb2.Decimal() # type: ignore[attr-defined] result.value = bytes - result.precision = 0 + # The unscaled value's digit count, which is what BigDecimal.precision() reports and what + # every other write path in this client family carries. Left at 0 here, this was one of + # three paths whose output a JVM consumer rewrites on its next touch: `precision()` is + # never less than 1, so 0 is a value the reference cannot produce, and its reader + # normalises it away. + # + # Counted arithmetically rather than as `len(str(abs(unscaled_datum)))`. CPython caps + # str <-> int conversion at 4300 digits (`int_max_str_digits`), so the string form raised + # `ValueError: Exceeds the limit (4300 digits) for integer string conversion` for a + # coefficient this function otherwise accepts - and raised it *after* `result.value` was + # already assigned. `decimal_to_protobuf(Decimal("1"), 4300)` was the first failing case, + # against a `_MAX_COEFFICIENT_DIGITS` of 646456993 here. + # + # `len(digits) + delta` is exact in both directions, and only because this function + # refuses an inexact narrowing: widening multiplies by 10**delta, which appends `delta` + # zeros with no carry, and narrowing only ever drops digits already proven to be zeros. + # A rounding rescale could carry (9.9 to scale 0 is 10, one digit becoming two) and would + # need the count taken after the fact. Verified equal to the string form across 269 + # value/scale combinations. + # + # Zero is the exception, since its digit tuple is `(0,)` at every scale. The reference + # agrees: `new BigDecimal("0").setScale(5000).precision()` is 1. + # + # Measured on the JDK, which is what the count has to match: + # BigDecimal("1").setScale(4300) -> precision 4301 + # BigDecimal("1").setScale(5000) -> precision 5001 + # BigDecimal("1").setScale(1000000) -> precision 1000001 + # BigDecimal("1.50").setScale(1) -> precision 2 + result.precision = 1 if unscaled_datum == 0 else len(digits) + delta result.scale = scale return result @@ -750,5 +983,46 @@ def protobuf_to_decimal(value: decimal_pb2.Decimal) -> Decimal: # type: ignore[ """ unscaled_datum = int.from_bytes(value.value, byteorder="big", signed=True) - decimal_context = Context(prec=value.precision if value.precision > 0 else MAX_PREC) - return decimal_context.create_decimal(unscaled_datum).scaleb(-value.scale, decimal_context) + # `precision` is deliberately not applied. Java reads it as + # `new BigDecimal(unscaled, scale, new MathContext(precision))`, but every client - this one + # included - writes it as the unscaled value's own digit count, which makes that MathContext + # a guaranteed no-op. It has an effect only on a message from a foreign producer carrying a + # *declared column* precision, and there its effect is to silently round data the producer + # sent exactly. Six of the seven clients already ignore it; this path was the one that did + # not, so the same message read here and through the CEL binding gave two different values + # (unscaled 125 at precision 2: 1.3E+2 here, 125 there). + # + # Emax/Emin are widened because the default +/-999999 is narrower than the int32 scale this + # message's field permits. + return _EXACT_CONTEXT.create_decimal(unscaled_datum).scaleb(-value.scale, _EXACT_CONTEXT) + + +def variant_to_protobuf(value: Variant) -> variant_pb2.Variant: # type: ignore[name-defined] + """ + Converts a Variant to a ``confluent.type.Variant`` Protobuf message. + + Args: + value (Variant): The Variant to convert. + + Returns: + The Protobuf value. + """ + result = variant_pb2.Variant() # type: ignore[attr-defined] + result.metadata = value.metadata + # standalone_value_bytes, not .value: a navigated sub-variant's own value starts at its + # position, and .value is the whole shared buffer. + result.value = value.standalone_value_bytes() + return result + + +def protobuf_to_variant(value: variant_pb2.Variant) -> Variant: # type: ignore[name-defined] + """ + Converts a ``confluent.type.Variant`` Protobuf message to a Variant. + + Args: + value (variant_pb2.Variant): The Protobuf value to convert. + + Returns: + The Variant value. + """ + return Variant(value.value, value.metadata) diff --git a/src/confluent_kafka/schema_registry/confluent/codegen.sh b/src/confluent_kafka/schema_registry/confluent/codegen.sh new file mode 100755 index 000000000..4161c1610 --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/codegen.sh @@ -0,0 +1,48 @@ +#!/usr/bin/env bash +# Regenerates the Python bindings for the vendored confluent value types. +# Run from the repo root. Requires protoc 35.1 to match the checked-in headers. +# +# Two post-processing steps, both load-bearing: +# +# 1. The runtime-version guard is stripped. protoc emits an import of +# google.protobuf.runtime_version plus a ValidateProtobufRuntimeVersion call, but +# requirements-protobuf.txt pins no protobuf version, and that module does not exist on +# older runtimes - so the guard would make these modules fail to import at all. Every +# checked-in _pb2.py here has it stripped for that reason. +# 2. Relative imports are rewritten to fully-qualified ones. protoc emits +# `from confluent.type import decimal_pb2`, which only resolves if `confluent` is a +# top-level package; here it lives under confluent_kafka.schema_registry. +# +# confluent/types/decimal.proto is a public-import stub for the path confluent.type.Decimal +# used to occupy. Its generated module re-exports Decimal and registers the old file name, so +# code and descriptors built against the old path keep working. +set -euo pipefail + +PKG=src/confluent_kafka/schema_registry +ABS=confluent_kafka.schema_registry + +# confluent/meta.proto is deliberately not regenerated here: its checked-in module came from a +# much older protoc and this one rewrites the whole file, which is churn unrelated to the value +# types. Regenerate it on purpose, not as a side effect of touching decimal or variant. +cd "$PKG" +protoc -I. --python_out=. confluent/type/decimal.proto confluent/type/variant.proto \ + confluent/types/decimal.proto + +for f in confluent/type/decimal_pb2.py confluent/type/variant_pb2.py \ + confluent/types/decimal_pb2.py; do + python3 - "$f" "$ABS" <<'PY' +import re, sys +path, abs_pkg = sys.argv[1], sys.argv[2] +s = open(path).read() +s = s.replace('from google.protobuf import runtime_version as _runtime_version\n', '') +s = re.sub(r'_runtime_version\.ValidateProtobufRuntimeVersion\((?:[^)]*)\)\n', '', s) +# `from confluent.type import decimal_pb2 as X` -> `import .confluent.type.decimal_pb2 as X` +s = re.sub(r'^from (confluent[\w.]*) import (\w+) as (\w+)$', + lambda m: 'import %s.%s.%s as %s' % (abs_pkg, m.group(1), m.group(2), m.group(3)), + s, flags=re.M) +# `from confluent.type.decimal_pb2 import *` (public import re-export) -> fully qualified +s = re.sub(r'^from (confluent[\w.]*) import \*$', + lambda m: 'from %s.%s import *' % (abs_pkg, m.group(1)), s, flags=re.M) +open(path, 'w').write(s) +PY +done diff --git a/src/confluent_kafka/schema_registry/confluent/type/__init__.py b/src/confluent_kafka/schema_registry/confluent/type/__init__.py new file mode 100644 index 000000000..50582affa --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2024 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. diff --git a/src/confluent_kafka/schema_registry/confluent/type/decimal.proto b/src/confluent_kafka/schema_registry/confluent/type/decimal.proto new file mode 100644 index 000000000..559dde662 --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/decimal.proto @@ -0,0 +1,15 @@ +syntax = "proto3"; + +package confluent.type; + +message Decimal { + + // The two's-complement representation of the unscaled integer value in big-endian byte order + bytes value = 1; + + // The precision (zero indicates unlimited precision) + uint32 precision = 2; + + // The scale + int32 scale = 3; +} \ No newline at end of file diff --git a/src/confluent_kafka/schema_registry/confluent/type/decimal_pb2.py b/src/confluent_kafka/schema_registry/confluent/type/decimal_pb2.py new file mode 100644 index 000000000..076b3843b --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/decimal_pb2.py @@ -0,0 +1,29 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: confluent/type/decimal.proto +# Protobuf Python Version: 7.35.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\n\x1c\x63onfluent/type/decimal.proto\x12\x0e\x63onfluent.type\":\n\x07\x44\x65\x63imal\x12\r\n\x05value\x18\x01 \x01(\x0c\x12\x11\n\tprecision\x18\x02 \x01(\r\x12\r\n\x05scale\x18\x03 \x01(\x05\x62\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'confluent.type.decimal_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_DECIMAL']._serialized_start = 48 + _globals['_DECIMAL']._serialized_end = 106 +# @@protoc_insertion_point(module_scope) diff --git a/src/confluent_kafka/schema_registry/confluent/type/decimal_utils.py b/src/confluent_kafka/schema_registry/confluent/type/decimal_utils.py new file mode 100644 index 000000000..4a79ae26f --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/decimal_utils.py @@ -0,0 +1,84 @@ +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Conversions between :class:`decimal.Decimal` and the ``confluent.type.Decimal`` proto +message - the Python counterpart of Java's ``io.confluent.protobuf.type.utils.DecimalUtils`` +(BigDecimal) and C#'s ``DecimalExtensions`` (System.Decimal). Independent of CEL: the Protobuf +serde uses these for ``confluent.type.Decimal`` fields, and the CEL layer reuses them. +""" + +from decimal import MAX_EMAX, MAX_PREC, MIN_EMIN, Context, Decimal + +from confluent_kafka.schema_registry.confluent.type import decimal_pb2 + +# Java builds `new BigDecimal(unscaledValue, scale)`, which is exact. `scaleb` otherwise uses +# the ambient thread-local context (28 significant digits by default) and would silently round +# an unscaled value wider than that. +# Emax/Emin are widened too: the default +/-999999 is narrower than the int32 scale the +# message's field permits, and a wide scale raised decimal.Overflow out of the conversion. +_EXACT_CONTEXT = Context(prec=MAX_PREC, Emax=MAX_EMAX, Emin=MIN_EMIN) + + +def from_proto_decimal(msg: decimal_pb2.Decimal) -> Decimal: + """Convert a ``confluent.type.Decimal`` message to a :class:`decimal.Decimal`. + + ``value`` is the unscaled integer as big-endian two's-complement bytes; ``scale`` is the + number of fractional digits (the value is ``unscaled * 10**-scale``). + """ + scale = int(msg.scale) + if not msg.value: + return Decimal(0).scaleb(-scale, context=_EXACT_CONTEXT) + unscaled = int.from_bytes(msg.value, "big", signed=True) + return Decimal(unscaled).scaleb(-scale, context=_EXACT_CONTEXT) + + +def to_proto_decimal(d: Decimal) -> decimal_pb2.Decimal: + """Convert a :class:`decimal.Decimal` to a ``confluent.type.Decimal`` message. + + Mirrors Java ``BigDecimal.unscaledValue()``/``scale()``: the scale is the number of + fractional digits (negative for values like ``1E+2``) and the value is the unscaled + integer as big-endian two's-complement bytes. + """ + sign, digits, exponent = d.as_tuple() + if not isinstance(exponent, int): + raise ValueError(f"cannot convert non-finite Decimal '{d}' to confluent.type.Decimal") + scale = -exponent + unscaled = int("".join(map(str, digits)) or "0") + if sign: + unscaled = -unscaled + value = unscaled_to_bytes(unscaled) + # Precision is the unscaled value's digit count, as Java's DecimalUtils.fromBigDecimal sets + # it. Safe here because the scale is derived from the value rather than requested, so + # len(digits) is exactly the digit count of the unscaled value being written - a reader + # that treats precision as a MathContext cannot then round it or shift its scale. + return decimal_pb2.Decimal(value=value, precision=len(digits), scale=scale) + + +def unscaled_to_bytes(unscaled: int) -> bytes: + """Minimal big-endian two's-complement encoding of an unscaled integer, matching + ``BigInteger.toByteArray()``. The single implementation for every writer in this client.""" + return unscaled.to_bytes(_twos_complement_length(unscaled), "big", signed=True) + + +def _twos_complement_length(n: int) -> int: + """Byte length of the minimal big-endian two's-complement form of ``n``, matching + ``BigInteger.toByteArray()``. + + ``bit_length()`` ignores the sign, so deriving the length from it alone over-allocates by a + byte at every negative power of two that is exactly a signed boundary: -128 needs one byte + (0x80) but reports a bit length of 8. A negative value's magnitude is taken from ``~n``, + which is one less, and one bit is reserved for the sign in both cases. + """ + bits = (n.bit_length() if n >= 0 else (~n).bit_length()) + 1 + return max(1, (bits + 7) // 8) diff --git a/src/confluent_kafka/schema_registry/confluent/type/variant.proto b/src/confluent_kafka/schema_registry/confluent/type/variant.proto new file mode 100644 index 000000000..897a35061 --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/variant.proto @@ -0,0 +1,14 @@ +syntax = "proto3"; + +package confluent.type; + +message Variant { + + // A dictionary of field names used by all objects in the variant value. + // Encoded as: version header, dictionary size, offset list, and UTF-8 string data. + bytes metadata = 1; + + // The variant value, which can be a primitive, object, or array. + // Encoded as: a one-byte header (basic type + type info) followed by the content bytes. + bytes value = 2; +} diff --git a/src/confluent_kafka/schema_registry/confluent/type/variant_pb2.py b/src/confluent_kafka/schema_registry/confluent/type/variant_pb2.py new file mode 100644 index 000000000..b6b9c0fd3 --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/variant_pb2.py @@ -0,0 +1,29 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: confluent/type/variant.proto +# Protobuf Python Version: 7.35.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\n\x1c\x63onfluent/type/variant.proto\x12\x0e\x63onfluent.type\"*\n\x07Variant\x12\x10\n\x08metadata\x18\x01 \x01(\x0c\x12\r\n\x05value\x18\x02 \x01(\x0c\x62\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'confluent.type.variant_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_VARIANT']._serialized_start = 48 + _globals['_VARIANT']._serialized_end = 90 +# @@protoc_insertion_point(module_scope) diff --git a/src/confluent_kafka/schema_registry/confluent/type/variant_utils.py b/src/confluent_kafka/schema_registry/confluent/type/variant_utils.py new file mode 100644 index 000000000..09b77b426 --- /dev/null +++ b/src/confluent_kafka/schema_registry/confluent/type/variant_utils.py @@ -0,0 +1,1202 @@ +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Codec for the Spark/Parquet Variant binary type (a metadata key-dictionary plus a +self-describing value stream) - the Python counterpart of Java's +``io.confluent.kafka.schemaregistry.type`` ``Variant`` / ``VariantFormat`` / ``VariantUtils``. + +The binary decode/encode is ported from Apache Spark's ``pyspark.sql.variant_utils`` (which +itself derives from ``org.apache.spark.types.variant.VariantUtil``) and extended with the +Parquet Variant additions Spark lacks: ``TIME`` (17), ``TIMESTAMP_NANOS`` tz/ntz (18/19), and +``UUID`` (20). + +Two behaviors deliberately match the Confluent Java reference rather than Spark: + +* ``to_json`` renders temporal types as ISO-8601 with ``T``/``Z`` and the seconds field + always present (the cross-language contract), not Python ``str()``. +* ``parse_json`` follows Java ``VariantUtils.fromJsonNode`` number handling - a JSON + fractional number becomes a ``DOUBLE`` (never a decimal), matching a default Jackson + ``ObjectMapper``; only integers wider than 64 bits fall back to a scale-0 decimal. +""" + +import base64 +import datetime +import decimal +import json +import math +import struct +import uuid as uuid_mod +from enum import Enum +from typing import Any, Dict, List, Optional, Tuple + +# --------------------------------------------------------------------------- +# Format constants (see VariantFormat.java). +# --------------------------------------------------------------------------- + +BASIC_TYPE_BITS = 2 +BASIC_TYPE_MASK = 0x3 +TYPE_INFO_MASK = 0x3F +MAX_SHORT_STR_SIZE = 0x3F + +# Exact/unbounded context so scaling an unscaled value with >28 significant digits +# is not silently rounded by the thread-local default context (prec=28) — matches +# java.math.BigDecimal's exact scaleb/setScale semantics. +_EXACT_CONTEXT = decimal.Context(prec=decimal.MAX_PREC, Emax=decimal.MAX_EMAX, Emin=decimal.MIN_EMIN) + +# Basic types (low 2 bits of the header byte). +PRIMITIVE = 0 +SHORT_STR = 1 +OBJECT = 2 +ARRAY = 3 + +# Primitive type codes (upper 6 bits of the header byte when basic type == PRIMITIVE). +NULL = 0 +TRUE = 1 +FALSE = 2 +INT1 = 3 +INT2 = 4 +INT4 = 5 +INT8 = 6 +DOUBLE = 7 +DECIMAL4 = 8 +DECIMAL8 = 9 +DECIMAL16 = 10 +DATE = 11 +TIMESTAMP = 12 +TIMESTAMP_NTZ = 13 +FLOAT = 14 +BINARY = 15 +LONG_STR = 16 +TIME = 17 +TIMESTAMP_NANOS = 18 +TIMESTAMP_NANOS_NTZ = 19 +UUID = 20 + +VERSION = 1 +VERSION_MASK = 0x0F + +U8_MAX = 0xFF +U16_MAX = 0xFFFF +U24_MAX = 0xFFFFFF +U24_SIZE = 3 +U32_SIZE = 4 + +I8_MAX = 0x7F +I8_MIN = -0x80 +I16_MAX = 0x7FFF +I16_MIN = -0x8000 +I32_MAX = 0x7FFFFFFF +I32_MIN = -0x80000000 +I64_MAX = 0x7FFFFFFFFFFFFFFF +I64_MIN = -0x8000000000000000 + +UUID_SIZE = 16 + +MAX_DECIMAL4_PRECISION = 9 +MAX_DECIMAL4_VALUE = 10**MAX_DECIMAL4_PRECISION +MAX_DECIMAL8_PRECISION = 18 +MAX_DECIMAL8_VALUE = 10**MAX_DECIMAL8_PRECISION +MAX_DECIMAL16_PRECISION = 38 +MAX_DECIMAL16_VALUE = 10**MAX_DECIMAL16_PRECISION + +_EPOCH_DATE = datetime.date(1970, 1, 1) +_EPOCH_UTC = datetime.datetime(1970, 1, 1, tzinfo=datetime.timezone.utc) +_EPOCH_NAIVE = datetime.datetime(1970, 1, 1) + + +class VariantError(ValueError): + """Raised for a malformed or unsupported Variant binary value.""" + + +class VariantType(Enum): + """The value type of a Variant, mirroring Java ``Variant.Type``. + + Integer, decimal, and timestamp widths are kept distinct here (as in Java); the CEL + layer collapses them into coarse labels (int/decimal/timestamp) where appropriate. + """ + + OBJECT = "OBJECT" + ARRAY = "ARRAY" + NULL = "NULL" + BOOLEAN = "BOOLEAN" + BYTE = "BYTE" + SHORT = "SHORT" + INT = "INT" + LONG = "LONG" + STRING = "STRING" + DOUBLE = "DOUBLE" + DECIMAL4 = "DECIMAL4" + DECIMAL8 = "DECIMAL8" + DECIMAL16 = "DECIMAL16" + DATE = "DATE" + TIMESTAMP_TZ = "TIMESTAMP_TZ" + TIMESTAMP_NTZ = "TIMESTAMP_NTZ" + FLOAT = "FLOAT" + BINARY = "BINARY" + TIME = "TIME" + TIMESTAMP_NANOS_TZ = "TIMESTAMP_NANOS_TZ" + TIMESTAMP_NANOS_NTZ = "TIMESTAMP_NANOS_NTZ" + UUID = "UUID" + + +# --------------------------------------------------------------------------- +# Low-level byte helpers. +# --------------------------------------------------------------------------- + + +def _check_index(pos: int, length: int) -> None: + if pos < 0 or pos >= length: + raise VariantError("malformed variant: index out of bounds") + + +def _read_long(data: bytes, pos: int, num_bytes: int, signed: bool) -> int: + _check_index(pos, len(data)) + _check_index(pos + num_bytes - 1, len(data)) + return int.from_bytes(data[pos : pos + num_bytes], byteorder="little", signed=signed) + + +def _get_type_info(value: bytes, pos: int) -> Tuple[int, int]: + basic_type = value[pos] & BASIC_TYPE_MASK + type_info = (value[pos] >> BASIC_TYPE_BITS) & TYPE_INFO_MASK + return basic_type, type_info + + +def _get_metadata_key(metadata: bytes, key_id: int) -> str: + _check_index(0, len(metadata)) + offset_size = ((metadata[0] >> 6) & 0x3) + 1 + dict_size = _read_long(metadata, 1, offset_size, signed=False) + if key_id >= dict_size: + raise VariantError("malformed variant: field id out of range") + string_start = 1 + (dict_size + 2) * offset_size + offset = _read_long(metadata, 1 + (key_id + 1) * offset_size, offset_size, signed=False) + next_offset = _read_long(metadata, 1 + (key_id + 2) * offset_size, offset_size, signed=False) + if offset > next_offset: + raise VariantError("malformed variant: non-monotonic metadata offsets") + _check_index(string_start + next_offset - 1, len(metadata)) + return metadata[string_start + offset : string_start + next_offset].decode("utf-8") + + +# --------------------------------------------------------------------------- +# Cross-language JSON temporal contract. +# +# ISO-8601 with a 0/3/6/9-digit fractional-second grouping (as in Java Instant.toString()) +# and the seconds field ALWAYS present. UTC instants append 'Z'; NTZ/time forms omit the +# zone. This intentionally deviates from Java LocalDateTime/LocalTime.toString() (which omit +# the seconds field when both seconds and fraction are zero); the Java reference is aligned +# to always emit seconds so NTZ stays consistent with the TZ form. +# --------------------------------------------------------------------------- + + +def _frac_nanos(nanos: int) -> str: + """Fractional-second suffix using Java's 0/3/6/9-digit grouping (empty if zero).""" + if nanos == 0: + return "" + if nanos % 1_000_000 == 0: + return ".%03d" % (nanos // 1_000_000) + if nanos % 1_000 == 0: + return ".%06d" % (nanos // 1_000) + return ".%09d" % nanos + + +def _ymd_hms(total_nanos: int, tz: Optional[datetime.timezone]) -> Tuple[datetime.datetime, int]: + """Split epoch-nanos into a (whole-second datetime, nano-of-second). Uses floor + semantics so negative instants match Java's Math.floorDiv/floorMod.""" + epoch_sec, nano = divmod(total_nanos, 1_000_000_000) # Python divmod floors, like Java + base = _EPOCH_UTC if tz is not None else _EPOCH_NAIVE + return base + datetime.timedelta(seconds=epoch_sec), nano + + +def _format_instant(total_nanos: int) -> str: + """ISO-8601 instant with 'Z', seconds always present - matches Instant.toString().""" + dt, nano = _ymd_hms(total_nanos, datetime.timezone.utc) + return "%04d-%02d-%02dT%02d:%02d:%02d%sZ" % ( + dt.year, + dt.month, + dt.day, + dt.hour, + dt.minute, + dt.second, + _frac_nanos(nano), + ) + + +def _format_local_datetime(total_nanos: int) -> str: + """ISO local date-time, seconds always present. This is the cross-language contract: it + deviates from Java LocalDateTime.toString() (which omits the seconds field when both + seconds and fraction are zero) - the Java reference is aligned to always emit seconds, + keeping NTZ consistent with the TZ (Instant) form.""" + dt, nano = _ymd_hms(total_nanos, None) + return "%04d-%02d-%02dT%02d:%02d:%02d%s" % ( + dt.year, + dt.month, + dt.day, + dt.hour, + dt.minute, + dt.second, + _frac_nanos(nano), + ) + + +# The range a TIME may occupy, in microseconds since midnight: 00:00:00 through 23:59:59.999999. +# RFC 3339's partial-time requires time-hour = 2DIGIT in 00-23, so a value at or past 24 hours (or +# negative) has no valid form. A variant TIME is an int64 of microseconds, so those are reachable +# and are refused rather than rendered. +_MIN_TIME_MICROS = 0 +_MAX_TIME_MICROS = 86_400_000_000 - 1 + +# The range a DATE may occupy, in days since the epoch: 0001-01-01 through 9999-12-31. RFC 3339's +# full-date requires date-fullyear = 4DIGIT, so an expanded or negative year is not a valid +# full-date. (Python's `date` already refuses those, but with an opaque OverflowError/ValueError; +# checking explicitly keeps the message and the bound the same as every other client's.) +_MIN_DATE_EPOCH_DAY = -719162 +_MAX_DATE_EPOCH_DAY = 2932896 + + +def _check_date_range(epoch_day: int) -> int: + if epoch_day < _MIN_DATE_EPOCH_DAY or epoch_day > _MAX_DATE_EPOCH_DAY: + raise VariantError( + "date epoch day (%d) must be in range [%d, %d]" % (epoch_day, _MIN_DATE_EPOCH_DAY, _MAX_DATE_EPOCH_DAY) + ) + return epoch_day + + +def _format_local_time(micros: int) -> str: + """ISO local time, seconds always present (see :func:`_format_local_datetime`).""" + if micros < _MIN_TIME_MICROS or micros > _MAX_TIME_MICROS: + raise VariantError( + "time microseconds of day (%d) must be in range [%d, %d]" % (micros, _MIN_TIME_MICROS, _MAX_TIME_MICROS) + ) + nano_of_day = micros * 1000 + seconds, nano = divmod(nano_of_day, 1_000_000_000) + hour, rem = divmod(seconds, 3600) + minute, second = divmod(rem, 60) + return "%02d:%02d:%02d%s" % (hour, minute, second, _frac_nanos(nano)) + + +def _format_double(d: float) -> str: + """A JSON number rendering of a double. Non-finite values render as the BAREWORD tokens + ``NaN``/``Infinity``/``-Infinity`` (the Confluent Java contract, diverging from Spark + which quotes them). Integral values render as ``N.0``; other values use Python's shortest + round-tripping ``repr``. (Java Double.toString scientific-notation edge cases for very + large/small magnitudes are a known minor divergence.)""" + if d != d: + return "NaN" + if d == float("inf"): + return "Infinity" + if d == float("-inf"): + return "-Infinity" + # `int()` erases the sign of -0.0, which Java's Double.toString preserves ("-0.0"). + if d.is_integer() and abs(d) < 1e16 and not (d == 0.0 and math.copysign(1.0, d) < 0): + return "%d.0" % int(d) + return repr(d) + + +def _format_float(f: float) -> str: + """A JSON number rendering of a 32-bit float. Emits the shortest decimal that round-trips + to the SAME float32 (matching Java ``Float.toString`` / Apache Arrow) rather than the + f64-shortest string produced by widening then formatting as a double. Integral values + render as ``N.0``; mirrors :func:`_format_double`'s non-finite handling (bareword + ``NaN``/``Infinity``/``-Infinity``).""" + if f != f: + return "NaN" + if f == float("inf"): + return "Infinity" + if f == float("-inf"): + return "-Infinity" + if f == int(f) and abs(f) < 1e16 and not (f == 0.0 and math.copysign(1.0, f) < 0): + return "%d.0" % int(f) + for p in range(1, 10): + s = "%.*g" % (p, f) + try: + narrowed = struct.unpack(" bytes: + """The value bytes from this node's start - what any write-back has to use. + + ``self.value`` is the whole buffer, shared across sub-variants, so a navigated + variant's own value begins at ``self.pos``. Handing ``self.value`` to an encoder + writes the *parent root* rather than the selected value: measured, a field navigated + out of ``{"a":1,"secret":"..."}`` came back as the whole document. Every other client + has this accessor (Go ``StandaloneValueBytes``, C++ ``standaloneValueBytes``, Rust + ``standalone_value_bytes``, C# ``StandaloneValueBytes``), and Java's ``getValueBuffer`` + is a positioned ``ByteBuffer``. + + Like all of those, this slices to the end of the buffer rather than to the node's exact + extent, so a navigated value still carries its later siblings' bytes. Decoding ignores + them - the encoding is self-delimiting. + """ + return self.value[self.pos :] if self.pos else self.value + + # -- equality ----------------------------------------------------------- + + def __eq__(self, other: object) -> bool: + """Equality is over the encoding: the metadata bytes and the value bytes from ``pos``. + + The same comparison a ``confluent.type.Variant`` protobuf message already gets, so a + variant read from a field and one built by ``variants.parseJson`` answer the same way. + Slicing from ``pos`` is what a navigated variant needs - its position is the start of + the value, not of the parent's header. + """ + if self is other: + return True + if not isinstance(other, Variant): + return NotImplemented + return self.standalone_value_bytes() == other.standalone_value_bytes() and self.metadata == other.metadata + + def __hash__(self) -> int: + return hash((self.standalone_value_bytes(), self.metadata)) + + # -- type --------------------------------------------------------------- + + def get_type(self) -> VariantType: + _check_index(self.pos, len(self.value)) + basic_type, type_info = _get_type_info(self.value, self.pos) + if basic_type == SHORT_STR: + return VariantType.STRING + if basic_type == OBJECT: + return VariantType.OBJECT + if basic_type == ARRAY: + return VariantType.ARRAY + mapping = { + NULL: VariantType.NULL, + TRUE: VariantType.BOOLEAN, + FALSE: VariantType.BOOLEAN, + INT1: VariantType.BYTE, + INT2: VariantType.SHORT, + INT4: VariantType.INT, + INT8: VariantType.LONG, + DOUBLE: VariantType.DOUBLE, + DECIMAL4: VariantType.DECIMAL4, + DECIMAL8: VariantType.DECIMAL8, + DECIMAL16: VariantType.DECIMAL16, + DATE: VariantType.DATE, + TIMESTAMP: VariantType.TIMESTAMP_TZ, + TIMESTAMP_NTZ: VariantType.TIMESTAMP_NTZ, + FLOAT: VariantType.FLOAT, + BINARY: VariantType.BINARY, + LONG_STR: VariantType.STRING, + TIME: VariantType.TIME, + TIMESTAMP_NANOS: VariantType.TIMESTAMP_NANOS_TZ, + TIMESTAMP_NANOS_NTZ: VariantType.TIMESTAMP_NANOS_NTZ, + UUID: VariantType.UUID, + } + result = mapping.get(type_info) + if result is None: + raise VariantError("unknown variant primitive type: %d" % type_info) + return result + + # -- scalar getters ----------------------------------------------------- + + def _primitive_info(self) -> Tuple[int, int]: + _check_index(self.pos, len(self.value)) + basic_type, type_info = _get_type_info(self.value, self.pos) + if basic_type != PRIMITIVE: + raise VariantError("expected a primitive variant value") + return basic_type, type_info + + def get_boolean(self) -> bool: + _, type_info = self._primitive_info() + if type_info not in (TRUE, FALSE): + raise VariantError("variant is not a boolean") + return type_info == TRUE + + def get_byte(self) -> int: + """8-bit integer (``INT1`` only) - mirrors Java ``getByte``. Wider integer widths + raise; use :meth:`get_short`/:meth:`get_int`/:meth:`get_long` for those.""" + _, type_info = self._primitive_info() + if type_info == INT1: + return _read_long(self.value, self.pos + 1, 1, signed=True) + raise VariantError("variant is not a byte-width integer") + + def get_short(self) -> int: + """16-bit integer, widening from ``INT1`` (byte) - mirrors Java ``getShort``.""" + _, type_info = self._primitive_info() + if type_info == INT1: + return _read_long(self.value, self.pos + 1, 1, signed=True) + if type_info == INT2: + return _read_long(self.value, self.pos + 1, 2, signed=True) + raise VariantError("variant is not a short-width integer") + + def get_int(self) -> int: + """32-bit integer, widening from ``INT1``/``INT2`` - mirrors Java ``getInt``.""" + _, type_info = self._primitive_info() + if type_info == INT1: + return _read_long(self.value, self.pos + 1, 1, signed=True) + if type_info == INT2: + return _read_long(self.value, self.pos + 1, 2, signed=True) + if type_info == INT4: + return _read_long(self.value, self.pos + 1, 4, signed=True) + raise VariantError("variant is not an int-width integer") + + def get_long(self) -> int: + """Raw integer for any integer-backed type (byte/short/int/long, date days, + timestamp micros, time micros, timestamp-nanos) - mirrors Java ``getLong``.""" + _, type_info = self._primitive_info() + if type_info == INT1: + return _read_long(self.value, self.pos + 1, 1, signed=True) + if type_info == INT2: + return _read_long(self.value, self.pos + 1, 2, signed=True) + if type_info in (INT4, DATE): + return _read_long(self.value, self.pos + 1, 4, signed=True) + if type_info in (INT8, TIMESTAMP, TIMESTAMP_NTZ, TIME, TIMESTAMP_NANOS, TIMESTAMP_NANOS_NTZ): + return _read_long(self.value, self.pos + 1, 8, signed=True) + raise VariantError("variant is not an integer-backed type") + + def get_float(self) -> float: + """32-bit float (``FLOAT`` only, exact) - mirrors Java ``getFloat``. Note the + returned Python ``float`` is 64-bit, but the value is decoded from 4 bytes.""" + _, type_info = self._primitive_info() + if type_info == FLOAT: + _check_index(self.pos + 4, len(self.value)) + return struct.unpack(" float: + """64-bit double (``DOUBLE`` only, exact) - mirrors Java ``getDouble``. Does not + widen a ``FLOAT``; use :meth:`get_float` for that.""" + _, type_info = self._primitive_info() + if type_info == DOUBLE: + _check_index(self.pos + 8, len(self.value)) + return struct.unpack(" decimal.Decimal: + _, type_info = self._primitive_info() + scale = self.value[self.pos + 1] + if type_info == DECIMAL4: + unscaled = _read_long(self.value, self.pos + 2, 4, signed=True) + _check_decimal(unscaled, scale, MAX_DECIMAL4_VALUE, MAX_DECIMAL4_PRECISION) + elif type_info == DECIMAL8: + unscaled = _read_long(self.value, self.pos + 2, 8, signed=True) + _check_decimal(unscaled, scale, MAX_DECIMAL8_VALUE, MAX_DECIMAL8_PRECISION) + elif type_info == DECIMAL16: + _check_index(self.pos + 17, len(self.value)) + unscaled = int.from_bytes(self.value[self.pos + 2 : self.pos + 18], byteorder="little", signed=True) + _check_decimal(unscaled, scale, MAX_DECIMAL16_VALUE, MAX_DECIMAL16_PRECISION) + else: + raise VariantError("variant is not a decimal") + return decimal.Decimal(unscaled).scaleb(-scale, context=_EXACT_CONTEXT) + + def get_binary(self) -> bytes: + _, type_info = self._primitive_info() + if type_info != BINARY: + raise VariantError("variant is not binary") + length = _read_long(self.value, self.pos + 1, U32_SIZE, signed=False) + start = self.pos + 1 + U32_SIZE + _check_index(start + length - 1, len(self.value)) + return bytes(self.value[start : start + length]) + + def get_uuid(self) -> uuid_mod.UUID: + _, type_info = self._primitive_info() + if type_info != UUID: + raise VariantError("variant is not a uuid") + start = self.pos + 1 + _check_index(start + UUID_SIZE - 1, len(self.value)) + return uuid_mod.UUID(bytes=bytes(self.value[start : start + UUID_SIZE])) # big-endian + + def get_string(self) -> str: + _check_index(self.pos, len(self.value)) + basic_type, type_info = _get_type_info(self.value, self.pos) + if basic_type == SHORT_STR: + start = self.pos + 1 + length = type_info + elif basic_type == PRIMITIVE and type_info == LONG_STR: + length = _read_long(self.value, self.pos + 1, U32_SIZE, signed=False) + start = self.pos + 1 + U32_SIZE + else: + raise VariantError("variant is not a string") + _check_index(start + length - 1, len(self.value)) + return self.value[start : start + length].decode("utf-8") + + # -- object / array navigation ----------------------------------------- + + def _object_info(self) -> Tuple[int, int, int, int, int, int]: + _check_index(self.pos, len(self.value)) + basic_type, type_info = _get_type_info(self.value, self.pos) + if basic_type != OBJECT: + raise VariantError("variant is not an object") + large_size = ((type_info >> 4) & 0x1) != 0 + size_bytes = U32_SIZE if large_size else 1 + num_fields = _read_long(self.value, self.pos + 1, size_bytes, signed=False) + id_size = ((type_info >> 2) & 0x3) + 1 + offset_size = (type_info & 0x3) + 1 + id_start = self.pos + 1 + size_bytes + offset_start = id_start + num_fields * id_size + data_start = offset_start + (num_fields + 1) * offset_size + return num_fields, id_size, offset_size, id_start, offset_start, data_start + + def _array_info(self) -> Tuple[int, int, int, int]: + _check_index(self.pos, len(self.value)) + basic_type, type_info = _get_type_info(self.value, self.pos) + if basic_type != ARRAY: + raise VariantError("variant is not an array") + large_size = ((type_info >> 2) & 0x1) != 0 + size_bytes = U32_SIZE if large_size else 1 + num_fields = _read_long(self.value, self.pos + 1, size_bytes, signed=False) + offset_size = (type_info & 0x3) + 1 + offset_start = self.pos + 1 + size_bytes + data_start = offset_start + (num_fields + 1) * offset_size + return num_fields, offset_size, offset_start, data_start + + def num_object_fields(self) -> int: + return self._object_info()[0] + + def num_array_elements(self) -> int: + return self._array_info()[0] + + def _field_id_and_offset(self, idx: int) -> Tuple[int, int]: + num_fields, id_size, offset_size, id_start, offset_start, data_start = self._object_info() + key_id = _read_long(self.value, id_start + id_size * idx, id_size, signed=False) + offset = _read_long(self.value, offset_start + offset_size * idx, offset_size, signed=False) + return key_id, data_start + offset + + def get_field_by_key(self, key: str) -> Optional["Variant"]: + """Returns the object field with the given key, or ``None`` if absent. Linear scan + for small objects, binary search past the threshold (fields are key-sorted).""" + num_fields, id_size, offset_size, id_start, offset_start, data_start = self._object_info() + if num_fields < self._BINARY_SEARCH_THRESHOLD: + for i in range(num_fields): + key_id = _read_long(self.value, id_start + id_size * i, id_size, signed=False) + if _get_metadata_key(self.metadata, key_id) == key: + offset = _read_long(self.value, offset_start + offset_size * i, offset_size, signed=False) + return Variant(self.value, self.metadata, data_start + offset) + return None + low, high = 0, num_fields - 1 + while low <= high: + mid = (low + high) >> 1 + mid_id = _read_long(self.value, id_start + id_size * mid, id_size, signed=False) + mid_key = _get_metadata_key(self.metadata, mid_id) + if mid_key < key: + low = mid + 1 + elif mid_key > key: + high = mid - 1 + else: + offset = _read_long(self.value, offset_start + offset_size * mid, offset_size, signed=False) + return Variant(self.value, self.metadata, data_start + offset) + return None + + def get_field_at_index(self, idx: int) -> Tuple[str, "Variant"]: + """Returns the (key, value) of the field at ``idx`` (fields are key-sorted).""" + key_id, value_pos = self._field_id_and_offset(idx) + return _get_metadata_key(self.metadata, key_id), Variant(self.value, self.metadata, value_pos) + + def get_element_at_index(self, index: int) -> Optional["Variant"]: + """Returns the array element at ``index``, or ``None`` if out of bounds.""" + num_fields, offset_size, offset_start, data_start = self._array_info() + if index < 0 or index >= num_fields: + return None + offset = _read_long(self.value, offset_start + offset_size * index, offset_size, signed=False) + return Variant(self.value, self.metadata, data_start + offset) + + # -- JSON --------------------------------------------------------------- + + def to_json(self) -> str: + """Serialize to a JSON string, matching Java ``VariantUtils.toJsonString``.""" + t = self.get_type() + if t == VariantType.OBJECT: + parts = [] + for i in range(self.num_object_fields()): + key, child = self.get_field_at_index(i) + parts.append(json.dumps(key, ensure_ascii=False) + ":" + child.to_json()) + return "{" + ",".join(parts) + "}" + if t == VariantType.ARRAY: + parts = [] + for i in range(self.num_array_elements()): + element = self.get_element_at_index(i) + if element is None: + raise VariantError("array element count exceeds the encoded elements") + parts.append(element.to_json()) + return "[" + ",".join(parts) + "]" + if t == VariantType.NULL: + return "null" + if t == VariantType.BOOLEAN: + return "true" if self.get_boolean() else "false" + if t == VariantType.STRING: + return json.dumps(self.get_string(), ensure_ascii=False) + if t in (VariantType.BYTE, VariantType.SHORT, VariantType.INT, VariantType.LONG): + return str(self.get_long()) + if t == VariantType.FLOAT: + return _format_float(self.get_float()) + if t == VariantType.DOUBLE: + return _format_double(self.get_double()) + if t in (VariantType.DECIMAL4, VariantType.DECIMAL8, VariantType.DECIMAL16): + # Fixed-point (never scientific), matching Java's toPlainString contract. + return format(self.get_decimal(), "f") + if t == VariantType.DATE: + return '"' + (_EPOCH_DATE + datetime.timedelta(days=_check_date_range(self.get_long()))).isoformat() + '"' + if t == VariantType.TIMESTAMP_TZ: + return '"' + _format_instant(self.get_long() * 1000) + '"' + if t == VariantType.TIMESTAMP_NTZ: + return '"' + _format_local_datetime(self.get_long() * 1000) + '"' + if t == VariantType.TIMESTAMP_NANOS_TZ: + return '"' + _format_instant(self.get_long()) + '"' + if t == VariantType.TIMESTAMP_NANOS_NTZ: + return '"' + _format_local_datetime(self.get_long()) + '"' + if t == VariantType.TIME: + return '"' + _format_local_time(self.get_long()) + '"' + if t == VariantType.BINARY: + return '"' + base64.b64encode(self.get_binary()).decode("ascii") + '"' + if t == VariantType.UUID: + return '"' + str(self.get_uuid()) + '"' + raise VariantError("unsupported variant type for JSON: %s" % t) + + +def _check_decimal(unscaled: int, scale: int, max_unscaled: int, max_scale: int) -> None: + if unscaled >= max_unscaled or unscaled <= -max_unscaled or scale > max_scale: + raise VariantError("malformed variant: decimal out of range") + + +# --------------------------------------------------------------------------- +# Module-level convenience API. +# --------------------------------------------------------------------------- + + +def from_bytes(value: bytes, metadata: bytes) -> Variant: + """Construct a Variant from its raw ``value`` + ``metadata`` byte strings.""" + return Variant(value, metadata) + + +def to_json_string(variant: Variant) -> str: + """Serialize a Variant to its JSON string form.""" + return variant.to_json() + + +def parse_json(json_str: str) -> Variant: + """Parse a JSON string into a Variant, matching Java ``VariantUtils.fromJsonNode``.""" + builder = VariantBuilder() + # Default float parsing (no parse_float=Decimal): a JSON fractional number becomes a + # Python float and is written as a DOUBLE, matching a default Jackson ObjectMapper. + builder._process_parsed_json(json.loads(json_str)) + value, metadata = builder._finalize() + return Variant(value, metadata) + + +# --------------------------------------------------------------------------- +# Variant builder, ported from Spark's VariantBuilder and exposed as a flat +# streaming writer (arrow-dotnet ``VariantValueWriter`` shape): a single object +# with an internal nesting stack. Each scalar/container append fills the "current +# slot" - the root, the next array element, or the current object field's value +# (after :meth:`append_key`). Object fields are sorted by key on :meth:`end_object` +# (canonical form); the metadata dictionary accumulates every key seen. +# +# The same internal machinery drives :func:`parse_json`: :meth:`_process_parsed_json` +# walks a parsed JSON tree using the same low-level writers and object/array +# finishers, so a programmatic build is byte-identical to ``parse_json`` of an +# equivalent value. +# +# Number handling in the JSON path follows Java VariantUtils.fromJsonNode (default +# Jackson ObjectMapper): a JSON fractional number becomes a DOUBLE (never a decimal); +# an integer becomes the smallest int1/2/4/8 that fits, or a scale-0 decimal when +# wider than 64 bits. +# --------------------------------------------------------------------------- + + +class _FieldEntry: + __slots__ = ("key", "id", "offset") + + def __init__(self, key: str, field_id: int, offset: int): + self.key = key + self.id = field_id + self.offset = offset + + +class _ObjectContext: + """Nesting-stack frame for an in-progress object.""" + + __slots__ = ("start", "fields", "pending_key", "pending_id", "has_pending_key") + + def __init__(self, start: int): + self.start = start + self.fields: List[_FieldEntry] = [] + self.pending_key: Optional[str] = None + self.pending_id = 0 + self.has_pending_key = False + + +class _ArrayContext: + """Nesting-stack frame for an in-progress array.""" + + __slots__ = ("start", "offsets") + + def __init__(self, start: int): + self.start = start + self.offsets: List[int] = [] + + +class VariantBuilder: + """A flat streaming writer for Variant values, with an internal nesting stack. + + Scalars (``append_*``) and containers (``start_object``/``start_array``) each fill + the current slot: the root, the next array element, or the value of the current + object field (set with :meth:`append_key`). Call :meth:`build` to obtain the + finished :class:`Variant`. + """ + + DEFAULT_SIZE_LIMIT = 16 * 1024 * 1024 + + def __init__(self, size_limit: int = DEFAULT_SIZE_LIMIT): + self.value = bytearray() + self.dictionary: Dict[str, int] = {} + self.dictionary_keys: List[bytes] = [] + self.size_limit = size_limit + self._stack: List[Any] = [] + self._root_written = False + + # -- public streaming API ---------------------------------------------- + + def build(self) -> Variant: + """Finalize and return the built :class:`Variant`. Raises if a container is + still open or nothing has been written.""" + if self._stack: + raise VariantError("cannot build with an open container") + if not self._root_written: + raise VariantError("cannot build an empty variant") + value, metadata = self._finalize() + return Variant(value, metadata) + + def append_null(self) -> None: + self._before_append() + self._append_null() + + def append_boolean(self, b: bool) -> None: + self._before_append() + self._append_boolean(bool(b)) + + def append_byte(self, value: int) -> None: + """Append an 8-bit integer (``INT1``).""" + self._before_append() + self._write_fixed_int(INT1, value, 1) + + def append_short(self, value: int) -> None: + """Append a 16-bit integer (``INT2``).""" + self._before_append() + self._write_fixed_int(INT2, value, 2) + + def append_int(self, value: int) -> None: + """Append a 32-bit integer (``INT4``).""" + self._before_append() + self._write_fixed_int(INT4, value, 4) + + def append_long(self, value: int) -> None: + """Append a 64-bit integer (``INT8``).""" + self._before_append() + self._write_fixed_int(INT8, value, 8) + + def append_float(self, value: float) -> None: + """Append a 32-bit float (``FLOAT``).""" + self._before_append() + self._check_capacity(1 + 4) + self.value.append(self._primitive_header(FLOAT)) + self.value.extend(struct.pack(" None: + """Append a 64-bit double (``DOUBLE``).""" + self._before_append() + self._append_double(value) + + def append_decimal(self, unscaled: Any, scale: Optional[int] = None) -> None: + """Append a decimal. Either ``append_decimal(unscaled_big_endian_bytes, scale)`` + with a big-endian two's-complement unscaled value, or the native overload + ``append_decimal(decimal.Decimal)`` (scale taken from the value).""" + self._before_append() + if isinstance(unscaled, (bytes, bytearray)): + if scale is None: + raise VariantError("scale is required when appending a decimal from bytes") + unscaled_int = int.from_bytes(bytes(unscaled), byteorder="big", signed=True) + self._write_decimal(unscaled_int, scale) + elif isinstance(unscaled, decimal.Decimal): + if scale is not None: + raise VariantError("scale must not be given with a Decimal value") + self._append_decimal(unscaled) + elif isinstance(unscaled, int): + if scale is None: + raise VariantError("scale is required when appending an unscaled integer") + self._write_decimal(unscaled, scale) + else: + raise VariantError("invalid append_decimal arguments") + + def append_string(self, s: str) -> None: + self._before_append() + self._append_string(s) + + def append_binary(self, data: bytes) -> None: + self._before_append() + data = bytes(data) + self._check_capacity(1 + U32_SIZE + len(data)) + self.value.append(self._primitive_header(BINARY)) + self.value.extend(len(data).to_bytes(U32_SIZE, byteorder="little")) + self.value.extend(data) + + def append_uuid(self, value: Any) -> None: + """Append a UUID. Accepts a :class:`uuid.UUID` or 16 big-endian bytes.""" + self._before_append() + if isinstance(value, uuid_mod.UUID): + raw = value.bytes # big-endian + else: + raw = bytes(value) + if len(raw) != UUID_SIZE: + raise VariantError("uuid must be 16 bytes") + self._check_capacity(1 + UUID_SIZE) + self.value.append(self._primitive_header(UUID)) + self.value.extend(raw) + + def append_date(self, days_since_epoch: int) -> None: + self._before_append() + self._write_fixed_int(DATE, days_since_epoch, 4) + + def append_time(self, micros_since_midnight: int) -> None: + """Append a ``TIME`` (TIME_NTZ) as microseconds since midnight.""" + self._before_append() + self._write_fixed_int(TIME, micros_since_midnight, 8) + + def append_timestamp_tz(self, micros: int) -> None: + self._before_append() + self._write_fixed_int(TIMESTAMP, micros, 8) + + def append_timestamp_ntz(self, micros: int) -> None: + self._before_append() + self._write_fixed_int(TIMESTAMP_NTZ, micros, 8) + + def append_timestamp_nanos_tz(self, nanos: int) -> None: + self._before_append() + self._write_fixed_int(TIMESTAMP_NANOS, nanos, 8) + + def append_timestamp_nanos_ntz(self, nanos: int) -> None: + self._before_append() + self._write_fixed_int(TIMESTAMP_NANOS_NTZ, nanos, 8) + + def start_object(self) -> None: + self._before_append() + self._stack.append(_ObjectContext(len(self.value))) + + def append_key(self, key: str) -> None: + if not self._stack or not isinstance(self._stack[-1], _ObjectContext): + raise VariantError("append_key called outside of an object") + ctx = self._stack[-1] + if ctx.has_pending_key: + raise VariantError("append_key called twice without an intervening value") + ctx.pending_key = key + ctx.pending_id = self._add_key(key) + ctx.has_pending_key = True + + def end_object(self) -> None: + if not self._stack or not isinstance(self._stack[-1], _ObjectContext): + raise VariantError("end_object without a matching start_object") + ctx = self._stack.pop() + if ctx.has_pending_key: + raise VariantError("end_object with a dangling append_key (no value)") + self._finish_writing_object(ctx.start, ctx.fields) + + def start_array(self) -> None: + self._before_append() + self._stack.append(_ArrayContext(len(self.value))) + + def end_array(self) -> None: + if not self._stack or not isinstance(self._stack[-1], _ArrayContext): + raise VariantError("end_array without a matching start_array") + ctx = self._stack.pop() + self._finish_writing_array(ctx.start, ctx.offsets) + + # -- current-slot bookkeeping ------------------------------------------ + + def _before_append(self) -> None: + """Register the slot that the value about to be written will occupy, recording its + offset in the enclosing container (or marking the root as written).""" + if not self._stack: + if self._root_written: + raise VariantError("cannot append multiple root values") + self._root_written = True + return + ctx = self._stack[-1] + if isinstance(ctx, _ObjectContext): + if not ctx.has_pending_key or ctx.pending_key is None: + raise VariantError("a value in an object must follow append_key") + ctx.fields.append(_FieldEntry(ctx.pending_key, ctx.pending_id, len(self.value) - ctx.start)) + ctx.pending_key = None + ctx.has_pending_key = False + else: # _ArrayContext + ctx.offsets.append(len(self.value) - ctx.start) + + # -- metadata finalization --------------------------------------------- + + def _finalize(self) -> Tuple[bytes, bytes]: + num_keys = len(self.dictionary_keys) + dictionary_string_size = sum(len(k) for k in self.dictionary_keys) + max_size = max(dictionary_string_size, num_keys) + if max_size > self.size_limit: + raise VariantError("variant size limit exceeded") + offset_size = _integer_size(max_size) + + offset_start = 1 + offset_size + string_start = offset_start + (num_keys + 1) * offset_size + if string_start + dictionary_string_size > self.size_limit: + raise VariantError("variant size limit exceeded") + + metadata = bytearray() + metadata.append(VERSION | ((offset_size - 1) << 6)) + metadata.extend(num_keys.to_bytes(offset_size, byteorder="little")) + current_offset = 0 + for key in self.dictionary_keys: + metadata.extend(current_offset.to_bytes(offset_size, byteorder="little")) + current_offset += len(key) + metadata.extend(current_offset.to_bytes(offset_size, byteorder="little")) + for key in self.dictionary_keys: + metadata.extend(key) + return bytes(self.value), bytes(metadata) + + # -- internal JSON-tree driver (used by parse_json) -------------------- + + def _process_parsed_json(self, parsed: Any) -> None: + if isinstance(parsed, dict): + fields = [] + start = len(self.value) + for key, val in parsed.items(): + field_id = self._add_key(key) + fields.append(_FieldEntry(key, field_id, len(self.value) - start)) + self._process_parsed_json(val) + self._finish_writing_object(start, fields) + elif isinstance(parsed, list): + offsets = [] + start = len(self.value) + for elem in parsed: + offsets.append(len(self.value) - start) + self._process_parsed_json(elem) + self._finish_writing_array(start, offsets) + elif isinstance(parsed, str): + self._append_string(parsed) + elif isinstance(parsed, bool): + # bool must precede int (bool is a subclass of int in Python). + self._append_boolean(parsed) + elif isinstance(parsed, int): + if not self._append_int(parsed): + # Wider than 64 bits: a scale-0 decimal, matching Java's BigInteger branch. + self._append_decimal(decimal.Decimal(parsed)) + elif isinstance(parsed, float): + self._append_double(parsed) + elif isinstance(parsed, decimal.Decimal): + self._append_decimal(parsed) + elif parsed is None: + self._append_null() + else: + raise VariantError("unsupported JSON value: %r" % type(parsed)) + + def _check_capacity(self, additional: int) -> None: + if len(self.value) + additional > self.size_limit: + raise VariantError("variant size limit exceeded") + + @staticmethod + def _primitive_header(type_code: int) -> int: + return (type_code << 2) | PRIMITIVE + + @staticmethod + def _short_string_header(size: int) -> int: + return (size << 2) | SHORT_STR + + @staticmethod + def _array_header(large_size: bool, offset_size: int) -> int: + return (int(large_size) << (BASIC_TYPE_BITS + 2)) | ((offset_size - 1) << BASIC_TYPE_BITS) | ARRAY + + @staticmethod + def _object_header(large_size: bool, id_size: int, offset_size: int) -> int: + return ( + (int(large_size) << (BASIC_TYPE_BITS + 4)) + | ((id_size - 1) << (BASIC_TYPE_BITS + 2)) + | ((offset_size - 1) << BASIC_TYPE_BITS) + | OBJECT + ) + + def _add_key(self, key: str) -> int: + if key in self.dictionary: + return self.dictionary[key] + field_id = len(self.dictionary_keys) + self.dictionary[key] = field_id + self.dictionary_keys.append(key.encode("utf-8")) + return field_id + + def _append_boolean(self, b: bool) -> None: + self._check_capacity(1) + self.value.append(self._primitive_header(TRUE if b else FALSE)) + + def _append_null(self) -> None: + self._check_capacity(1) + self.value.append(self._primitive_header(NULL)) + + def _append_string(self, s: str) -> None: + text = s.encode("utf-8") + long_str = len(text) > MAX_SHORT_STR_SIZE + self._check_capacity((1 + U32_SIZE if long_str else 1) + len(text)) + if long_str: + self.value.append(self._primitive_header(LONG_STR)) + self.value.extend(len(text).to_bytes(U32_SIZE, byteorder="little")) + else: + self.value.append(self._short_string_header(len(text))) + self.value.extend(text) + + def _append_int(self, i: int) -> bool: + # Capacity is checked against the width actually chosen, not the widest one: Java's + # appendByte/appendShort/appendInt/appendLong each check `1 + their own width`, so + # reserving 9 bytes up front would reject a one-byte write that fits. + if I8_MIN <= i <= I8_MAX: + code, width = INT1, 1 + elif I16_MIN <= i <= I16_MAX: + code, width = INT2, 2 + elif I32_MIN <= i <= I32_MAX: + code, width = INT4, 4 + elif I64_MIN <= i <= I64_MAX: + code, width = INT8, 8 + else: + return False + self._check_capacity(1 + width) + self.value.append(self._primitive_header(code)) + self.value.extend(i.to_bytes(width, byteorder="little", signed=True)) + return True + + def _write_fixed_int(self, type_code: int, value: int, width: int) -> None: + """Write a fixed-width signed little-endian integer primitive.""" + self._check_capacity(1 + width) + try: + payload = int(value).to_bytes(width, byteorder="little", signed=True) + except OverflowError: + raise VariantError("integer value out of range for a %d-byte width" % width) + self.value.append(self._primitive_header(type_code)) + self.value.extend(payload) + + def _write_decimal(self, unscaled: int, scale: int) -> None: + """Write a decimal primitive from an unscaled integer and scale, choosing the + smallest of DECIMAL4/8/16 that fits.""" + if scale < 0: + raise VariantError("cannot encode decimal with negative scale") + # A 38-digit coefficient needs at most 127 bits, so anything wider is out of range + # without rendering it. str() on a wider one raises CPython's own 4300-digit + # ValueError, which names an interpreter limit rather than this API's contract. + if unscaled.bit_length() > 128: + raise VariantError("decimal exceeds maximum precision (38)") + precision = len(str(abs(unscaled))) + if scale <= MAX_DECIMAL4_PRECISION and precision <= MAX_DECIMAL4_PRECISION: + code, width = DECIMAL4, 4 + elif scale <= MAX_DECIMAL8_PRECISION and precision <= MAX_DECIMAL8_PRECISION: + code, width = DECIMAL8, 8 + elif scale <= MAX_DECIMAL16_PRECISION and precision <= MAX_DECIMAL16_PRECISION: + code, width = DECIMAL16, 16 + else: + raise VariantError("decimal exceeds maximum precision (38)") + # Java's appendDecimal checks `2 + width` inside each branch, after the width is known. + self._check_capacity(2 + width) + self.value.append(self._primitive_header(code)) + self.value.append(scale) + self.value.extend(unscaled.to_bytes(width, byteorder="little", signed=True)) + + def _append_decimal(self, d: decimal.Decimal) -> None: + sign, digits, exponent = d.as_tuple() + if not isinstance(exponent, int): + raise VariantError("cannot encode non-finite decimal") + # Before int(), not after: the digit string for a coefficient past CPython's 4300-digit + # limit raises its ValueError, which is not the error this API documents. The digit + # count is already in hand, so the encoding's own limit is the cheaper check. + if len(digits) > MAX_DECIMAL16_PRECISION: + raise VariantError("decimal exceeds maximum precision (38)") + unscaled = int("".join(map(str, digits)) or "0") + if sign: + unscaled = -unscaled + self._write_decimal(unscaled, -exponent) + + def _append_double(self, f: float) -> None: + self._check_capacity(1 + 8) + self.value.append(self._primitive_header(DOUBLE)) + self.value.extend(struct.pack(" None: + data_size = len(self.value) - start + num_offsets = len(offsets) + large_size = num_offsets > U8_MAX + size_bytes = U32_SIZE if large_size else 1 + offset_size = _integer_size(data_size) + header_size = 1 + size_bytes + (num_offsets + 1) * offset_size + self._check_capacity(header_size) + self.value.extend(bytearray(header_size)) + self.value[start + header_size :] = bytes(self.value[start : start + data_size]) + offset_start = start + 1 + size_bytes + self.value[start : start + 1] = bytes([self._array_header(large_size, offset_size)]) + self.value[start + 1 : offset_start] = num_offsets.to_bytes(size_bytes, byteorder="little") + offset_list = bytearray() + for offset in offsets: + offset_list.extend(offset.to_bytes(offset_size, byteorder="little")) + offset_list.extend(data_size.to_bytes(offset_size, byteorder="little")) + self.value[offset_start : offset_start + len(offset_list)] = offset_list + + def _finish_writing_object(self, start: int, fields: List[_FieldEntry]) -> None: + num_fields = len(fields) + fields.sort(key=lambda f: f.key) + max_id = max((f.id for f in fields), default=0) + data_size = len(self.value) - start + large_size = num_fields > U8_MAX + size_bytes = U32_SIZE if large_size else 1 + id_size = _integer_size(max_id) + offset_size = _integer_size(data_size) + header_size = 1 + size_bytes + num_fields * id_size + (num_fields + 1) * offset_size + self._check_capacity(header_size) + self.value.extend(bytearray(header_size)) + self.value[start + header_size :] = bytes(self.value[start : start + data_size]) + self.value[start : start + 1] = bytes([self._object_header(large_size, id_size, offset_size)]) + self.value[start + 1 : start + 1 + size_bytes] = num_fields.to_bytes(size_bytes, byteorder="little") + id_start = start + 1 + size_bytes + offset_start = id_start + num_fields * id_size + id_list = bytearray() + offset_list = bytearray() + for field in fields: + id_list.extend(field.id.to_bytes(id_size, byteorder="little")) + offset_list.extend(field.offset.to_bytes(offset_size, byteorder="little")) + offset_list.extend(data_size.to_bytes(offset_size, byteorder="little")) + self.value[id_start : id_start + len(id_list)] = id_list + self.value[offset_start : offset_start + len(offset_list)] = offset_list + + +def _integer_size(value: int) -> int: + if value <= U8_MAX: + return 1 + if value <= U16_MAX: + return 2 + if value <= U24_MAX: + return U24_SIZE + return U32_SIZE diff --git a/src/confluent_kafka/schema_registry/confluent/types/__init__.py b/src/confluent_kafka/schema_registry/confluent/types/__init__.py index 50582affa..92d632767 100644 --- a/src/confluent_kafka/schema_registry/confluent/types/__init__.py +++ b/src/confluent_kafka/schema_registry/confluent/types/__init__.py @@ -11,3 +11,11 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. + +# Deprecated: import from confluent_kafka.schema_registry.confluent.type instead. +# +# This package held the generated confluent.type.Decimal bindings until they moved to the +# canonical confluent/type path - the one the Java client registers and ProtobufSchema declares. +# What remains is generated from confluent/types/decimal.proto, a stub that declares nothing and +# publicly imports the canonical file, so Decimal stays importable under its old name and a +# descriptor built against the old import path still resolves. diff --git a/src/confluent_kafka/schema_registry/confluent/types/decimal.proto b/src/confluent_kafka/schema_registry/confluent/types/decimal.proto index 9abfeea5c..aaa622076 100644 --- a/src/confluent_kafka/schema_registry/confluent/types/decimal.proto +++ b/src/confluent_kafka/schema_registry/confluent/types/decimal.proto @@ -1,17 +1,7 @@ syntax = "proto3"; -package confluent.type; - -option go_package="../types"; - -message Decimal { - - // The two's-complement representation of the unscaled integer value in big-endian byte order - bytes value = 1; - - // The precision (zero indicates unlimited precision) - uint32 precision = 2; - - // The scale - int32 scale = 3; -} \ No newline at end of file +// The path this file used to occupy, kept so a descriptor generated against the old import path +// still resolves. It declares nothing and re-exports the canonical file, so confluent.type.Decimal +// is visible through it without a second declaration of the symbol - which a descriptor pool +// refuses ("duplicate symbol"). Read-only compatibility: the clients emit the canonical path. +import public "confluent/type/decimal.proto"; diff --git a/src/confluent_kafka/schema_registry/confluent/types/decimal_pb2.py b/src/confluent_kafka/schema_registry/confluent/types/decimal_pb2.py index 4c4cf812d..f7eae16c6 100644 --- a/src/confluent_kafka/schema_registry/confluent/types/decimal_pb2.py +++ b/src/confluent_kafka/schema_registry/confluent/types/decimal_pb2.py @@ -1,6 +1,8 @@ # -*- coding: utf-8 -*- # Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE # source: confluent/types/decimal.proto +# Protobuf Python Version: 7.35.1 """Generated protocol buffer code.""" from google.protobuf import descriptor as _descriptor @@ -13,16 +15,16 @@ _sym_db = _symbol_database.Default() +import confluent_kafka.schema_registry.confluent.type.decimal_pb2 as confluent_dot_type_dot_decimal__pb2 +from confluent_kafka.schema_registry.confluent.type.decimal_pb2 import * + DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( - b'\n\x1d\x63onfluent/types/decimal.proto\x12\x0e\x63onfluent.type\":\n\x07\x44\x65\x63imal\x12\r\n\x05value\x18\x01 \x01(\x0c\x12\x11\n\tprecision\x18\x02 \x01(\r\x12\r\n\x05scale\x18\x03 \x01(\x05\x42\nZ\x08../typesb\x06proto3' + b'\n\x1d\x63onfluent/types/decimal.proto\x1a\x1c\x63onfluent/type/decimal.protoP\x00\x62\x06proto3' ) -_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, globals()) -_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'confluent.types.decimal_pb2', globals()) -if _descriptor._USE_C_DESCRIPTORS == False: - - DESCRIPTOR._options = None - DESCRIPTOR._serialized_options = b'Z\010../types' - _DECIMAL._serialized_start = 49 - _DECIMAL._serialized_end = 107 +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'confluent.types.decimal_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None # @@protoc_insertion_point(module_scope) diff --git a/src/confluent_kafka/schema_registry/rules/cel/cel_executor.py b/src/confluent_kafka/schema_registry/rules/cel/cel_executor.py index cdc44c095..0522a2da5 100644 --- a/src/confluent_kafka/schema_registry/rules/cel/cel_executor.py +++ b/src/confluent_kafka/schema_registry/rules/cel/cel_executor.py @@ -24,6 +24,7 @@ from confluent_kafka.schema_registry import RuleKind, Schema from confluent_kafka.schema_registry.rule_registry import RuleRegistry +from confluent_kafka.schema_registry.rules.cel import protobuf_result_writer from confluent_kafka.schema_registry.rules.cel.cel_field_presence import InterpretedRunner from confluent_kafka.schema_registry.rules.cel.constraints import _msg_to_cel, _scalar_field_value_to_cel from confluent_kafka.schema_registry.rules.cel.extra_func import EXTRA_FUNCS @@ -31,6 +32,38 @@ log = logging.getLogger(__name__) + +def _to_plain_containers(value: Any) -> Any: + """Replaces celpy's container types with the plain ones, all the way down. + + ``celtypes.MapType`` is a dict subclass, so an Avro record it stands in for looks usable - + but it overrides ``get`` to *raise* KeyError when the key is absent and the default is + ``None``, where ``dict.get`` returns ``None``. fastavro fills an omitted field with + ``datum.get(name, field.get("default"))``, so a field whose declared default is ``null`` + passes ``None`` as that default and raises instead of taking the default. The effect was + that a message-level transform omitting a nullable field failed with a bare + ``KeyError: 'nullable'``, while a field with a non-null default worked - a split with no + reason behind it. Handing fastavro a plain dict gets the JVM's behaviour for free. + + Keys are normalised too: a celpy ``StringType`` is a str subclass and hashes alike, so this + is for the benefit of anything downstream that checks the type rather than the value. + + A tuple stays a tuple. fastavro selects a union branch from a ``(record_name, value)`` + pair, which is the only way to disambiguate two branches of the same shape, and + ``common/avro.py`` preserves that pair through the field-level walk for exactly that + reason. ``_value_to_cel`` has no tuple arm, so such a pair reaches a rule unconverted and + comes back out of an identity transform unchanged - and flattening it to a list left + fastavro with a value it refuses outright ("['B', {'x': 5}] (type list) do not match ..."). + """ + if isinstance(value, dict): + return {str(k): _to_plain_containers(v) for k, v in value.items()} + if isinstance(value, tuple): + return tuple(_to_plain_containers(v) for v in value) + if isinstance(value, list): + return [_to_plain_containers(v) for v in value] + return value + + # A date logical type annotates an Avro int, where the int stores the number # of days from the unix epoch, 1 January 1970 (ISO calendar). DAYS_SHIFT = datetime.date(1970, 1, 1).toordinal() @@ -69,7 +102,21 @@ def execute(self, ctx: RuleContext, msg: Any, args: Any) -> Any: return msg expr = expr[index + 1 :] - return self.execute_rule(ctx, expr, args) + return self._write_back(ctx, msg, self.execute_rule(ctx, expr, args)) + + def _write_back(self, ctx: RuleContext, msg: Any, result: Any) -> Any: + """Shapes a rule result back into the form the serializer for this format expects. + + An Avro record is a dict in this client, so the celpy map a rule returns is nearly + usable as-is - but only nearly, hence ``_to_plain_containers``. A protobuf message is + not a dict at all: the result has to be rebuilt into one, which is also where the + transform's replace semantics live. + """ + if ctx.rule.kind == RuleKind.CONDITION: + return result + if isinstance(msg, message.Message): + return protobuf_result_writer.convert(result, msg) + return _to_plain_containers(result) def execute_rule(self, ctx: RuleContext, expr: str, args: Any) -> Any: schema = ctx.target @@ -85,6 +132,12 @@ def execute_rule(self, ctx: RuleContext, expr: str, args: Any) -> Any: ast = self._env.compile(expr) prog = self._env.program(ast, functions=self._funcs) self._cache.set(expr, script_type, schema, prog) + # `now` is bound lazily, fresh per evaluation. Only inject when the + # expression references it — substring check matches protovalidate- + # python's pattern. Each rule evaluation sees a freshly-captured UTC + # instant. + if "now" in expr and "now" not in args: + args["now"] = celtypes.TimestampType(datetime.datetime.now(tz=datetime.timezone.utc)) result = prog.evaluate(args) if isinstance(result, celtypes.BoolType): return bool(result) diff --git a/src/confluent_kafka/schema_registry/rules/cel/cel_field_executor.py b/src/confluent_kafka/schema_registry/rules/cel/cel_field_executor.py index f07dcfaea..d97ebf477 100644 --- a/src/confluent_kafka/schema_registry/rules/cel/cel_field_executor.py +++ b/src/confluent_kafka/schema_registry/rules/cel/cel_field_executor.py @@ -33,8 +33,10 @@ def new_transform(self, ctx: RuleContext) -> FieldTransform: return self._field_transform def _field_transform(self, ctx: RuleContext, field_ctx: FieldContext, field_value: Any) -> Any: - if field_value is None: - return None + # No null guard here, matching the reference: whether an absent value reaches a rule is + # each format's walk to decide, not the executor's. The protobuf walk skips an unset + # field before calling this (a field with presence that is unset has no value to + # transform); the Avro walk passes the null branch through so a rule can guard on it. if not field_ctx.is_primitive(): return field_value args = { diff --git a/src/confluent_kafka/schema_registry/rules/cel/cel_field_presence.py b/src/confluent_kafka/schema_registry/rules/cel/cel_field_presence.py index 597f1b5fb..37297b9be 100644 --- a/src/confluent_kafka/schema_registry/rules/cel/cel_field_presence.py +++ b/src/confluent_kafka/schema_registry/rules/cel/cel_field_presence.py @@ -14,9 +14,10 @@ # limitations under the License. import threading -from typing import Any +from typing import Any, Optional import celpy +import lark _has_state = threading.local() @@ -33,6 +34,50 @@ def in_has() -> bool: return getattr(_has_state, "in_has", False) +# Method-call macros that the standard CEL evaluator handles directly via +# `member_dot_arg`. We must never intercept these — they're not user-registered +# functions and must keep their stdlib semantics. +_RESERVED_MACROS = frozenset(["map", "filter", "all", "exists", "exists_one", "reduce", "min"]) + + +def _extract_namespace_path(tree: Any) -> Optional[str]: + """ + If ``tree`` is a pure identifier chain (e.g., ``decimals`` or ``foo.bar``) + with no function calls / indexing / literals, return the dotted path + string. Otherwise return None. + + Used by the namespace-aware ``member_dot_arg`` override below to detect + rule fragments like ``decimals.ge(a, b)`` and look up + ``"decimals.ge"`` as a flat function-registry key — which is what + cel-java/go/cpp do for namespaced extensions (decimals.*, variants.*, + timestamp.*), but which celpy doesn't natively support. + """ + if not isinstance(tree, lark.Tree): + return None + if tree.data == "member_dot": + # member_dot: member "." IDENT (a value-position access, no parens) + # Two children: left (member tree), right (IDENT token). + left = _extract_namespace_path(tree.children[0]) + if left is None: + return None + right = tree.children[1] + if not isinstance(right, lark.Token) or right.type != "IDENT": + return None + return f"{left}.{right.value}" + if tree.data in ("member", "primary"): + # Pass-through: single child wrapping the actual node. + if len(tree.children) != 1: + return None + return _extract_namespace_path(tree.children[0]) + if tree.data == "ident": + # ident: IDENT + child = tree.children[0] + if isinstance(child, lark.Token) and child.type == "IDENT": + return child.value + return None + return None + + class InterpretedRunner(celpy.InterpretedRunner): def evaluate(self, context: Any) -> Any: class Evaluator(celpy.Evaluator): @@ -42,6 +87,62 @@ def macro_has_eval(self, exprlist: Any) -> celpy.celtypes.BoolType: _has_state.in_has = False return result + def member_dot_arg(self, tree: Any) -> Any: + """ + Adds namespace-aware function dispatch to celpy. + + Standard celpy treats ``x.foo(args)`` as a method call on + ``x``, looking up bare ``foo`` in the function registry. That + makes dotted function names like ``decimals.ge(a, b)`` (used + by the schema-registry CEL extensions, matching cel-java / + cel-go / cel-cpp) impossible. + + Override: if the receiver is a pure identifier path + (``decimals``, ``timestamp``, ``foo.bar``) and + ``{path}.{method}`` exists in the function registry, dispatch + that flat function directly with the explicit args. Otherwise + fall through to celpy's standard method dispatch (which keeps + ``"hi".startsWith("h")``, ``ts.getDate()``, ``list.all(...)`` + etc. working). + """ + if isinstance(tree, lark.Tree) and len(tree.children) >= 2: + member_tree = tree.children[0] + method_token = tree.children[1] + if ( + isinstance(method_token, lark.Token) + and method_token.type == "IDENT" + and method_token.value not in _RESERVED_MACROS + ): + path = _extract_namespace_path(member_tree) + if path is not None: + candidate = f"{path}.{method_token.value}" + funcs = getattr(self.activation, "functions", None) + if funcs is not None and candidate in funcs: + # The dotted-function name is registered — + # dispatch it directly. Precedence: a + # registered `{path}.{method}` always wins + # over standard method-call dispatch. This + # lets the namespaced extensions + # (`decimals.*`, `variants.*`) work even when + # the path's root identifier also names + # something else in scope. + func = funcs[candidate] + if len(tree.children) == 3: + args = list(self.visit(tree.children[2])) + else: + args = [] + return func(*args) + return super().member_dot_arg(tree) + e = Evaluator(ast=self.ast, activation=self.new_activation()) value = e.evaluate(context) return value + + +def _is_bound_variable(activation: Any, name: str) -> bool: + """Return True if `name` resolves to a bound variable in the activation.""" + try: + activation.resolve_variable(name) + return True + except Exception: + return False diff --git a/src/confluent_kafka/schema_registry/rules/cel/cel_validator.py b/src/confluent_kafka/schema_registry/rules/cel/cel_validator.py index c57268b02..2fed58f90 100644 --- a/src/confluent_kafka/schema_registry/rules/cel/cel_validator.py +++ b/src/confluent_kafka/schema_registry/rules/cel/cel_validator.py @@ -23,6 +23,7 @@ from confluent_kafka.schema_registry.rules.cel.cel_executor import _value_to_cel from confluent_kafka.schema_registry.rules.cel.cel_field_presence import InterpretedRunner from confluent_kafka.schema_registry.rules.cel.constraints import _field_value_to_cel, _msg_to_cel +from confluent_kafka.schema_registry.rules.cel.decimal_funcs import decimal_boundary_value from confluent_kafka.schema_registry.rules.cel.extra_func import EXTRA_FUNCS from confluent_kafka.schema_registry.serde import RuleError, ValidationRule, ValidationRuleExecutor @@ -119,6 +120,9 @@ def _to_cel(schema: Any, value: Any) -> Any: scalars, repeated fields and maps faithfully), and the format's schema object otherwise (unused — Avro/JSON values are converted structurally). """ + decimal_value = decimal_boundary_value(value) + if decimal_value is not None: + return decimal_value if isinstance(value, message.Message): return _msg_to_cel(value) if isinstance(schema, descriptor.FieldDescriptor): diff --git a/src/confluent_kafka/schema_registry/rules/cel/constraints.py b/src/confluent_kafka/schema_registry/rules/cel/constraints.py index a910de0bf..a47b23c6f 100644 --- a/src/confluent_kafka/schema_registry/rules/cel/constraints.py +++ b/src/confluent_kafka/schema_registry/rules/cel/constraints.py @@ -42,15 +42,41 @@ def make_key_path(field_name: str, key: celtypes.Value) -> str: return f"{field_name}[{string_format.format_value(key)}]" # type: ignore[str-bytes-safe] +# A CEL timestamp is a ``datetime`` and a CEL duration a ``timedelta``, both of which resolve +# to microseconds, so the low three digits of a nanos field cannot survive the conversion. The +# write-back does not need them to: an echoed value can be copied from the message it was read +# from, which is what the decimal and variant bindings already let it do (see +# ``_set_message``'s "echoed unchanged" path in protobuf_result_writer). These two carry the +# same source, so an identity rule - or a rule that rewrites some other field and merely passes +# this one along - leaves the nanos exactly as the producer wrote them. +# +# Measured before this: a rule rewriting only a sibling field turned nanos 1 into 0, while the +# decimal in the same message came back byte-identical. A *computed* timestamp still lands on +# the microsecond ceiling, which is inherent to the type and matches what ``timestamp(x, 9)`` +# and ``string(ts)`` already document. +def _with_source(value, msg: message.Message): + """Tags a converted value with the message it came from, for the write-back.""" + value.msg = msg + return value + + def make_duration(msg: message.Message) -> celtypes.DurationType: - return celtypes.DurationType( - seconds=msg.seconds, - nanos=msg.nanos, + return _with_source( + celtypes.DurationType( + seconds=msg.seconds, + nanos=msg.nanos, + ), + msg, ) def make_timestamp(msg: message.Message) -> celtypes.TimestampType: - return make_duration(msg) + celtypes.TimestampType(1970, 1, 1) # type: ignore[return-value] + # The intermediate duration carries this same message; it is discarded here, and only the + # timestamp's own tag is read back. + return _with_source( + make_duration(msg) + celtypes.TimestampType(1970, 1, 1), # type: ignore[operator] + msg, + ) def unwrap(msg: message.Message) -> celtypes.Value: diff --git a/src/confluent_kafka/schema_registry/rules/cel/decimal_funcs.py b/src/confluent_kafka/schema_registry/rules/cel/decimal_funcs.py new file mode 100644 index 000000000..6e2d0d98c --- /dev/null +++ b/src/confluent_kafka/schema_registry/rules/cel/decimal_funcs.py @@ -0,0 +1,912 @@ +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""CEL bindings for the {@code decimal} constructor and {@code decimals.*} operators. + +celpy has no overload-set concept — one function per name, internal arity + type +dispatch. The {@code decimal} constructor handles both shapes +({@code decimal(dyn)} and {@code decimal(bytes, scale)}) in a single Python +callable that branches on {@code len(args)}. + +Decimal division uses {@code decimal.Context(prec=38, rounding=ROUND_HALF_UP)} — +matches Flink SQL's MC_DIVIDE and the Java reference implementation. Add/sub/mul +use Python's default exact arithmetic (BigDecimal-like). +""" + +import decimal +import typing +from decimal import Decimal + +import celpy +from celpy import celtypes + +from confluent_kafka.schema_registry.rules.cel.timestamp_funcs import format_timestamp + +try: + from confluent_kafka.schema_registry.confluent.type import decimal_pb2 + from confluent_kafka.schema_registry.confluent.type.decimal_utils import from_proto_decimal as _from_proto_decimal + + _PROTO_DECIMAL_CLS: typing.Any = decimal_pb2.Decimal +except ImportError: + _PROTO_DECIMAL_CLS = None + _from_proto_decimal = None # type: ignore[assignment] + + +# 38-digit precision with HALF_UP rounding — matches Flink/PostgreSQL NUMERIC +# division. +# Emax/Emin are widened alongside: the default +/-999999 is far narrower than the exponent +# range the constructor accepts (BigDecimal's signed-int scale), so dividing a legitimately +# constructed value such as decimal("1e1000000") overflowed where Java returns a result. +_DIV_CONTEXT = decimal.Context( + prec=38, + rounding=decimal.ROUND_HALF_UP, + Emax=decimal.MAX_EMAX, + Emin=decimal.MIN_EMIN, +) + +# Exact/unbounded context for operations Java computes exactly (add/sub/mul/mod, +# setScale/quantize, scaleb) — matches java.math.BigDecimal's exact semantics rather +# than the thread-local default context (prec=28). Only div/sqrt cap at 38 (_DIV_CONTEXT). +_EXACT_CONTEXT = decimal.Context(prec=decimal.MAX_PREC, Emax=decimal.MAX_EMAX, Emin=decimal.MIN_EMIN) + + +_INT32_MIN = -(2**31) +_INT32_MAX = 2**31 - 1 + + +def _require_int_scale(scale: typing.Any, fn: str) -> int: + """A scale argument as an int, mirroring Java's ``requireIntScale``. + + Java declares every scale parameter as ``Long`` and narrows it with ``Math.toIntExact``, + so a double or a bool has no matching overload there and an out-of-int-range value is an + error rather than a wildly wrong Decimal. Go, C#, C++ and Rust carry the same check. + A *negative* scale is legitimate - ``BigDecimal.setScale(-2)`` rounds to hundreds - so + only the type and the width are constrained here. + """ + # UintType alongside the bools because it subclasses ``int`` too, so the ``int`` arm would + # accept it. ``uint`` is a distinct CEL type and no declared overload takes one - the scale + # is ``SimpleType.INT`` in the reference - so ``decimals.round(decimal("1"), 1u)`` has to + # fail. cel-java, cel-go and cel-es all report "no matching overload" at compile time; + # celpy is untyped, so the nearest equivalent is refusing here, as the Rust client does. + if isinstance(scale, (bool, celtypes.BoolType, celtypes.UintType)) or not isinstance( + scale, (int, celtypes.IntType) + ): + raise celpy.CELEvalError(f"{fn}: scale must be int, got {type(scale).__name__}") + s = int(scale) + if s < _INT32_MIN or s > _INT32_MAX: + raise celpy.CELEvalError(f"{fn}: scale out of int range: {s}") + return s + + +def _drop_negative_zero(d: Decimal) -> Decimal: + """A zero without a sign, because BigDecimal has no negative zero. + + ``BigDecimal("-0").toPlainString()`` is ``"0"`` and its signum is 0, while Python's Decimal + keeps the sign and renders ``"-0"``. ``abs`` preserves the scale, so ``-0.00`` becomes + ``0.00`` rather than ``0`` - matching ``BigDecimal("-0.00").toPlainString()``. + """ + return abs(d) if not d and d.is_signed() else d + + +def _from_bytes_scale(value: typing.Any, scale: typing.Any) -> Decimal: + """Construct a Decimal from raw two's-complement big-endian bytes + scale. + + The scale itself costs nothing here - ``scaleb`` only sets the exponent, so a coefficient + at an extreme scale stays compact and the width guard fires later, where the digits are + actually needed. The *coefficient* is the width risk on this path, and it is checked + before ``int.from_bytes`` builds it: one byte carries about 2.41 decimal digits. + """ + raw = _coerce_bytes(value) + s = _require_int_scale(scale, "decimal(bytes, scale)") + if len(raw) == 0: + return Decimal(0).scaleb(-s, context=_EXACT_CONTEXT) + _require_sane_width(int(len(raw) * 2.408) + 1, "decimal(bytes, scale)", "the coefficient", _SANE_COEFFICIENT) + return Decimal(int.from_bytes(raw, "big", signed=True)).scaleb(-s, context=_EXACT_CONTEXT) + + +def _coerce_bytes(v: typing.Any) -> bytes: + if isinstance(v, (bytes, bytearray)): + return bytes(v) + if isinstance(v, memoryview): + return v.tobytes() + if isinstance(v, celtypes.BytesType): + return bytes(v) + raise celpy.CELEvalError(f"decimal: expected bytes for the (bytes, scale) overload, got " f"{type(v).__name__}") + + +def _decimal_from_string(text: str, original: typing.Any) -> Decimal: + """Parse a Decimal the way Java's ``new BigDecimal(String)`` does. + + Python's ``Decimal(str)`` is far more permissive than Java's + ``BigDecimal(String)`` / ``BigDecimal.valueOf(double)``, silently accepting + inputs Java rejects with ``NumberFormatException``: + + * non-finite values — ``"NaN"``, ``"sNaN"``, ``"Infinity"``, ``"inf"``, + ``"-inf"`` (and the ``str()`` of a NaN/Inf float); + * underscore digit-grouping such as ``"1_000"`` (→ 1000); + * surrounding whitespace (Python strips it; internal whitespace is already + rejected by ``Decimal``). + + Reject those to match Java, while still accepting every legitimate finite + decimal (integers, ``"1e40"``, negatives, ``"-0"``, a leading ``+``, ...). + """ + if "_" in text or text != text.strip(): + raise celpy.CELEvalError(f"decimal: invalid number '{original}'") + try: + d = Decimal(text) + except decimal.InvalidOperation as ex: + raise celpy.CELEvalError(f"decimal: invalid number '{original}'") from ex + if not d.is_finite(): + raise celpy.CELEvalError(f"decimal: invalid number '{original}'") + # BigDecimal holds its scale in a signed int and rejects a literal whose exponent will not + # fit, so `1e-2147483648` is a NumberFormatException there while Python's Decimal accepts + # it happily. Measured against the JVM, the accepted band is symmetric: |exponent| <= + # INT32_MAX, with both +/-2147483648 refused. Left unchecked, rendering such a value as + # fixed-point would try to materialise billions of digits. + exponent = d.as_tuple().exponent + if not isinstance(exponent, int) or exponent < -_INT32_MAX or exponent > _INT32_MAX: + raise celpy.CELEvalError(f"decimal: invalid number '{original}'") + return _drop_negative_zero(d) + + +def _decimal(*args: typing.Any) -> Decimal: + """Runtime dispatch backing the {@code decimal(...)} constructor. + + Two arities: + * {@code decimal(dyn)} — convert any supported value to Decimal. + * {@code decimal(bytes, int)} — explicit bytes + scale construction. + """ + if len(args) == 2: + return _from_bytes_scale(args[0], args[1]) + if len(args) != 1: + raise celpy.CELEvalError(f"decimal: expected 1 or 2 args, got {len(args)}") + v = args[0] + if v is None: + raise celpy.CELEvalError("decimal: cannot convert null to Decimal") + if isinstance(v, Decimal): + return v + if _PROTO_DECIMAL_CLS is not None and isinstance(v, _PROTO_DECIMAL_CLS): + return _from_proto_decimal(v) + # Generic proto Message duck-typing — accept any message whose descriptor + # full_name is confluent.type.Decimal (covers DynamicMessage or alternate + # generated bindings). + if hasattr(v, "DESCRIPTOR") and getattr(v.DESCRIPTOR, "full_name", "") == "confluent.type.Decimal": + return _from_proto_decimal(v) + # celpy binds a proto-message field into CEL as a MessageType wrapper (a MapType that + # keeps the underlying message on ``.msg``), so `decimal(message.decField)` for a + # confluent.type.Decimal field arrives here rather than as a raw message. Unwrap it. + proto_msg = getattr(v, "msg", None) + if ( + proto_msg is not None + and getattr(getattr(proto_msg, "DESCRIPTOR", None), "full_name", "") == "confluent.type.Decimal" + ): + return _from_proto_decimal(proto_msg) + if isinstance(v, (bool, celtypes.BoolType)): + # bool is a subclass of int in Python, and celtypes.BoolType subclasses int rather + # than bool, so both have to be named here or a CEL bool becomes Decimal(1). Java has + # no decimal(bool) overload at all. + raise celpy.CELEvalError("decimal: cannot convert bool to Decimal") + if isinstance(v, int): + return Decimal(v) + if isinstance(v, float): + # Java uses BigDecimal.valueOf(double), which throws on NaN/Infinity. + # str() of a non-finite float ("nan"/"inf"/"-inf") builds a poisoned + # Decimal in Python, so validate through the same finite check. + return _decimal_from_string(str(v), v) + if isinstance(v, (str, celtypes.StringType)): + return _decimal_from_string(str(v), v) + if isinstance(v, (bytes, bytearray, memoryview, celtypes.BytesType)): + raise celpy.CELEvalError( + "decimal: raw bytes need a scale; use decimal(bytes, scale) or set " + "useLogicalTypeConverters=true on the Avro client so decimal fields " + "arrive as Decimal" + ) + raise celpy.CELEvalError(f"decimal: cannot convert {type(v).__name__} to Decimal") + + +# ---- comparison ---- + + +def decimal_boundary_value(v: typing.Any) -> typing.Optional[Decimal]: + """A ``confluent.type.Decimal`` message as a :class:`~decimal.Decimal`, else ``None``. + + Presents a decimal-shaped bound value as this client's in-CEL decimal. Without it a bare + decimal field and ``decimal(...)`` are two different things at runtime: the field is a celpy + ``MessageType`` wrapper, so ``this == decimal("12.34")`` compares a wrapper against a + ``Decimal`` and answers False, and ``string(this)`` / ``double(this)`` fail outright with + ``TypeError: float() argument must be a string or a real number, not 'MessageType'``. + + Converting at the boundary also makes ``==`` scale-insensitive for free, since + ``Decimal.__eq__`` is numeric: 12.34 and 12.340 are the same number in two encodings, and + comparing the protobuf fields one by one would call them unequal. + + Only reaches a value bound directly. A decimal reached by selection instead + (``this.amount``) is resolved inside celpy, past any boundary. + """ + if isinstance(v, Decimal): + return v + if _from_proto_decimal is None: + return None + if _PROTO_DECIMAL_CLS is not None and isinstance(v, _PROTO_DECIMAL_CLS): + return _from_proto_decimal(v) + if getattr(getattr(v, "DESCRIPTOR", None), "full_name", "") == "confluent.type.Decimal": + return _from_proto_decimal(v) + # celpy wraps a proto message as a MessageType keeping the message on ``.msg``. + proto_msg = getattr(v, "msg", None) + if ( + proto_msg is not None + and getattr(getattr(proto_msg, "DESCRIPTOR", None), "full_name", "") == "confluent.type.Decimal" + ): + return _from_proto_decimal(proto_msg) + return None + + +def _decimals_eq(a: typing.Any, b: typing.Any) -> celtypes.BoolType: + return celtypes.BoolType(_d(a).compare(_d(b)) == 0) + + +def _decimals_lt(a: typing.Any, b: typing.Any) -> celtypes.BoolType: + return celtypes.BoolType(_d(a).compare(_d(b)) < 0) + + +def _decimals_le(a: typing.Any, b: typing.Any) -> celtypes.BoolType: + return celtypes.BoolType(_d(a).compare(_d(b)) <= 0) + + +def _decimals_gt(a: typing.Any, b: typing.Any) -> celtypes.BoolType: + return celtypes.BoolType(_d(a).compare(_d(b)) > 0) + + +def _decimals_ge(a: typing.Any, b: typing.Any) -> celtypes.BoolType: + return celtypes.BoolType(_d(a).compare(_d(b)) >= 0) + + +# ---- arithmetic ---- + + +# The width ceiling for a computation. Deliberately *not* BigDecimal's - BigInteger tops out +# at Integer.MAX_VALUE bits, which is 646456993 decimal digits, and reproducing that bound is +# neither achievable across six libraries nor the point. This is a round number chosen so no +# single rule evaluation can exhaust memory: 10**7 digits is ~4 MB of libmpdec coefficient +# (packed 19 digits to a 64-bit word) and ~10 MB rendered. Values above it are refused as a +# rule error, which is the one thing a resource exhaustion cannot be turned into after the +# fact. Java is the only client in the family that fails cleanly on width; this stands in for +# that, as a bound rather than as a domain model. +_SANE_WIDTH = 10_000_000 + +# A far tighter bound on what can be *encoded*, which is a different resource. The wire form +# is the unscaled integer in base 256, and decimal <-> binary radix conversion is quadratic in +# every client: measured in the C++ client, its digit-string codec takes 0.04 s at 10**4 +# digits, 4.2 s at 10**5 and ~420 s at 10**6, and mpdecimal's own mpd_qexport_u32 is only +# about 10x better with the same quadratic shape. So a value can be cheap to hold, cheap to +# compute with, and still unserialisable. +# +# 4300 is not arbitrary: it is CPython's own int_max_str_digits, the limit it puts on +# str <-> int conversion for exactly this reason. This client already could not encode a wider +# coefficient - `int("9" * 5000)` raises ValueError "Exceeds the limit (4300 digits) for +# integer string conversion" - so the bound is pre-existing and the only thing added is that +# it now reads as a decimal error instead of a CPython internal one. Held as a literal rather +# than read from `sys` so the accepted set does not shift with a host's own setting. +_SANE_COEFFICIENT = 4300 + + +def _adjusted_of(d: Decimal) -> int: + """``d.adjusted()``, guarded for a non-finite value the way ``_exponent_of`` is.""" + if not d.is_finite(): + raise celpy.CELEvalError(f"decimal: not a finite number '{d}'") + return d.adjusted() + + +def _rescaled_digits(target_scale: int, d: Decimal) -> int: + """Digits in the coefficient ``d`` would have at ``target_scale``. + + Only *expanding* a scale costs anything - the coefficient grows by the difference. + Coarsening one is free at any distance, and the earlier ``abs(shift) + digits`` form + refused it wrongly. Measured, all instant and all one digit wide: + + * ``1.23`` at scale -1000000, -100000000, -2000000000 -> 0E+1000000 ... 0E+2000000000 + * ``1e-1000000`` and ``1e-100000000`` at scale 0 -> 0 + + against ``1.23`` at scale 100000000, which is 952 MB and a 100000001-digit coefficient. + Java agrees on both sides: ``BigDecimal("1.23").setScale(-100000000)`` is precision 1. + + Computed from ``(exponent, adjusted)`` rather than from the value, so the digits the + caller is about to refuse are never built - ``as_tuple()`` on a wide value is itself the + allocation being guarded (9.2 GB for a 2**31-digit coefficient). + """ + exponent = _exponent_of(d) + digits = _adjusted_of(d) - exponent + 1 + return max(1, digits + target_scale + exponent) + + +def _plain_form_length(d: Decimal) -> int: + """Characters in ``d``'s plain (non-scientific) rendering, to within a couple. + + Unlike a rescale, this *does* pay for the exponent in both directions: a positive + exponent writes that many trailing zeros and a negative one that many leading zeros, so + ``0E-2147483647`` renders as two billion characters even though its coefficient is one + digit. + """ + exponent = _exponent_of(d) + digits = _adjusted_of(d) - exponent + 1 + return digits + abs(exponent) + + +def _require_sane_width(needed: int, fn: str, what: str, limit: int = _SANE_WIDTH) -> None: + """Refuse a positional form too wide to build. + + Three unrelated-looking things reduce to this one quantity, because each has to + materialise a value in positional form: + + * **aligning two exponents** - ``add`` and ``sub`` expand the narrower operand into the + wider one's frame before computing a single digit; ``remainder`` is the same family but + is bounded by its integral quotient instead (see :func:`_decimals_mod`); + * **rescaling** - ``round``/``trunc``/``floor``/``ceil`` produce a coefficient at the + target scale; + * **rendering** - ``string()`` writes every digit out. + + ``mul``, ``div``, comparison, negation and ``abs`` are absent deliberately: ``mul`` adds + the exponents and multiplies the coefficients, ``div`` holds the coefficient to the + context precision and lets the exponent absorb the difference, and libmpdec's comparison + short-circuits on the adjusted exponent. Measured on operands 1e2147483647 and 3, peak + RSS: ``mul``, ``div``, ``<``, ``==``, ``compare``, ``min``, ``neg``, ``abs`` all 13 MB; + ``add`` 1738 MB, ``sub`` 1738 MB, ``remainder`` 1733 MB; and + ``add(1e2147483647, 1e-2147483647)`` 3125 MB. So the guard follows *alignment*, not + arithmetic - a single expression over two cheaply constructed operands is enough. + """ + if needed > limit: + raise celpy.CELEvalError(f"{fn}: {what} needs {needed} digits, past this client's {limit}-digit limit") + + +def _operand_width(target_scale: int, d: Decimal) -> int: + """Digits ``d`` needs once expanded to ``target_scale``. + + A **zero** contributes one digit whatever the distance, because expanding a zero appends no + digits - and that is what decides several of these cases, since alignment expands only the + operand whose scale is coarser. Measured on libmpdec, and the reference agrees on every + row: + + ========================== ========================== =============================== + expression libmpdec JDK + ========================== ========================== =============================== + ``0E+2e9 + 0E-2e9`` free, 1 digit precision 1 + ``0E+2e9 + 1`` free, 1 digit precision 1 + ``0E+2e9 mod 1E-2e9`` free, 1 digit precision 1 + ``0E-2e9 + 1`` **1601 MB**, 2e9+1 digits ArithmeticException + ========================== ========================== =============================== + + The last row is the one that must still be refused, and the difference is purely which + operand expands: aligning to scale 0 expands the zero (free), aligning to scale 2e9 + expands the *one* (2e9 digits). + """ + if not d: + return 1 + exponent = _exponent_of(d) + digits = _adjusted_of(d) - exponent + 1 + return digits + target_scale + exponent + + +def _require_additive_domain(x: Decimal, y: Decimal, fn: str) -> None: + """Addition and subtraction align both operands on the finer scale, so the frame is the + widest either of them needs there - computed per operand, because a zero costs nothing to + expand however far it moves (see :func:`_operand_width`).""" + target_scale = -min(_exponent_of(x), _exponent_of(y)) + needed = max(_operand_width(target_scale, x), _operand_width(target_scale, y)) + 1 + _require_sane_width(needed, fn, "aligning the operands") + + +def _decimals_add(a: typing.Any, b: typing.Any) -> Decimal: + x, y = _d(a), _d(b) + _require_additive_domain(x, y, "decimals.add") + return _EXACT_CONTEXT.add(x, y) + + +def _decimals_sub(a: typing.Any, b: typing.Any) -> Decimal: + x, y = _d(a), _d(b) + _require_additive_domain(x, y, "decimals.sub") + return _EXACT_CONTEXT.subtract(x, y) + + +def _decimals_mul(a: typing.Any, b: typing.Any) -> Decimal: + # No width guard: multiplication adds the exponents and multiplies the coefficients, so + # the result is as compact as its operands. It was guarded here once, on a prediction of + # BigDecimal's own domain errors; that prediction is what this design stopped doing, and + # the operations it guarded turned out to be the cheap ones. + return _EXACT_CONTEXT.multiply(_d(a), _d(b)) + + +def _decimals_div(a: typing.Any, b: typing.Any) -> Decimal: + try: + return _DIV_CONTEXT.divide(_d(a), _d(b)) + except decimal.DivisionByZero as ex: + raise celpy.CELEvalError("decimals.div: division by zero") from ex + except decimal.DecimalException as ex: + raise celpy.CELEvalError(f"decimals.div: {ex}") from ex + + +def _decimals_mod(a: typing.Any, b: typing.Any) -> Decimal: + """Remainder with the sign of the dividend (truncated division), matching + Java BigDecimal.remainder and SQL MOD. A zero divisor raises the canonical + message. + """ + da, db = _d(a), _d(b) + if db == 0: + raise celpy.CELEvalError("decimals.mod: division by zero") + # The remainder itself is small - its magnitude is bounded by both operands - but the + # *integral quotient* has to be produced to get there, and that is the width. Not the + # aligned frame add and sub are guarded on: libmpdec short-circuits when the operands' + # magnitudes are close or the dividend is the smaller, so the frame over-refuses. Measured + # (peak RSS), with the aligned frame in the last column for contrast: + # + # 1e2147483647 mod 3 2^31 quotient digits 1733 MB 2^31 frame + # 1.5 mod 1e-2147483647 2^31 1519 MB 2^31 + # 1e2147483647 mod 1e-2147483647 4.3e9 2568 MB 4.3e9 + # 1e-2147483647 mod 1e2147483647 0 13 MB 4.3e9 <- free + # 1e2147483647 mod 1e2147483000 647 13 MB 4.3e9 <- free + # 1e40 mod 3 40 13 MB + # + # so the last two are what the frame would have cost us, and both are values the JVM + # accepts: `1e-2147483647 mod 1e2147483647` is the dividend itself at precision 1. + # A zero dividend has a quotient of zero whatever the scales, and the adjusted exponent + # says nothing useful about it - a zero keeps whatever scale it was built with, so + # `0E+2e9 mod 1E-2e9` estimated 4e9 digits for a result that is just zero. Measured free on + # libmpdec, and the JDK returns 0 at precision 1. + quotient_digits = 1 if not da else max(0, _adjusted_of(da) - _adjusted_of(db)) + 1 + _require_sane_width(quotient_digits, "decimals.mod", "the integral quotient") + return _EXACT_CONTEXT.remainder(da, db) + + +# ---- selection ---- + + +def _decimals_greatest(a: typing.Any, b: typing.Any) -> Decimal: + return max(_d(a), _d(b)) + + +def _decimals_least(a: typing.Any, b: typing.Any) -> Decimal: + return min(_d(a), _d(b)) + + +# ---- square root ---- + + +def _apply_preferred_scale(value: Decimal, preferred_scale: int, fn: str) -> Decimal: + """``value`` rewritten to the reference's preferred scale. + + Strip trailing zeros down to - never below - ``preferred_scale``, then pad back up to it + when the natural scale is smaller. Only ever called on an exact result: padding an + inexact one would claim digits it does not have. + + A zero takes the preferred scale outright, in both directions, because ``BigDecimal`` + returns ``zeroValueOf(preferredScale)`` for it - a strip loop guarded on a non-zero + coefficient cannot lower a zero's scale, so the zero case has to be separate. Measured: + ``sqrt(0.000)`` is scale 1 in the reference, and ``0 / 3.00`` is scale -2. + """ + if not value: + return _quantize(value, preferred_scale, decimal.ROUND_HALF_UP, fn) + minimal = value.normalize(context=_EXACT_CONTEXT) + minimal_scale = -_exponent_of(minimal) + # The preferred scale does not override the context precision. The reference pads toward + # it only while the result still fits in ``mc.precision`` significant digits, and stops + # short otherwise: ``1.<40 zeros> / 1`` is scale 37 there and not the preferred 40, and + # ``1.<100 zeros> / 8`` is scale 38 because 0.125 already spends 3 of the 38 on digits + # that are not padding. ``_quantize`` runs in ``_EXACT_CONTEXT``, so nothing else caps it - + # the raw preferred scale would have padded ``sqrt(1.<100 zeros>)`` to 51 digits. + headroom = _DIV_CONTEXT.prec - len(minimal.as_tuple().digits) + target = max(minimal_scale, min(preferred_scale, minimal_scale + headroom)) + return _quantize(minimal, target, decimal.ROUND_HALF_UP, fn) + + +def _decimals_sqrt(a: typing.Any) -> Decimal: + """Square root with 38-digit HALF_UP precision (same context as div). + + A negative input raises the canonical ``decimals.sqrt: square root of + negative number`` message (no complex result); zero passes through to 0. + + An exact root is then rewritten to the reference's preferred scale. libmpdec and + ``BigDecimal`` disagree here, and only here: the decimal arithmetic spec's ideal + exponent for square root is ``floor(exponent / 2)``, while ``BigDecimal.sqrt`` uses + ``scale / 2`` truncated toward zero. Those are the same for an even scale and one apart + for an odd one, so ``sqrt(9.0)`` is ``3.0`` on libmpdec and ``3`` on the reference. Note + that division needs no such step: the spec's ideal exponent for divide is + ``exponent(dividend) - exponent(divisor)``, which *is* the reference's preferred scale. + """ + d = _d(a) + if d < 0: + raise celpy.CELEvalError("decimals.sqrt: square root of negative number") + root = _DIV_CONTEXT.sqrt(d) + # Inexact roots keep all 38 digits: a trailing zero there is significant. + if _EXACT_CONTEXT.multiply(root, root) != d: + return root + scale = -_exponent_of(d) + # `scale / 2` in Java truncates toward zero, so a negative scale halves toward zero too; + # Python's `//` floors, which would answer -2 where the reference answers -1. + preferred = scale // 2 if scale >= 0 else -(-scale // 2) + return _apply_preferred_scale(root, preferred, "decimals.sqrt") + + +# ---- unary ---- + + +def _decimals_neg(a: typing.Any) -> Decimal: + return _d(a).copy_negate() + + +def _decimals_abs(a: typing.Any) -> Decimal: + return _d(a).copy_abs() + + +def _decimals_sign(a: typing.Any) -> celtypes.IntType: + d = _d(a) + if d == 0: + return celtypes.IntType(0) + return celtypes.IntType(1 if d > 0 else -1) + + +# ---- rounding family ---- + + +def _exponent_of(d: Decimal) -> int: + """The decimal's exponent, as an int. + + ``Decimal.as_tuple().exponent`` is ``int | Literal['n', 'N', 'F']`` - the strings stand + for NaN, sNaN and Infinity - so it cannot be compared or negated as it comes. ``_d`` + rejects a non-finite value before this is reached; the check keeps that guarantee local + rather than assumed. + """ + exponent = d.as_tuple().exponent + if not isinstance(exponent, int): + raise celpy.CELEvalError(f"decimal: not a finite number '{d}'") + return exponent + + +def _quantize(d: Decimal, scale: int, rounding: str, fn: str) -> Decimal: + """``d`` at ``scale``, or a rule error. The single Guard B site. + + Two things have to hold, and they pull in opposite directions. + + The quantizer is built in ``_EXACT_CONTEXT``, not the ambient one. The default context's + Emin of -999999 made ``Decimal(1).scaleb(1000000)`` raise ``decimal.Overflow``, so a + negative scale past a million was refused for values the JVM rounds happily: + ``BigDecimal("1e1000000").setScale(-1000000)`` is a no-op, and ``setScale(-1000000)`` on + 1.23 gives 0E+1000000. ``Overflow`` is not an ``InvalidOperation`` either, so it escaped + the handler below as a raw Python exception rather than a rule error. + + Building it in ``_EXACT_CONTEXT`` alone goes too far the other way: libmpdec then honours + any int32 scale, and quantizing to one materialises the whole coefficient. A scale of + 2**31-1 costs 918 MB inside ``quantize`` (2**31 digits, packed 19 to a 64-bit word) and + 9.2 GB the moment anything calls ``as_tuple()`` on the result - which ``_exponent_of`` + and the protobuf writer's ``_set_decimal`` both do. So the shift is bounded, by + :data:`_SANE_WIDTH`. + + Every rounding call site goes through here, the one-argument forms included. Three of the + five did not, and each was a multi-GB allocation reachable from a rule with no scale + argument at all: ``round(x)``, ``floor(x)`` and ``ceil(x)`` on a value with a large + negative exponent quantize to scale 0. That is the failure mode of a guard hung off one + helper rather than off the operation. + """ + # Zero is one digit at any scale. Rescaling it never expands anything - measured, both + # directions are free, and its result stays compact - and BigDecimal agrees: + # `new BigDecimal(BigInteger.ZERO, 2147483647)` is precision 1. Without this, the width + # formula reads the exponent and refuses `round(decimal(b"", 2147483647))`, a false + # rejection of a value the reference handles. + if d: + _require_sane_width(_rescaled_digits(scale, d), fn, f"a scale of {scale}") + try: + return d.quantize( + Decimal(1).scaleb(-scale, context=_EXACT_CONTEXT), + rounding=rounding, + context=_EXACT_CONTEXT, + ) + except (decimal.InvalidOperation, decimal.Overflow, decimal.Underflow) as e: + raise celpy.CELEvalError(f"{fn}: cannot represent a scale of {scale}") from e + + +def _decimals_round(*args: typing.Any) -> Decimal: + """Round to the given scale (HALF_UP). One-arg form rounds to integer.""" + if len(args) == 1: + return _quantize(_d(args[0]), 0, decimal.ROUND_HALF_UP, "decimals.round") + if len(args) == 2: + scale = _require_int_scale(args[1], "decimals.round") + return _quantize(_d(args[0]), scale, decimal.ROUND_HALF_UP, "decimals.round") + raise celpy.CELEvalError(f"decimals.round: expected 1 or 2 args, got {len(args)}") + + +def _decimals_trunc(*args: typing.Any) -> Decimal: + """Truncate to the given scale (toward zero). One-arg form truncates to integer. + + Matches Flink's TRUNCATE early-return: if the target scale is at-or-finer + than the input's current scale, return the input unchanged. Without this + guard, ``quantize`` would zero-pad and the string representation would + diverge from Flink (numerically identical, but ``string(trunc(d, n>=cur))`` + output would differ). + """ + if len(args) == 1: + d = _d(args[0]) + # current scale = -exponent. Early-return if 0 >= current_scale. + if _exponent_of(d) >= 0: + return d + return _quantize(d, 0, decimal.ROUND_DOWN, "decimals.trunc") + if len(args) == 2: + d = _d(args[0]) + scale = _require_int_scale(args[1], "decimals.trunc") + if scale >= -_exponent_of(d): + return d + return _quantize(d, scale, decimal.ROUND_DOWN, "decimals.trunc") + raise celpy.CELEvalError(f"decimals.trunc: expected 1 or 2 args, got {len(args)}") + + +def _decimals_floor(a: typing.Any) -> Decimal: + return _quantize(_d(a), 0, decimal.ROUND_FLOOR, "decimals.floor") + + +def _decimals_ceil(a: typing.Any) -> Decimal: + return _quantize(_d(a), 0, decimal.ROUND_CEILING, "decimals.ceil") + + +def _d(v: typing.Any) -> Decimal: + """Coerce a rule-argument value to Decimal for operator dispatch. + + Non-finite values are rejected here as well as in the string constructor: Java's and + Rust's decimals cannot represent NaN or an infinity at all, so no operator should see + one. An already-Decimal argument - from arithmetic, or from a decimal field - is the + path that skipped the constructor's check. + """ + if isinstance(v, Decimal): + if not v.is_finite(): + raise celpy.CELEvalError(f"decimal: not a finite number '{v}'") + return v + return _decimal(v) + + +# ---- string(Decimal) — extend celpy stdlib's string(...) ---- + +# Capture celpy's stdlib string callable at import time so we can delegate to +# it for non-Decimal inputs. celpy uses StringType(value) as the conversion; +# treating it as the underlying coercion gives us the standard semantics for +# int/uint/double/bytes/timestamp/duration/string args. +_STDLIB_STRING = celtypes.StringType + + +def _cel_type_name(v: typing.Any) -> str: + """A CEL type name, for the "no matching overload" messages below.""" + if v is None: + return "null" + # BoolType subclasses int and UintType is distinct from int, so both come before the int arm. + if isinstance(v, (bool, celtypes.BoolType)): + return "bool" + if isinstance(v, celtypes.UintType): + return "uint" + if isinstance(v, (bytes, bytearray, celtypes.BytesType)): + return "bytes" + if isinstance(v, (str, celtypes.StringType)): + return "string" + if isinstance(v, (int, celtypes.IntType)): + return "int" + if isinstance(v, (float, celtypes.DoubleType)): + return "double" + if isinstance(v, (list, tuple, celtypes.ListType)): + return "list" + if isinstance(v, (dict, celtypes.MapType)): + return "map" + return type(v).__name__ + + +def _string(v: typing.Any) -> celtypes.StringType: + """Extension of CEL stdlib {@code string(...)} with Decimal and Timestamp arms. + + Returns ``Decimal.toPlainString()``-equivalent form (Python's + ``format(d, 'f')``) for Decimal inputs; renders a Timestamp through + :func:`~confluent_kafka.schema_registry.rules.cel.timestamp_funcs.format_timestamp`, + because celpy's own ``TimestampType.__str__`` silently drops the + sub-second component. + + Registering a ``string`` entry *replaces* celpy's, rather than extending it the way the + typed runtimes do, so every remaining type is this function's responsibility - and + ``StringType`` is just ``str``, which renders them as Python rather than as CEL. Measured: + a bool came back ``"True"`` where the reference gives ``"true"``, and null, a list and a + map - none of which CEL declares a ``string`` overload for - came back as ``"None"`` and + Python container reprs instead of the reference's "found no matching overload". + """ + if isinstance(v, celtypes.TimestampType): + return celtypes.StringType(format_timestamp(v)) + # A decimal reached by selection (`this.amount`) is a celpy MessageType wrapper, not a + # Decimal. Java's string() resolves it through asDecimalOrNull, which accepts the + # confluent.type.Decimal message form as well as its own decimal. + d = decimal_boundary_value(v) + if d is not None: + # Guard C. `format(d, "f")` writes every digit of the positional form, and that form + # can be enormous for a value that was cheap to compute: `div` holds its coefficient + # to 38 digits while its exponent runs free, so + # `decimals.div(decimal("1e-2147483647"), decimal("1e2147483647"))` costs nothing and + # renders as four billion characters. Measured: rendering a 10**8-digit value takes + # 204 MB. No zero shortcut here, unlike the rescale guard - a zero at an extreme + # scale renders as that many zeros. + _require_sane_width(_plain_form_length(d), "string", "the plain form") + return celtypes.StringType(format(_drop_negative_zero(d), "f")) + if isinstance(v, (bool, celtypes.BoolType)): + return celtypes.StringType("true" if v else "false") + if isinstance(v, (bytes, bytearray, celtypes.BytesType)): + return celtypes.StringType(bytes(v).decode("utf-8")) + if isinstance(v, (str, celtypes.StringType)): + return _STDLIB_STRING(v) + if isinstance( + v, + ( + int, + celtypes.IntType, + celtypes.UintType, + float, + celtypes.DoubleType, + celtypes.DurationType, + ), + ): + # StringType's own fallback for a non-string source is `str(source)`, but its signature + # only admits string and bytes forms; go through str() so the call type-checks. + return _STDLIB_STRING(str(v)) + raise celpy.CELEvalError(f"found no matching overload for 'string' applied to ({_cel_type_name(v)})") + + +# ---- double(Decimal) — extend celpy stdlib's double(...) ---- + +# Capture celpy's stdlib double callable so we can delegate non-Decimal inputs. +_STDLIB_DOUBLE = celtypes.DoubleType + + +def _double(v: typing.Any) -> celtypes.DoubleType: + """Extension of CEL stdlib {@code double(...)} with a Decimal arm. + + Narrowing conversion (``float(Decimal)``) for Decimal inputs — may lose + precision, and out-of-range magnitudes become ``inf``; delegates to celpy's + stdlib double coercion for everything else. + + Bool and bytes are refused first: ``DoubleType`` is ``float``, and Python's ``float`` + accepts both (``float(True)`` is 1.0, ``float(b"12")`` is 12.0) where CEL declares no such + overload and the reference reports one. Everything else celpy already rejects - a list, a + map and an unparseable string all raise out of ``float``. + """ + d = decimal_boundary_value(v) + if d is not None: + return celtypes.DoubleType(float(d)) + if isinstance(v, (bool, celtypes.BoolType, bytes, bytearray, celtypes.BytesType)): + raise celpy.CELEvalError(f"found no matching overload for 'double' applied to ({_cel_type_name(v)})") + return _STDLIB_DOUBLE(v) + + +def _as_decimal_or_none(o: typing.Any) -> typing.Optional[Decimal]: + """``o`` as a Decimal if it *is* one, else ``None``. + + Deliberately narrow — unlike :func:`_d`, it does not coerce ints, floats or strings. This runs + on every ``==`` in every rule, so turning ``1 == "1"`` into a decimal comparison would be + wrong, and it must stay cheap for the common case. + """ + return decimal_boundary_value(o) + + +def _has_decimal(o: typing.Any) -> bool: + """Whether ``o`` is a Decimal or holds one at any depth. + + Only consulted once both operands are containers, so it never runs on the scalar path. + """ + if _as_decimal_or_none(o) is not None: + return True + if isinstance(o, (list, tuple)): + return any(_has_decimal(e) for e in o) + if isinstance(o, dict): + return any(_has_decimal(v) for v in o.values()) + return False + + +def _cel_equals(a: typing.Any, b: typing.Any) -> bool: + """CEL ``==`` with decimals made numeric, as a plain bool. + + A decimal operand may be a :class:`~decimal.Decimal` or a ``confluent.type.Decimal`` message, + and comparing the latter structurally - field by field over unscaled bytes and scale - calls + 12.34 and 12.340 unequal even though they are the same number. + + Containers are handled too, but only when a decimal is actually inside one of them: the base + implementation recurses with its own equality, so a Decimal nested in a list or map was + compared structurally and ``[a] == [b]`` disagreed with ``a == b`` on the same values. Gating + on :func:`_has_decimal` leaves every decimal-free comparison on the base path untouched, and + each element pair recurses back through here so non-decimal elements keep base semantics. + """ + da = _as_decimal_or_none(a) + db = _as_decimal_or_none(b) + if da is not None and db is not None: + return da == db + if da is not None or db is not None: + # A decimal is never equal to a non-decimal. + return False + if isinstance(a, (list, tuple)) and isinstance(b, (list, tuple)) and (_has_decimal(a) or _has_decimal(b)): + return len(a) == len(b) and all(_cel_equals(x, y) for x, y in zip(a, b)) + if isinstance(a, dict) and isinstance(b, dict) and (_has_decimal(a) or _has_decimal(b)): + return len(a) == len(b) and all(k in b and _cel_equals(v, b[k]) for k, v in a.items()) + return bool(celpy.evaluation.bool_eq(a, b)) + + +def _decimal_aware_eq(a: typing.Any, b: typing.Any) -> typing.Any: + if ( + _as_decimal_or_none(a) is None + and _as_decimal_or_none(b) is None + and not _has_decimal(a) + and not _has_decimal(b) + ): + # No decimal anywhere: hand it straight back to the base implementation, errors and all. + return celpy.evaluation.bool_eq(a, b) + return celtypes.BoolType(_cel_equals(a, b)) + + +def _decimal_aware_ne(a: typing.Any, b: typing.Any) -> typing.Any: + if ( + _as_decimal_or_none(a) is None + and _as_decimal_or_none(b) is None + and not _has_decimal(a) + and not _has_decimal(b) + ): + return celpy.evaluation.bool_ne(a, b) + return celtypes.BoolType(not _cel_equals(a, b)) + + +def _decimal_aware_in(item: typing.Any, container: typing.Any) -> typing.Any: + """``in`` has to follow ``==`` or the two contradict each other.""" + if not _has_decimal(item) and not _has_decimal(container): + return celpy.evaluation.operator_in(item, container) + if isinstance(container, dict): + return celtypes.BoolType(any(_cel_equals(item, k) for k in container)) + if isinstance(container, (list, tuple)): + return celtypes.BoolType(any(_cel_equals(item, e) for e in container)) + return celpy.evaluation.operator_in(item, container) + + +# The CEL operators, overridden so a decimal compares numerically. celpy resolves functions +# through a ChainMap that consults these before its own base_functions. +DECIMAL_OPERATOR_FUNCS: typing.Dict[str, typing.Any] = { + "_==_": _decimal_aware_eq, + "_!=_": _decimal_aware_ne, + "_in_": _decimal_aware_in, +} + + +# Typed as Any rather than celpy.CELFunction: these functions return this client's own +# Decimal and Variant values, which are not in celpy's declared return union - the CEL +# surface is extended with opaque types celpy does not know. celpy dispatches them fine +# at runtime; only its annotation is narrower than what an extension can return. +DECIMAL_FUNCS: typing.Dict[str, typing.Any] = { + "decimal": _decimal, + "decimals.eq": _decimals_eq, + "decimals.lt": _decimals_lt, + "decimals.le": _decimals_le, + "decimals.gt": _decimals_gt, + "decimals.ge": _decimals_ge, + "decimals.add": _decimals_add, + "decimals.sub": _decimals_sub, + "decimals.mul": _decimals_mul, + "decimals.div": _decimals_div, + "decimals.mod": _decimals_mod, + "decimals.greatest": _decimals_greatest, + "decimals.least": _decimals_least, + "decimals.sqrt": _decimals_sqrt, + "decimals.neg": _decimals_neg, + "decimals.abs": _decimals_abs, + "decimals.sign": _decimals_sign, + "decimals.round": _decimals_round, + "decimals.trunc": _decimals_trunc, + "decimals.floor": _decimals_floor, + "decimals.ceil": _decimals_ceil, + # string(Decimal) overrides celpy stdlib — the wrapper falls through to + # stdlib for non-Decimal inputs. + "string": _string, + # double(Decimal) overrides celpy stdlib — same fall-through pattern. + "double": _double, +} diff --git a/src/confluent_kafka/schema_registry/rules/cel/extra_func.py b/src/confluent_kafka/schema_registry/rules/cel/extra_func.py index fafd75db0..f88ab5980 100644 --- a/src/confluent_kafka/schema_registry/rules/cel/extra_func.py +++ b/src/confluent_kafka/schema_registry/rules/cel/extra_func.py @@ -177,8 +177,17 @@ def is_uuid(string: celtypes.Value) -> celpy.Result: def make_extra_funcs(locale: str) -> typing.Dict[str, celpy.CELFunction]: + # Local import keeps the cel-python package import light when these + # extended-type modules aren't needed. + from confluent_kafka.schema_registry.rules.cel.decimal_funcs import ( + DECIMAL_FUNCS, + DECIMAL_OPERATOR_FUNCS, + ) + from confluent_kafka.schema_registry.rules.cel.timestamp_funcs import TIMESTAMP_FUNCS + from confluent_kafka.schema_registry.rules.cel.variant_funcs import VARIANT_FUNCS + string_fmt = string_format.StringFormat(locale) - return { + funcs: typing.Dict[str, celpy.CELFunction] = { # Missing standard functions "format": string_fmt.format, # protovalidate specific functions @@ -190,6 +199,12 @@ def make_extra_funcs(locale: str) -> typing.Dict[str, celpy.CELFunction]: "isHostname": is_hostname, "isUuid": is_uuid, } + # Extended types — decimal / timestamp / variant. + funcs.update(DECIMAL_FUNCS) + funcs.update(DECIMAL_OPERATOR_FUNCS) + funcs.update(TIMESTAMP_FUNCS) + funcs.update(VARIANT_FUNCS) + return funcs EXTRA_FUNCS = make_extra_funcs("en_US") diff --git a/src/confluent_kafka/schema_registry/rules/cel/protobuf_result_writer.py b/src/confluent_kafka/schema_registry/rules/cel/protobuf_result_writer.py new file mode 100644 index 000000000..2b886948f --- /dev/null +++ b/src/confluent_kafka/schema_registry/rules/cel/protobuf_result_writer.py @@ -0,0 +1,537 @@ +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Writes the result of a message-level ``CEL`` transform back into a protobuf message. + +A ``CEL`` rule that returns a map is returning **the whole new message**: the transform has +replace semantics, not merge. The result map is therefore rebuilt into a fresh message, which +gives three behaviours that a rule author needs to know about and that every client has to +match: + +* a field the rule does not name is **dropped** - a rule naming only the field it changes + discards the rest; +* a ``null`` in the map **clears** its field; +* echoing a field that was absent **materialises** it, because reading it produced a value. + Preserve absence with ``has(x) ? x : null``. + +Without this the executor handed back a raw ``celpy`` map, which the protobuf serializer +cannot write - a decimal, a timestamp and a variant are messages in protobuf and celpy has no +rendering for any of them. + +**Mechanism note.** The JVM client rebuilds by rendering the result to JSON and parsing it +back (``ProtobufResultWriter`` plus ``ProtobufSchema.fromJson``). This builds the message +directly against the descriptor instead: Python has no ``fromJson`` on the protobuf schema, +and going through JSON would mean base64-encoding every bytes field and formatting every +timestamp as RFC 3339 only for ``ParseDict`` to parse them straight back. Direct construction +is fewer conversions and fewer places to lose fidelity. The behaviours the JVM client gets +for free from the JSON mapping - null clearing a field, and a key matching either the +declared name or the JSON name - are reproduced explicitly below. +""" + +import datetime +import decimal +import math +from typing import Any, Mapping, Optional + +import celpy.celtypes as celtypes +from google.protobuf import descriptor, message + +from confluent_kafka.schema_registry.common.protobuf import ( + MAX_ENCODABLE_COEFFICIENT_DIGITS, + _is_repeated, +) +from confluent_kafka.schema_registry.confluent.type.decimal_utils import ( + unscaled_to_bytes, +) +from confluent_kafka.schema_registry.confluent.type.variant_utils import Variant +from confluent_kafka.schema_registry.rules.cel.constraints import _WRAPPER_TYPES + +__all__ = ["convert"] + +_DECIMAL_TYPE_NAME = "confluent.type.Decimal" +_VARIANT_TYPE_NAME = "confluent.type.Variant" +_DURATION_TYPE_NAME = "google.protobuf.Duration" +_TIMESTAMP_TYPE_NAME = "google.protobuf.Timestamp" + +_EPOCH = datetime.datetime(1970, 1, 1, tzinfo=datetime.timezone.utc) + + +def convert(result: Any, msg: Any) -> Any: + """Rebuild ``msg``'s type from a CEL result map, or return ``result`` unchanged. + + ``msg`` is the message the rule ran against; its concrete class is reused so the caller + gets back the type it passed in rather than a dynamic message. + """ + if not isinstance(result, Mapping) or not isinstance(msg, message.Message): + return result + out = type(msg)() + _fill(out, result) + return out + + +def _fill(out: message.Message, values: Mapping) -> None: + """Applies a result map to ``out``, one entry per declared field. + + Two entries can name the same slot, and applying both would leave the outcome to the + order the rule happened to write them in. ``JsonFormat`` refuses both shapes, and the two + have *opposite* null handling, which is the part worth stating: + + * **The same field twice.** ``_find_field`` accepts a field's declared name and its JSON + name, so ``total_amount`` and ``totalAmount`` are one field. ``mergeField`` tests + ``builder.hasField`` before its null early-return, so a null after a value is refused + ("Field p.M.total_amount has already been set.") while a null after a null is not. + * **Two members of one oneof.** Setting a member clears its siblings, so applying both + kept whichever came last - `{a: 1, b: 2}` kept b and `{b: 2, a: 1}` kept a. + ``mergeOneofField`` refuses this ("Cannot set field p.M.b because another field p.M.a + belonging to the same oneof has already been set"), but only after returning early for + a null, so a null does *not* count - which agrees with this writer's own rule that a + null clears rather than sets. + + Measured against protobuf-java 4.35.1. A proto3 ``optional`` field sits in a synthetic + oneof of exactly one member, so it can never collide with a sibling. + """ + desc = out.DESCRIPTOR + # field number -> the result key that set it; oneof name -> the field that filled it. + set_by: dict = {} + oneof_by: dict = {} + for key, value in values.items(): + name = str(key) + fd = _find_field(desc, name) + if fd is None: + # A key the schema does not declare has nowhere to go. Dropping it matches the + # JVM client, whose JSON parse ignores unknown fields. + continue + # Before the null branch, because that is where the JVM's hasField test sits. + first = set_by.get(fd.number) + if first is not None: + a, b = sorted((first, name)) + raise ValueError(f"result names field '{fd.full_name}' twice, as '{a}' and '{b}'") + if _is_null(value): + # An explicit null clears the field, which is how a rule preserves an absent + # value across a transform that echoes it. + out.ClearField(fd.name) + continue + set_by[fd.number] = name + oneof = fd.containing_oneof + if oneof is not None: + sibling = oneof_by.get(oneof.full_name) + if sibling is not None and sibling != fd.name: + a, b = sorted((sibling, fd.name)) + raise ValueError(f"result sets more than one member of oneof '{oneof.full_name}': " f"'{a}' and '{b}'") + oneof_by[oneof.full_name] = fd.name + _set_field(out, fd, value) + + +def _find_field(desc: descriptor.Descriptor, name: str) -> Optional[descriptor.FieldDescriptor]: + """Resolves a result key to a field by declared name, then by JSON name. + + A rule may legitimately return either, so matching only the declared name would silently + skip a field like ``total_amount``. + """ + fd = desc.fields_by_name.get(name) + if fd is not None: + return fd + return desc.fields_by_camelcase_name.get(name) + + +def _is_null(value: Any) -> bool: + return value is None or isinstance(value, celtypes.NullType) + + +def _set_field(out: message.Message, fd: descriptor.FieldDescriptor, value: Any) -> None: + # protobuf >=7 dropped the instance .label attribute; the client's own helper covers + # both runtimes. + if _is_repeated(fd): + if fd.message_type is not None and fd.message_type.GetOptions().map_entry: + _set_map(out, fd, value) + return + _set_repeated(out, fd, value) + return + if fd.type == descriptor.FieldDescriptor.TYPE_MESSAGE: + _set_message(getattr(out, fd.name), fd, value) + return + setattr(out, fd.name, _scalar(fd, value)) + + +def _set_map(out: message.Message, fd: descriptor.FieldDescriptor, value: Any) -> None: + if not isinstance(value, Mapping): + # Silence here dropped the field from the rebuilt message, turning a rule-authoring + # type error into lost data. The JVM's message-level path rejects the same mismatch, + # because it writes through a protobuf JSON parse. + raise ValueError(f"cannot write {type(value).__name__} to map field '{fd.name}'") + target = getattr(out, fd.name) + value_fd = fd.message_type.fields_by_name["value"] + for k, v in value.items(): + if _is_null(v): + # Dropping the entry reported success while deleting it. A protobuf map value + # cannot be null and the JVM's write-back parse says exactly that: "Map value + # cannot be null." - the counterpart of the repeated-element check below. + raise ValueError(f"cannot write a null value to map field '{fd.name}'") + if value_fd.type == descriptor.FieldDescriptor.TYPE_MESSAGE: + _set_message(target[k], value_fd, v) + else: + target[k] = _scalar(value_fd, v) + + +def _set_repeated(out: message.Message, fd: descriptor.FieldDescriptor, value: Any) -> None: + # A Mapping is iterable, so without naming it here a map result silently wrote the map's + # *keys* as the list - corruption rather than loss. A scalar or a string wrote an empty + # list. Both are rule-authoring type errors that the JVM's protobuf JSON parse rejects. + if isinstance(value, (str, bytes, Mapping)) or not hasattr(value, "__iter__"): + raise ValueError(f"cannot write {type(value).__name__} to repeated field '{fd.name}'") + target = getattr(out, fd.name) + del target[:] + for item in value: + if _is_null(item): + # Dropping it changed the list's length and hid the mistake. protobuf JSON says + # "Repeated field elements cannot be null in field: ..." and refuses the document. + raise ValueError(f"cannot write null to repeated field '{fd.name}'") + if fd.type == descriptor.FieldDescriptor.TYPE_MESSAGE: + _set_message(target.add(), fd, item) + else: + target.append(_scalar(fd, item)) + + +def _set_message(target: message.Message, fd: descriptor.FieldDescriptor, value: Any) -> None: + """Writes one message-valued field, inverting how the CEL binding read it. + + The three value types do not arrive as maps of their own fields once a rule has computed + one: a decimal comes back as a Python ``Decimal``, a timestamp as a ``datetime`` and a + variant as a ``Variant``. Echoed unchanged they arrive as the binding's own wrapper, which + still holds the original message and can be copied outright. + """ + full_name = fd.message_type.full_name + + # Echoed unchanged: the binding wrapper kept the message it read. + inner = getattr(value, "msg", None) + if isinstance(inner, message.Message): + target.CopyFrom(inner) + return + if isinstance(value, message.Message): + target.CopyFrom(value) + return + + if full_name == _DECIMAL_TYPE_NAME and isinstance(value, decimal.Decimal): + _set_decimal(target, value) + return + if full_name == _TIMESTAMP_TYPE_NAME and isinstance(value, datetime.datetime): + _set_timestamp(target, value) + return + if full_name == _VARIANT_TYPE_NAME and isinstance(value, Variant): + target.metadata = bytes(value.metadata) + # standalone_value_bytes, not .value: a navigated sub-variant's own value starts at its + # position, and .value is the whole shared buffer - writing it back reconstructs the + # parent document instead of the selected value. + target.value = bytes(value.standalone_value_bytes()) + return + + # A wrapper, or a Duration. The CEL binding unwraps these on the way in - a StringValue + # field is bound as a plain string, a Duration as a CEL duration (see + # _MSG_TYPE_URL_TO_CTOR in constraints.py) - so the inverse has to put them back. Without + # it an identity transform over such a field wrote an empty message and the value was + # silently lost. The JVM gets this for free: its message-level write-back goes through + # protobuf JSON, whose parser reads "hello" into a StringValue and "3s" into a Duration. + if full_name in _WRAPPER_TYPES: + _set_wrapper(target, value) + return + if full_name == _DURATION_TYPE_NAME and isinstance(value, datetime.timedelta): + _set_duration(target, value) + return + + # A nested message the rule rebuilt field by field. + if isinstance(value, Mapping): + _fill(target, value) + return + + # A message the rule echoed unchanged rather than rebuilding. + if isinstance(value, message.Message) and value.DESCRIPTOR.full_name == full_name: + target.CopyFrom(value) + return + + raise ValueError(f"cannot write {type(value).__name__} to {full_name} (field '{fd.name}')") + + +def _set_duration(target: message.Message, value: datetime.timedelta) -> None: + """Sets a ``google.protobuf.Duration`` from a timedelta. + + Split by truncation toward zero, not by floor division: a Duration's ``seconds`` and + ``nanos`` must carry the same sign, whereas timedelta normalises to a non-negative + microseconds component (-3.5s is stored as days=-1, seconds=86396, microseconds=500000). + """ + total_us = (value.days * 86400 + value.seconds) * 1_000_000 + value.microseconds + sign = -1 if total_us < 0 else 1 + magnitude = abs(total_us) + target.seconds = sign * (magnitude // 1_000_000) + target.nanos = sign * (magnitude % 1_000_000) * 1000 + + +def _set_wrapper(target: message.Message, value: Any) -> None: + """Sets a wrapper's single ``value`` field from the scalar the rule returned. + + Narrowed by ``_scalar``, the same as a plain scalar field of that type. This arm had its + own copy of the conversions, and the copy was the unguarded one: an Int32Value took 1.9 as + 1 and a BoolValue took the string "false" as *true*, while the identical plain fields + refused both. The JVM makes no such distinction - JsonFormat's parseWrapperFieldValue + hands the value to the same parseFieldValue a plain field goes through, so the wrapper's + accept/reject set is identical. Measured against protobuf-java 4.35.1: + + * Int32Value <- 1.9, 2147483648, true -> refused; 2.0 -> 2 + * BoolValue <- 0 -> "Invalid bool value: 0" + * FloatValue <- 1.0e40 -> "Out of range float value" + * DoubleValue <- true -> "Not a double value: true" + """ + target.value = _scalar(target.DESCRIPTOR.fields_by_name["value"], value) + + +def _set_decimal(target: message.Message, value: decimal.Decimal) -> None: + sign, digits, exponent = value.as_tuple() + if not isinstance(exponent, int): + raise ValueError("cannot write a non-finite decimal to " + _DECIMAL_TYPE_NAME) + # `scale` is an int32 on the wire. The operators no longer predict which results + # BigDecimal could hold - each library's own exponent range is delegated and documented - + # so a value outside that range now reaches here instead: `decimals.mul` on two + # 1e2147483647 operands is exact and cheap, and needs a scale of -4294967294. Assigning + # it raises a bare `ValueError: Value out of range` from the protobuf runtime; this says + # what the value was and which field could not hold it. + if not (-(2**31) <= -exponent <= 2**31 - 1): + raise ValueError( + f"decimal needs a scale of {-exponent}, which does not fit the int32 scale field " + f"of {_DECIMAL_TYPE_NAME}" + ) + # The coefficient goes out in base 256, and str <-> int radix conversion is quadratic, so + # CPython caps it: `int("9" * 5000)` raises ValueError "Exceeds the limit (4300 digits) for + # integer string conversion". That cap is the real ceiling on what this client can encode - + # far below any width a rule can compute - and it reached callers as a CPython internal + # error naming neither the field nor the decimal. Checked first so it reads as one. + if len(digits) > _MAX_COEFFICIENT_DIGITS: + raise ValueError( + f"decimal coefficient has {len(digits)} digits, past the " + f"{_MAX_COEFFICIENT_DIGITS} this client can encode into {_DECIMAL_TYPE_NAME}" + ) + unscaled = int("".join(str(d) for d in digits) or "0") + if sign: + unscaled = -unscaled + # Negated exponent, negative included - see set_decimal_message in common/protobuf.py. + # Precision is the unscaled value's digit count, as Java's ProtobufResultWriter sets it + # (`m.put("precision", dec.precision())`). Safe to set here because this writer never + # rescales: len(digits) is exactly the digit count of the unscaled value being written, + # so the MathContext the reader builds from it cannot round the value or shift its scale. + target.value = unscaled_to_bytes(unscaled) + target.precision = len(digits) + target.scale = -exponent + + +def _set_timestamp(target: message.Message, value: datetime.datetime) -> None: + if value.tzinfo is None: + value = value.replace(tzinfo=datetime.timezone.utc) + delta = value - _EPOCH + target.seconds = delta.days * 86400 + delta.seconds + target.nanos = delta.microseconds * 1000 + + +def _text(fd: descriptor.FieldDescriptor, value: Any) -> str: + """``value`` as a string field's value, or a rule error.""" + if not isinstance(value, str): + raise ValueError(f"cannot write {type(value).__name__} to string field '{fd.name}'") + return str(value) + + +def _boolean(fd: descriptor.FieldDescriptor, value: Any) -> bool: + """``value`` as a bool field's value, or a rule error. + + Truthiness is not the rule: it read the string "false" as true, and accepted the + numbers protobuf JSON refuses. + """ + if not isinstance(value, (bool, celtypes.BoolType)): + raise ValueError(f"cannot write {type(value).__name__} to bool field '{fd.name}'") + return bool(value) + + +# The widest float32 magnitude, with the 1e-6 slack JsonFormat.parseFloat allows. CEL has one +# floating type, so writing to a `float` field is a narrowing that can overflow; float(1e40) +# gave inf, where the JVM says "Out of range float value: 1.0e40". +_FLOAT32_LIMIT = 3.4028234663852886e38 * (1 + 1e-6) + + +def _is_finite(value: Any) -> bool: + """Whether the value itself is finite. Only a float or a Decimal can be otherwise; a + Python int always is, and ``math.isfinite`` would raise OverflowError on a wide one.""" + if isinstance(value, decimal.Decimal): + return value.is_finite() + if isinstance(value, float): + return math.isfinite(value) + return True + + +def _floating(fd: descriptor.FieldDescriptor, value: Any) -> float: + """``value`` as a float field's value, or a rule error. + + A bool is an int subclass in Python and ``celtypes.BoolType`` subclasses int, so both + spellings have to be named before the numeric check - the same reason ``_integral`` + names them. + + Overflow is judged on the *source*, not the result. ``float()`` saturates a finite but + too-large value to an infinity, and the range check below reads that infinity as one the + rule asked for and lets it through - so ``Decimal("1e1000")`` was silently written as inf + to a double field, and a double field had no range check at all. A wide Python int is the + same case reported differently: ``float(10**400)`` raises OverflowError, which escaped as + a raw Python exception rather than a rule error. + + An *explicitly* non-finite value does pass, because protobuf JSON has canonical spellings + for those and the JVM's parser takes them. Measured against protobuf-java 4.35.1: + + * double <- 1e308 -> 1.0E308 + * double <- 1e309, 1e1000, -1e1000 -> "Out of range double value" + * double <- "Infinity", "-Infinity", "NaN" -> accepted as-is + * float <- 1e39, 1e1000 -> "Out of range float value" + * float <- "Infinity", "NaN" -> accepted as-is + """ + if isinstance(value, (bool, celtypes.BoolType)): + raise ValueError(f"cannot write bool to float field '{fd.name}'") + if not isinstance(value, (int, float, decimal.Decimal)): + raise ValueError(f"cannot write {type(value).__name__} to float field '{fd.name}'") + source_is_finite = _is_finite(value) + try: + as_float = float(value) + except OverflowError as e: + raise ValueError(f"out of range value for float field '{fd.name}': {value}") from e + if not source_is_finite: + return as_float + if not math.isfinite(as_float): + raise ValueError(f"out of range value for float field '{fd.name}': {value}") + if fd.type == descriptor.FieldDescriptor.TYPE_FLOAT and abs(as_float) > _FLOAT32_LIMIT: + raise ValueError(f"out of range float value for field '{fd.name}': {as_float}") + return as_float + + +def _scalar(fd: descriptor.FieldDescriptor, value: Any) -> Any: + """Narrows a celpy value to what protobuf's setter accepts. + + A field takes a value of its own kind, and nothing else. Narrowing unconditionally + accepted wrong-typed results and silently changed their meaning: ``bytes(5)`` + fabricated five NUL bytes out of a number, ``bool("false")`` wrote **true**, and + ``float(True)`` wrote 1.0. A number and a bool are both writable to a *string* field + with `str`, which is worse still - no error, and a rule-authoring mistake becomes data. + + The JVM's write-back renders the result map to JSON and parses it with protobuf's own + JSON parser, so that parser's rejections are the contract, and this matches all of them + (measured against protobuf-java 4.35.1): + + * bool <- 0, "TRUE", "" -> "Invalid bool value" + * bytes <- 5, [97, 98] -> refused + * float <- true -> "Not a double value: true" + * int <- 1.9, true -> "Not an int32 value" + + That parser is also *lenient* in one direction, which this deliberately does not follow: + it stringifies a number or a bool into a string field (1 -> "1"), reads the exact + strings "true"/"false" as a bool, and reads a numeric string as a number. Those are + artifacts of crossing a JSON transport, which this writer does not do - it builds + against the descriptor (see the module note) - and every coercion of that kind turns a + rule-authoring mistake into silently wrong data instead of an error. Refusing them is + the cross-client contract; the JVM accepting a base64 *string* for a bytes field is the + same artifact, so a CEL string is never reinterpreted as bytes either. + + An **enum** is the one exception, and not a coercion: a symbol name is protobuf JSON's + canonical form for an enum and CEL has no enum type, so a string is the only way a rule + can name a symbol. ``_integral`` passes it through for protobuf's own setter to resolve. + """ + if fd.type == descriptor.FieldDescriptor.TYPE_BYTES: + if not isinstance(value, (bytes, bytearray, memoryview)): + raise ValueError(f"cannot write {type(value).__name__} to bytes field '{fd.name}'") + return bytes(value) + if fd.type == descriptor.FieldDescriptor.TYPE_STRING: + return _text(fd, value) + if fd.type == descriptor.FieldDescriptor.TYPE_BOOL: + return _boolean(fd, value) + if fd.type in ( + descriptor.FieldDescriptor.TYPE_FLOAT, + descriptor.FieldDescriptor.TYPE_DOUBLE, + ): + return _floating(fd, value) + # Integer-valued fields (and enums, which take an int). int() silently truncated, so a CEL + # double of 1.9 landed as 1 and a fractional Decimal lost its fraction. The JVM's + # message-level write-back goes through a protobuf JSON parse, which refuses a non-integral + # value ("Not an int32 value: 1.9") while accepting an integral one (2.0 -> 2), and + # range-checks the result. Measured against protobuf-java 4.35.1. + return _integral(fd, value) + + +# Inclusive value ranges for protobuf's integer scalar types, which its JSON parser enforces. +_INT_RANGES = { + descriptor.FieldDescriptor.TYPE_INT32: (-(2**31), 2**31 - 1), + descriptor.FieldDescriptor.TYPE_SINT32: (-(2**31), 2**31 - 1), + descriptor.FieldDescriptor.TYPE_SFIXED32: (-(2**31), 2**31 - 1), + descriptor.FieldDescriptor.TYPE_UINT32: (0, 2**32 - 1), + descriptor.FieldDescriptor.TYPE_FIXED32: (0, 2**32 - 1), + descriptor.FieldDescriptor.TYPE_INT64: (-(2**63), 2**63 - 1), + descriptor.FieldDescriptor.TYPE_SINT64: (-(2**63), 2**63 - 1), + descriptor.FieldDescriptor.TYPE_SFIXED64: (-(2**63), 2**63 - 1), + descriptor.FieldDescriptor.TYPE_UINT64: (0, 2**64 - 1), + descriptor.FieldDescriptor.TYPE_FIXED64: (0, 2**64 - 1), + descriptor.FieldDescriptor.TYPE_ENUM: (-(2**31), 2**31 - 1), +} + + +# Digits in the widest protobuf integer, 2**64-1. A value whose leading digit sits past this +# cannot fit any of the ranges below, so its magnitude settles the question before the digits +# are ever built. +_MAX_INT_DIGITS = 20 + +# The encodable-coefficient ceiling, imported rather than redefined: it used to be a second +# constant named _MAX_COEFFICIENT_DIGITS, colliding with the BigInteger-capacity one in +# common/protobuf.py that carries a different value. +_MAX_COEFFICIENT_DIGITS = MAX_ENCODABLE_COEFFICIENT_DIGITS + + +def _integral(fd: descriptor.FieldDescriptor, value: Any) -> Any: + """``value`` as an int for an integer-valued field, or a rule error. + + A fractional value is a rule-authoring mistake rather than something to round: the JVM + rejects it, and truncating would write a different number than the rule computed. An + integral float or Decimal is accepted, as protobuf JSON accepts ``2.0`` for an int32. + """ + if isinstance(value, (bool, celtypes.BoolType)): + # Both spellings: bool is an int subclass in Python, and celtypes.BoolType subclasses + # int rather than bool - so naming only `bool` caught the case CEL never produces and + # let a real CEL `true` through as 1. protobuf JSON refuses true for an integer field. + raise ValueError(f"cannot write bool to integer field '{fd.name}'") + if isinstance(value, decimal.Decimal): + if not value.is_finite(): + raise ValueError(f"cannot write non-finite {value} to integer field '{fd.name}'") + # Magnitude first, before int() materialises the digits. The widest protobuf integer + # is 2**64-1, twenty digits, so anything with a larger adjusted exponent is out of + # range for every one of them - and int(Decimal("1e100000000")) would spend minutes + # building a hundred million digits just to reach that same rejection, which the JVM + # reports off the token without building the number. + if value.adjusted() > _MAX_INT_DIGITS - 1: + raise ValueError(f"value {value} is out of range for field '{fd.name}'") + if value != value.to_integral_value(): + raise ValueError(f"cannot write non-integral {value} to integer field '{fd.name}'") + as_int = int(value) + elif isinstance(value, float): + if not value.is_integer(): + raise ValueError(f"cannot write non-integral {value!r} to integer field '{fd.name}'") + as_int = int(value) + elif isinstance(value, int): + as_int = int(value) + else: + # Not a number at all - left for protobuf's own setter to reject, which names the + # field and the offending type. + return value + + bounds = _INT_RANGES.get(fd.type) + if bounds is not None and not (bounds[0] <= as_int <= bounds[1]): + raise ValueError(f"value {as_int} is out of range for field '{fd.name}'") + return as_int diff --git a/src/confluent_kafka/schema_registry/rules/cel/timestamp_funcs.py b/src/confluent_kafka/schema_registry/rules/cel/timestamp_funcs.py new file mode 100644 index 000000000..88b8f7ba8 --- /dev/null +++ b/src/confluent_kafka/schema_registry/rules/cel/timestamp_funcs.py @@ -0,0 +1,245 @@ +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""CEL bindings for the {@code timestamp} constructor. + +celpy already provides a stdlib ``timestamp(string)`` (RFC 3339 parsing) plus +the standard timestamp operators (``<``, ``>``, ``==``, ``-``, ``+ duration``, +``.getDate()`` etc.). What we add, by overriding the ``timestamp`` name itself: + + * ``timestamp(int) -> timestamp`` — epoch **seconds**, the overload every + other CEL implementation declares (cel-java's ``int64_to_timestamp``, plus + Go/C++/C#). celpy's base ``timestamp`` is ``TimestampType`` itself, which + accepts a ``datetime``, a ``str``, or an int *plus at least two more args* + (datetime components) but rejects a lone int -- so without this a + single-int call would error on Python only. + * ``timestamp(dyn) -> timestamp`` — the shapes a format decoder produces + that the base implementation doesn't handle: a proto ``Timestamp``, and a + naive ``datetime`` (Avro ``local-timestamp-*``), which is refused rather + than silently read at UTC. + * ``timestamp(int, int) -> timestamp`` — an epoch value at a Flink-style + decimal precision: 0 seconds, 3 millis, 6 micros, 9 nanos. + +Every other form is forwarded to the base implementation verbatim, including +the *datetime components* form (``timestamp(2009, 2, 13)``), which needs three +or more args and so never collides with the two-arg precision form above. + +There is no ``timestamp.of`` namespace any more: these are overloads of the standard +constructor in all seven clients. +""" + +import typing +from datetime import datetime as Datetime +from datetime import timedelta, timezone + +import celpy +from celpy import celtypes + +try: + from google.protobuf.timestamp_pb2 import Timestamp as _ProtoTimestamp +except ImportError: # pragma: no cover + _ProtoTimestamp = None # type: ignore[assignment] + +# celpy's stdlib ``timestamp`` binding. Registering our own "timestamp" entry +# *replaces* it (``Activation.functions`` is a ChainMap where local functions +# shadow the base ones), so we capture the base callable here and delegate to it +# for every form it already handles. +try: + from celpy.evaluation import base_functions as _base_functions + + # celpy's own callable, declared as returning the whole CEL value union rather than a + # TimestampType specifically - so the annotation follows what it hands back. + _BASE_TIMESTAMP: typing.Callable[..., typing.Any] = _base_functions.get("timestamp", celtypes.TimestampType) +except ImportError: # pragma: no cover + _BASE_TIMESTAMP = celtypes.TimestampType + + +_PRECISION_SECONDS = 0 +_PRECISION_MILLIS = 3 +_PRECISION_MICROS = 6 +_PRECISION_NANOS = 9 + +_EPOCH_UTC = Datetime(1970, 1, 1, tzinfo=timezone.utc) + + +def _from_epoch(value: int, precision: int) -> celtypes.TimestampType: + """Construct from an epoch numeric value at a decimal precision. + + Splits the epoch value into whole microseconds using exact integer floor + division (Python ``//`` matches Java ``Math.floorDiv``), then builds the + datetime as ``_EPOCH_UTC + timedelta`` -- so the result floors toward + negative infinity for both positive and negative epochs, and never loses + precision to float rounding (the old ``Datetime.fromtimestamp(value / 1e9)`` + rounded half-to-even to the microsecond). ``datetime`` resolution is one + microsecond, so nanos below that are floored away -- an inherent limit of + the CEL timestamp type, matching Java, not a rounding discrepancy. + + Precisions outside {0, 3, 6, 9} are rejected rather than generalized to + "any p means 10^-p": with the unit a number rather than a name, that check + is the only thing between a typo and a silently wrong instant. + """ + if precision == _PRECISION_SECONDS: + micros = value * 1_000_000 + elif precision == _PRECISION_MILLIS: + micros = value * 1_000 + elif precision == _PRECISION_MICROS: + micros = value + elif precision == _PRECISION_NANOS: + micros = value // 1_000 + else: + raise celpy.CELEvalError( + f"timestamp: unknown precision {precision}; expected 0 (seconds), " "3 (millis), 6 (micros) or 9 (nanos)" + ) + return celtypes.TimestampType(_EPOCH_UTC + timedelta(microseconds=micros)) + + +def _from_proto_timestamp(t: typing.Any) -> celtypes.TimestampType: + """Decode a google.protobuf.Timestamp into a CEL TimestampType. + + Uses exact integer arithmetic: whole seconds plus the nanos field floored + to microseconds (``nanos // 1000``), mirroring the Java reference rather + than the float ``seconds + nanos / 1e9`` that lost precision. + + A nanos field outside the proto contract's ``[0, 999999999]`` normalizes into + the neighbouring instant rather than being rejected, as cel-java does with the + same message; validating it would refuse values the reference accepts. + """ + seconds = int(t.seconds) + nanos = int(t.nanos) + try: + return celtypes.TimestampType(_EPOCH_UTC + timedelta(seconds=seconds, microseconds=nanos // 1_000)) + except (OverflowError, ValueError) as e: + # Normalized like the epoch overloads: an instant ``datetime`` cannot hold is a + # rule error, not a raw Python exception escaping the evaluation. + raise celpy.CELEvalError(f"timestamp: proto Timestamp out of range: {seconds}s {nanos}ns") from e + + +def _timestamp_one(v: typing.Any) -> celtypes.TimestampType: + """The one-argument ``timestamp(dyn)`` dispatch.""" + if v is None: + raise celpy.CELEvalError("timestamp: cannot convert null to Timestamp") + # ``celtypes.BoolType`` subclasses ``int``, *not* ``bool`` (its MRO is + # BoolType -> int -> object), so it has to be named explicitly here or a CEL + # bool falls through to the epoch-seconds branch below and means epoch 1. + if isinstance(v, (bool, celtypes.BoolType)): + raise celpy.CELEvalError("timestamp: cannot convert bool to Timestamp") + if isinstance(v, celtypes.TimestampType): + return v + if isinstance(v, Datetime): + if v.tzinfo is None: + # Avro local-timestamp-* logical types produce naive datetimes that + # carry no timezone — refuse rather than silently picking UTC. + raise celpy.CELEvalError( + "timestamp: naive datetime (no timezone) cannot be converted. " + "Use the regular timestamp-* logical type (UTC by spec), or pass " + "an offset-adjusted epoch value via timestamp(value, precision)." + ) + return celtypes.TimestampType(v) + if _ProtoTimestamp is not None and isinstance(v, _ProtoTimestamp): + return _from_proto_timestamp(v) + # Generic proto Timestamp duck-typing for DynamicMessage / alternate + # generated bindings. + if hasattr(v, "DESCRIPTOR") and getattr(v.DESCRIPTOR, "full_name", "") == "google.protobuf.Timestamp": + return _from_proto_timestamp(v) + if isinstance(v, celtypes.UintType): + # Refused rather than let through: UintType subclasses ``int``, so the arm below would + # take it as epoch seconds, and ``uint`` is a distinct CEL type that no declared + # overload accepts. The typed runtimes report "no matching overload" before evaluating. + raise celpy.CELEvalError("timestamp: expected an int epoch, got uint") + if isinstance(v, (int, celtypes.IntType)): + # A bare int is epoch seconds, matching cel-java's int64_to_timestamp + # and Go/C++/C#. Any other unit needs the two-arg precision form. + try: + return _from_epoch(int(v), _PRECISION_SECONDS) + except (OverflowError, ValueError, OSError) as e: + raise celpy.CELEvalError(f"timestamp: epoch seconds value out of range: {int(v)}") from e + # str (lenient RFC 3339) and anything else the base implementation handles. + return _BASE_TIMESTAMP(v) + + +def _timestamp(*args: typing.Any) -> celtypes.TimestampType: + """CEL stdlib ``timestamp(...)`` plus the epoch-seconds, dyn and precision overloads. + + Three or more args is celpy's *datetime components* form + (``timestamp(2009, 2, 13)``) and is forwarded to the base implementation + verbatim; two args is the epoch + precision form; one arg dispatches on the + value's Python type. + """ + if len(args) == 2: + value, precision = args + # Bools and uints before ints: both subclass int (see _timestamp_one). The reference + # declares this overload (INT, INT), so a uint on either argument has no overload. + if isinstance(value, (bool, celtypes.BoolType, celtypes.UintType)) or not isinstance( + value, (int, celtypes.IntType) + ): + raise celpy.CELEvalError(f"timestamp: epoch value must be int, got {type(value).__name__}") + if isinstance(precision, (bool, celtypes.BoolType, celtypes.UintType)) or not isinstance( + precision, (int, celtypes.IntType) + ): + raise celpy.CELEvalError(f"timestamp: precision must be int, got {type(precision).__name__}") + try: + return _from_epoch(int(value), int(precision)) + except (OverflowError, ValueError, OSError) as e: + # Normalized the same way the one-argument overload does, so an out-of-range epoch + # surfaces as a rule error rather than escaping as a raw Python exception. + raise celpy.CELEvalError(f"timestamp: epoch value out of range: {int(value)}") from e + if len(args) == 1: + return _timestamp_one(args[0]) + return _BASE_TIMESTAMP(*args) + + +def format_timestamp(t: Datetime) -> str: + """Render a timestamp the way every other client's ``string(...)`` does. + + celpy's ``TimestampType.__str__`` formats with + ``strftime("%Y-%m-%dT%H:%M:%S%z")`` -- no ``%f`` -- so it drops the + sub-second component entirely: ``string(timestamp("...T22:13:20.123Z"))`` + came back as ``2023-11-14T22:13:20Z``. The stored value was always right + (comparisons and ``getMilliseconds()`` agreed with the other clients); + only the rendering was lossy, which made it a silent divergence rather + than an error. + + The fraction is emitted in whole 3-digit groups, matching protobuf's + ``Timestamps.toString`` (the Java reference) and the Go, C++, JS and C# + clients: no fraction when it is zero, otherwise 3 digits when the value is + a whole millisecond and 6 when it is not. Java also has a 9-digit + (nanosecond) group, which ``datetime`` cannot represent -- its resolution + is one microsecond -- so a nanosecond-precision value renders 6 digits + here. That is the same pre-existing limit that floors the value itself in + :func:`_from_epoch`, not something this formatting introduces. + + The instant is rendered in UTC with a ``Z`` suffix regardless of the + offset it carries, as the Java reference does. + """ + utc = t.astimezone(timezone.utc) + # Formatted from the numeric components rather than through strftime: %Y is + # platform-dependent below year 1000 (glibc emits "1" where BSD emits "0001"), and RFC + # 3339's date-fullyear is 4DIGIT, which is also what Java's Instant.toString gives. + text = f"{utc.year:04d}-{utc.month:02d}-{utc.day:02d}T" f"{utc.hour:02d}:{utc.minute:02d}:{utc.second:02d}" + micros = utc.microsecond + if micros: + if micros % 1000 == 0: + text += f".{micros // 1000:03d}" + else: + text += f".{micros:06d}" + return text + "Z" + + +# Typed as Any rather than celpy.CELFunction: these functions return this client's own +# Decimal and Variant values, which are not in celpy's declared return union - the CEL +# surface is extended with opaque types celpy does not know. celpy dispatches them fine +# at runtime; only its annotation is narrower than what an extension can return. +TIMESTAMP_FUNCS: typing.Dict[str, typing.Any] = { + "timestamp": _timestamp, +} diff --git a/src/confluent_kafka/schema_registry/rules/cel/variant_funcs.py b/src/confluent_kafka/schema_registry/rules/cel/variant_funcs.py new file mode 100644 index 000000000..f4f009304 --- /dev/null +++ b/src/confluent_kafka/schema_registry/rules/cel/variant_funcs.py @@ -0,0 +1,437 @@ +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""CEL bindings for the ``variant(...)`` constructor and the ``variants.*`` accessor +functions - the Python counterpart of Java's ``rules/cel/builtin`` variant glue. + +celpy has no overload-set concept (one callable per name, internal arity/type dispatch) and +no opaque type, so a :class:`Variant` flows through CEL as a plain Python object, exactly as +``decimal.Decimal`` does for the decimal functions. + +Null model (matching the Java reference / Spark Variant semantics): + +* CEL null (Python ``None``) = *absent*: a missing field, an out-of-bounds index, a + type-mismatched receiver, or a non-Variant input. +* a Variant whose top type is ``NULL`` = *present, but the value is variant-null*. + +``variants.field``/``path``/``index`` return CEL null on a miss; ``variants.isNull`` is true +only for a real Variant with top type NULL. Distinguish the two with +``result == null`` (absent) vs ``variants.isNull(result)`` (variant-null). +""" + +import typing +from datetime import datetime, timedelta, timezone + +import celpy +from celpy import celtypes + +from confluent_kafka.schema_registry.confluent.type import variant_utils as vu +from confluent_kafka.schema_registry.confluent.type.variant_utils import Variant, VariantType +from confluent_kafka.schema_registry.rules.cel import variant_path + +try: + from confluent_kafka.schema_registry.confluent.type import variant_pb2 + + _PROTO_VARIANT_CLS: typing.Any = variant_pb2.Variant +except ImportError: + _PROTO_VARIANT_CLS = None + +_VARIANT_PROTO_NAME = "confluent.type.Variant" +_EPOCH_UTC = datetime(1970, 1, 1, tzinfo=timezone.utc) +_INT32_MAX = 2**31 - 1 + +# VariantType -> the coarse label variants.type returns, matching Java variantTypeName: +# integer widths collapse to "int", float/double to "double", decimal widths to "decimal", +# and all four timestamp variants to "timestamp". +_TYPE_LABELS = { + VariantType.OBJECT: "object", + VariantType.ARRAY: "array", + VariantType.NULL: "null", + VariantType.BOOLEAN: "boolean", + VariantType.BYTE: "int", + VariantType.SHORT: "int", + VariantType.INT: "int", + VariantType.LONG: "int", + VariantType.FLOAT: "double", + VariantType.DOUBLE: "double", + VariantType.DECIMAL4: "decimal", + VariantType.DECIMAL8: "decimal", + VariantType.DECIMAL16: "decimal", + VariantType.DATE: "date", + VariantType.TIME: "time", + VariantType.TIMESTAMP_TZ: "timestamp", + VariantType.TIMESTAMP_NTZ: "timestamp", + VariantType.TIMESTAMP_NANOS_TZ: "timestamp", + VariantType.TIMESTAMP_NANOS_NTZ: "timestamp", + VariantType.STRING: "string", + VariantType.BINARY: "bytes", + VariantType.UUID: "uuid", +} + +_INT_TYPES = (VariantType.BYTE, VariantType.SHORT, VariantType.INT, VariantType.LONG) +_DECIMAL_TYPES = (VariantType.DECIMAL4, VariantType.DECIMAL8, VariantType.DECIMAL16) +_TIMESTAMP_TYPES = ( + VariantType.TIMESTAMP_TZ, + VariantType.TIMESTAMP_NTZ, + VariantType.TIMESTAMP_NANOS_TZ, + VariantType.TIMESTAMP_NANOS_NTZ, +) +_MICROS_TIMESTAMP_TYPES = (VariantType.TIMESTAMP_TZ, VariantType.TIMESTAMP_NTZ) + + +def _coerce_bytes(v: typing.Any) -> bytes: + if isinstance(v, (bytes, bytearray)): + return bytes(v) + if isinstance(v, memoryview): + return v.tobytes() + if isinstance(v, celtypes.BytesType): + return bytes(v) + raise celpy.CELEvalError(f"variant: expected bytes, got {type(v).__name__}") + + +def _variant_or_absent(value: bytes, metadata: bytes) -> typing.Optional[Variant]: + """Build a Variant, or ``None`` when there is no metadata at all -- a protobuf field left + unset, or an Avro variant record whose byte fields are empty. The Variant constructor reads + the metadata version byte, so an empty buffer has to be caught before construction; callers + report the ``None`` as CEL null, which every ``variants.*`` accessor already propagates.""" + if not metadata: + return None + return Variant(value, metadata) + + +def _to_variant(v: typing.Any) -> typing.Optional[Variant]: + """Runtime dispatch backing ``variant(dyn)``: accept the shapes proto/Avro decoders + produce. Rejects strings (use ``variants.parseJson``). CEL null passes through as CEL + null (matching the Java reference: ``variant(null) -> CEL null``).""" + if v is None: + return None + if isinstance(v, Variant): + return v + # A confluent.type.Variant proto message: the generated class, or any message whose + # descriptor full name matches (covers DynamicMessage / alternate bindings). + if _PROTO_VARIANT_CLS is not None and isinstance(v, _PROTO_VARIANT_CLS): + return _variant_or_absent(_coerce_bytes(v.value), _coerce_bytes(v.metadata)) + if getattr(getattr(v, "DESCRIPTOR", None), "full_name", "") == _VARIANT_PROTO_NAME: + return _variant_or_absent(_coerce_bytes(v.value), _coerce_bytes(v.metadata)) + # celpy binds a proto-message field as a wrapper that keeps the message on ``.msg``. + proto_msg = getattr(v, "msg", None) + if ( + proto_msg is not None + and getattr(getattr(proto_msg, "DESCRIPTOR", None), "full_name", "") == _VARIANT_PROTO_NAME + ): + return _variant_or_absent(_coerce_bytes(proto_msg.value), _coerce_bytes(proto_msg.metadata)) + # An Avro variant-logical field reaches CEL as a map with {"metadata", "value"} byte + # entries (celpy MapType is a dict subclass). + if isinstance(v, dict): + md = v.get("metadata") + val = v.get("value") + if md is not None and val is not None: + return _variant_or_absent(_coerce_bytes(val), _coerce_bytes(md)) + if md is not None or val is not None: + missing = "value" if val is None else "metadata" + raise celpy.CELEvalError(f"variant: cannot convert map to Variant: missing '{missing}' entry") + if isinstance(v, (str, celtypes.StringType)): + raise celpy.CELEvalError( + "variant: cannot convert string to Variant; use variants.parseJson(s) for " + "strict JSON parsing or variants.tryParseJson(s) for soft mode" + ) + raise celpy.CELEvalError(f"variant: cannot convert {type(v).__name__} to Variant") + + +def _variant(*args: typing.Any) -> typing.Optional[Variant]: + """The ``variant(...)`` constructor. ``variant(dyn)`` runtime-dispatches (CEL null in -> + CEL null out); ``variant(bytes, bytes)`` builds directly from (value, metadata) bytes.""" + if len(args) == 2: + metadata = _coerce_bytes(args[1]) + if not metadata: + # Passing empty metadata explicitly is a rule-authoring mistake rather than an + # absent field, so it is reported instead of yielding null. + raise celpy.CELEvalError("variant(value, metadata): metadata is empty, so there is no variant to read") + return Variant(_coerce_bytes(args[0]), metadata) + if len(args) != 1: + raise celpy.CELEvalError(f"variant: expected 1 or 2 args, got {len(args)}") + return _to_variant(args[0]) + + +def _parse_json(s: typing.Any) -> Variant: + """``variants.parseJson(string)`` - strict; raises on malformed JSON.""" + if not isinstance(s, (str, celtypes.StringType)): + raise celpy.CELEvalError("variants.parseJson: expected a string") + try: + return vu.parse_json(str(s)) + except (vu.VariantError, ValueError) as ex: + raise celpy.CELEvalError(f"variants.parseJson: {ex}") from ex + + +def _try_parse_json(s: typing.Any) -> typing.Optional[Variant]: + """``variants.tryParseJson(string)`` - soft; CEL null on any parse failure. + + Soft about *parsing*, not about the argument type: Java declares this binding over + ``String`` exactly as it does the strict form, so a non-string argument is a rule error + there. Stringifying instead turned ``tryParseJson(123)`` into a numeric variant. + """ + if not isinstance(s, (str, celtypes.StringType)): + raise celpy.CELEvalError("variants.tryParseJson: expected a string") + try: + return vu.parse_json(str(s)) + except Exception: # noqa: BLE001 - soft form: any parse failure -> CEL null + return None + + +def _type(v: typing.Any) -> typing.Any: + """``variants.type(Variant)`` - the type label as a string; propagates CEL null.""" + variant = _require_variant_or_null(v, "variants.type") + if variant is None: + return None + return celtypes.StringType(_TYPE_LABELS[variant.get_type()]) + + +def _is_null(o: typing.Any) -> celtypes.BoolType: + """``variants.isNull(dyn)`` - true iff input is a Variant whose top type is NULL. + + Coerces like every other accessor. ``isinstance(o, Variant)`` alone answered False for the + shapes a variant-typed *field* decodes to - a celpy MessageType wrapping a + confluent.type.Variant, or the map an Avro variant record yields - which the dyn declaration + admits, so a bare variant holding an explicit JSON null reported "not null". CEL null and any + non-variant stay False rather than raising: this predicate never errors. + """ + if isinstance(o, Variant): + return celtypes.BoolType(o.get_type() == VariantType.NULL) + try: + coerced = _require_variant_or_null(o, "variants.isNull") + except Exception: + return celtypes.BoolType(False) + return celtypes.BoolType(coerced is not None and coerced.get_type() == VariantType.NULL) + + +def _require_variant_or_null(o: typing.Any, fn: str) -> typing.Optional[Variant]: + """A ``variants.*`` navigation argument: CEL null passes through as ``None``; a real + Variant is returned; anything else is a hard error (the dyn signature lets a misused + non-Variant reach the binding).""" + if o is None: + return None + if isinstance(o, Variant): + return o + # Otherwise accept the shapes a variant-typed *field* decodes to -- a proto + # confluent.type.Variant message (celpy wraps one as a MessageType), or the mapping an + # Avro variant record decodes to -- so such a field can be used without a variant(...) + # call, as on the other clients. An unrecognized shape still gets a clear message. + try: + return _to_variant(o) + except celpy.CELEvalError as ex: + raise celpy.CELEvalError(f"{fn}: expected Variant, got {type(o).__name__}") from ex + return o + + +def _path(o: typing.Any, path: typing.Any) -> typing.Optional[Variant]: + """``variants.path(dyn, string)`` - navigate a JSONPath subset; CEL null on a miss; + malformed path raises.""" + # Java declares this `(DYN, STRING)`, so a non-string path has no matching overload there + # whatever the receiver holds - and Go, JS and C++ enforce it the same way, in the declared + # overload. celpy has no overload-set concept, so the check has to be explicit; without it + # `str(path)` looked up the path "1" for variants.path(v, 1). Checked *before* the receiver + # for the same reason: the argument error does not depend on the receiver's shape, so a + # null receiver must not turn a mistyped call into CEL null. Same contract as + # variants.field, variants.index, variants.as and variants.tryAs. + if not isinstance(path, (str, celtypes.StringType)): + raise celpy.CELEvalError(f"variants.path: expected a string path, got {type(path).__name__}") + v = _require_variant_or_null(o, "variants.path") + if v is None: + return None + try: + return variant_path.walk(v, str(path)) + except ValueError as ex: + raise celpy.CELEvalError(f"variants.path: {ex}") from ex + + +def _field(o: typing.Any, key: typing.Any) -> typing.Optional[Variant]: + """``variants.field(dyn, string)`` - object field by key; CEL null on a miss or a + non-object receiver.""" + # Java binds this as (Object, String), so a non-string key has no matching overload; + # str(key) would have looked up "1" for variants.field(v, 1). Same contract as + # variants.index and variants.parseJson. + if not isinstance(key, (str, celtypes.StringType)): + raise celpy.CELEvalError(f"variants.field: expected a string key, got {type(key).__name__}") + v = _require_variant_or_null(o, "variants.field") + if v is None or v.get_type() != VariantType.OBJECT: + return None + return v.get_field_by_key(str(key)) + + +def _index(o: typing.Any, idx: typing.Any) -> typing.Optional[Variant]: + """``variants.index(dyn, int)`` - array element by index; CEL null on out-of-bounds or a + non-array receiver.""" + # The index is checked before the receiver, the way variants.field checks its key. Java + # declares this overload as (DYN, INT), so a double or a bool fails to bind whatever the + # receiver turns out to hold; checking the receiver first made the argument's own type + # depend on it, and `variants.index(anObject, 1.5)` answered CEL null instead of + # reporting the wrong type. int() would then have quietly floored 1.9 to element 1 and + # read true as element 1. + # UintType alongside the bools because it subclasses ``int`` too; the reference declares + # variants.index as (DYN, INT), so `variants.index(v, 1u)` has no matching overload. + if isinstance(idx, (bool, celtypes.BoolType, celtypes.UintType)) or not isinstance(idx, (int, celtypes.IntType)): + raise celpy.CELEvalError(f"variants.index: expected an int index, got {type(idx).__name__}") + v = _require_variant_or_null(o, "variants.index") + if v is None or v.get_type() != VariantType.ARRAY: + return None + i = int(idx) + if i < 0 or i > _INT32_MAX: + return None + return v.get_element_at_index(i) + + +# The CEL timestamp range in epoch seconds, matching the reference's +# TimestampUtils.MIN/MAX_EPOCH_SECOND. +_MIN_EPOCH_SECOND = -62135596800 +_MAX_EPOCH_SECOND = 253402300799 + + +def _variant_get_timestamp(v: Variant) -> typing.Optional[celtypes.TimestampType]: + """Backing for ``variants.as(v, 'timestamp')``. + + MICROS-precision variants (TIMESTAMP_TZ/NTZ) store microseconds since the epoch and + are used as-is. NANOS-precision variants (TIMESTAMP_NANOS_TZ/NTZ) store nanoseconds: + the Java reference (TimestampUtils.fromEpochNanos) splits them with + ``Math.floorDiv``/``Math.floorMod`` into (seconds, nanos) and keeps the full + nanosecond field in a protobuf ``Timestamp``. celpy's ``TimestampType`` is a subclass + of ``datetime.datetime``, whose finest resolution is one microsecond, so the epoch + value is floored to microseconds. Python ``//`` is floor division (matching Java + ``Math.floorDiv``), so pre-epoch (negative) values round toward negative infinity + exactly as Java does -- e.g. -1 ns -> -1 us, never 0. The only residual difference + from Java is the sub-microsecond nanoseconds that a ``datetime`` cannot represent; + this is an inherent limit of the CEL timestamp type, not a rounding discrepancy. + + Returns ``None`` when the value falls outside the CEL timestamp range (0001-9999). A + variant timestamp spans the whole int64 range, so that is reachable from data; ``None`` + rather than an exception because the caller decides -- ``variants.as`` raises and names the + range, ``variants.tryAs`` answers CEL null, the same split those two already apply to a type + mismatch. ``datetime`` cannot represent such a value at all, so the arithmetic below would + raise ``OverflowError`` and escape ``tryAs`` as well. + """ + raw = v.get_long() + micros = raw if v.get_type() in _MICROS_TIMESTAMP_TYPES else raw // 1000 + seconds = micros // 1_000_000 + if seconds < _MIN_EPOCH_SECOND or seconds > _MAX_EPOCH_SECOND: + return None + return celtypes.TimestampType(_EPOCH_UTC + timedelta(microseconds=micros)) + + +def _variant_as(o: typing.Any, type_str: str, null_on_error: bool) -> typing.Any: + """Backing for ``variants.as`` (strict) and ``variants.tryAs`` (soft). Extracts a typed + value; on a type mismatch the strict form raises and the soft form returns CEL null. + Types with no CEL extraction (object/array/null/date/time/uuid) always raise.""" + fn = "variants.tryAs" if null_on_error else "variants.as" + v = _require_variant_or_null(o, fn) + if v is None: + return None + t = v.get_type() + if type_str == "string": + if t == VariantType.STRING: + return celtypes.StringType(v.get_string()) + elif type_str == "int": + if t in _INT_TYPES: + return celtypes.IntType(v.get_long()) + elif type_str == "double": + if t == VariantType.FLOAT: + return celtypes.DoubleType(float(v.get_float())) + if t == VariantType.DOUBLE: + return celtypes.DoubleType(v.get_double()) + elif type_str == "boolean": + if t == VariantType.BOOLEAN: + return celtypes.BoolType(v.get_boolean()) + elif type_str == "decimal": + if t in _DECIMAL_TYPES: + return v.get_decimal() + elif type_str == "timestamp": + if t in _TIMESTAMP_TYPES: + ts = _variant_get_timestamp(v) + if ts is not None: + return ts + if null_on_error: + return None + raise celpy.CELEvalError( + f"variants.as: timestamp {v.get_long()} is outside " + "0001-01-01T00:00:00Z..9999-12-31T23:59:59.999999999Z" + ) + elif type_str == "bytes": + if t == VariantType.BINARY: + return celtypes.BytesType(v.get_binary()) + elif type_str in ("object", "array", "null", "date", "time", "uuid"): + # Not extractable as a CEL scalar - always an error, even in the soft form. + raise celpy.CELEvalError( + f"variants.as: type '{type_str}' is not supported for extraction " + "(use variants.type/variants.path/variants.field/variants.index instead)" + ) + else: + if null_on_error: + return None + raise celpy.CELEvalError( + f"variants.as: unknown type '{type_str}' (expected one of: string, int, " + "double, boolean, decimal, timestamp, bytes)" + ) + # Recognized type string, but the variant's actual type does not match. + if null_on_error: + return None + raise celpy.CELEvalError(f"variants.as: variant is not {type_str}-typed (type={t.value})") + + +def _require_type_name(type_str: typing.Any, fn: str) -> str: + """The type-name argument as a string. + + Java declares both overloads as (DYN, STRING), so a non-string second argument fails to + bind. `str()` accepted anything, which mattered most for the soft form: `variants.tryAs(v, + 1)` stringified to "1", took the unknown-type branch and returned CEL **null**. Null is + tryAs's answer for a type *mismatch*, not for a call that names no type at all, so a + wrong-typed argument was indistinguishable from a variant of the wrong shape. + """ + if not isinstance(type_str, (str, celtypes.StringType)): + raise celpy.CELEvalError(f"{fn}: expected a string type name, got {type(type_str).__name__}") + return str(type_str) + + +def _as(o: typing.Any, type_str: typing.Any) -> typing.Any: + """``variants.as(dyn, string)`` - typed extraction; raises on type mismatch.""" + return _variant_as(o, _require_type_name(type_str, "variants.as"), null_on_error=False) + + +def _try_as(o: typing.Any, type_str: typing.Any) -> typing.Any: + """``variants.tryAs(dyn, string)`` - typed extraction; CEL null on type mismatch.""" + return _variant_as(o, _require_type_name(type_str, "variants.tryAs"), null_on_error=True) + + +def _to_json(v: typing.Any) -> typing.Any: + """``variants.toJson(Variant)`` - serialize to a JSON string; propagates CEL null.""" + variant = _require_variant_or_null(v, "variants.toJson") + if variant is None: + return None + return celtypes.StringType(variant.to_json()) + + +# Typed as Any rather than celpy.CELFunction: these functions return this client's own +# Decimal and Variant values, which are not in celpy's declared return union - the CEL +# surface is extended with opaque types celpy does not know. celpy dispatches them fine +# at runtime; only its annotation is narrower than what an extension can return. +VARIANT_FUNCS: typing.Dict[str, typing.Any] = { + "variant": _variant, + "variants.parseJson": _parse_json, + "variants.tryParseJson": _try_parse_json, + "variants.type": _type, + "variants.isNull": _is_null, + "variants.path": _path, + "variants.field": _field, + "variants.index": _index, + "variants.as": _as, + "variants.tryAs": _try_as, + "variants.toJson": _to_json, +} diff --git a/src/confluent_kafka/schema_registry/rules/cel/variant_path.py b/src/confluent_kafka/schema_registry/rules/cel/variant_path.py new file mode 100644 index 000000000..b5a21829a --- /dev/null +++ b/src/confluent_kafka/schema_registry/rules/cel/variant_path.py @@ -0,0 +1,175 @@ +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""The JSONPath subset used by ``variants.path(v, path)`` - a port of Java's +``VariantPath``. Supports: + +* ``$`` - root +* ``$.field`` / ``$.field.subfield`` - object field by identifier name +* ``$[i]`` - array element by non-negative integer index +* ``$["quoted key"]`` / ``$['quoted key']`` - quoted key for non-identifier names + +Resolution failures (missing field, out-of-bounds index, type mismatch) return ``None`` +from :func:`walk`; malformed paths raise :class:`ValueError` at parse time. + +Identifier names follow ``[letter_][letter digit _]*``, where letter and digit are +Unicode-aware (``str.isalpha`` / ``str.isalnum``), so accented and non-Latin names are +identifiers too; use the quoted form for any other key. +Negative indices are rejected (no RFC 9535 ``len + i`` semantics). + +Quoted-key escapes recognize only ``\\\\`` (a literal backslash) and backslash + the +enclosing quote; any other escape is a parse error rather than being silently decoded +(option B, matching the Java reference). Non-ASCII characters may be written literally; +for keys needing escapes beyond these two, use ``variants.field(v, key)`` with a regular +CEL string. +""" + +from functools import lru_cache +from typing import List, Optional, Tuple + +from confluent_kafka.schema_registry.confluent.type.variant_utils import Variant, VariantType + +# A parsed path is a list of segments. Each segment is a ("field", key) or ("index", idx) pair. +Segment = Tuple[str, object] + + +def walk(root: Variant, path: str) -> Optional[Variant]: + """Walk ``root`` following ``path``. Returns the resolved Variant, or ``None`` if any + segment fails to resolve. Raises :class:`ValueError` on a malformed path.""" + current: Optional[Variant] = root + for kind, arg in parse(path): + if current is None: + return None + if kind == "field": + current = ( + current.get_field_by_key(arg) # type: ignore[arg-type] + if current.get_type() == VariantType.OBJECT + else None + ) + else: # "index" + current = ( + current.get_element_at_index(arg) # type: ignore[arg-type] + if current.get_type() == VariantType.ARRAY + else None + ) + return current + + +@lru_cache(maxsize=1000) +def parse(path: str) -> Tuple[Segment, ...]: + """Parse ``path`` into segments. Cached, since rules usually pass a literal path that + recurs per record. ``lru_cache`` does not cache exceptions, so a malformed path raises + on every call (matching the Java LoadingCache behavior).""" + if not path: + raise ValueError("variant path must start with '$'") + cur = _Cursor(path) + if cur.peek() != "$": + raise ValueError("variant path must start with '$', got: " + path) + cur.next() + out: List[Segment] = [] + while cur.has_more(): + ch = cur.peek() + if ch == ".": + cur.next() + out.append(("field", _read_ident(cur, path))) + elif ch == "[": + cur.next() + if not cur.has_more(): + raise ValueError("unexpected end of input after '[' in variant path: " + path) + if cur.peek() in ("\"", "'"): + out.append(("field", _read_quoted_key(cur, path))) + else: + out.append(("index", _read_index(cur, path))) + if not cur.has_more() or cur.next() != "]": + raise ValueError("expected ']' in variant path: " + path) + else: + raise ValueError("unexpected character '" + ch + "' in variant path: " + path) + return tuple(out) + + +def _read_ident(cur: "_Cursor", path: str) -> str: + if not cur.has_more() or not (cur.peek().isalpha() or cur.peek() == "_"): + raise ValueError("expected identifier (starting with a letter or '_') after '.' in variant path: " + path) + start = cur.pos + cur.next() + while cur.has_more(): + ch = cur.peek() + if ch.isalnum() or ch == "_": + cur.next() + else: + break + return cur.src[start : cur.pos] + + +def _read_quoted_key(cur: "_Cursor", path: str) -> str: + quote = cur.next() + out = [] + while cur.has_more(): + ch = cur.next() + if ch == "\\": + # Only two escapes are recognized: a doubled backslash for a literal backslash, + # and backslash + the enclosing quote for a literal quote. Any other escape - + # including a would-be Unicode escape like backslash-u00e9 - is a parse error + # rather than being silently decoded to the wrong key. Literal characters + # (including non-ASCII) need no escaping and pass through as-is. + if not cur.has_more(): + raise ValueError("unterminated escape at end of quoted key in variant path: " + path) + esc = cur.next() + if esc == "\\" or esc == quote: + out.append(esc) + else: + raise ValueError( + "unsupported escape '\\" + esc + "' in quoted key of variant path " + "(only '\\\\' and '\\" + quote + "' are allowed): " + path + ) + elif ch == quote: + return "".join(out) + else: + out.append(ch) + raise ValueError("unterminated quoted key in variant path: " + path) + + +def _read_index(cur: "_Cursor", path: str) -> int: + if cur.has_more() and cur.peek() == "-": + raise ValueError("negative indices are not supported in variant path: " + path) + start = cur.pos + # ASCII digits only, and bounded like Java's Integer.parseInt: a wider index + # is an error rather than an arbitrarily large Python int. + while cur.has_more() and "0" <= cur.peek() <= "9": + cur.next() + if cur.pos == start: + raise ValueError("expected integer index in variant path: " + path) + index = int(cur.src[start : cur.pos]) + if index > 2147483647: + raise ValueError("index out of int range in variant path: " + path) + return index + + +class _Cursor: + __slots__ = ("src", "pos") + + def __init__(self, src: str): + self.src = src + self.pos = 0 + + def has_more(self) -> bool: + return self.pos < len(self.src) + + def peek(self) -> str: + return self.src[self.pos] + + def next(self) -> str: + ch = self.src[self.pos] + self.pos += 1 + return ch diff --git a/src/confluent_kafka/src/Admin.c b/src/confluent_kafka/src/Admin.c index 1ea879a5b..38630989d 100644 --- a/src/confluent_kafka/src/Admin.c +++ b/src/confluent_kafka/src/Admin.c @@ -1185,7 +1185,7 @@ Admin_describe_configs(Handle *self, PyObject *args, PyObject *kwargs) { #ifdef Py_GIL_DISABLED Py_XDECREF(owned_resources); #endif - Py_XDECREF(ConfigResource_type); /* from lookup() */ + Py_XDECREF(ConfigResource_type); /* from lookup() */ /* Release our extra ref only on failure; on success the opaque keeps * it (see options_to_c()). */ if (future_incremented && !result) @@ -1360,8 +1360,8 @@ static PyObject *Admin_incremental_alter_configs(Handle *self, if (rkqu) rd_kafka_queue_destroy(rkqu); /* drop ref from get_background */ Py_XDECREF(owned_resources); - Py_XDECREF(ConfigResource_type); /* from lookup() */ - Py_XDECREF(ConfigEntry_type); /* from lookup() */ + Py_XDECREF(ConfigResource_type); /* from lookup() */ + Py_XDECREF(ConfigEntry_type); /* from lookup() */ /* Release our extra ref only on failure; on success the opaque keeps * it (see options_to_c()). */ if (future_incremented && !result) @@ -1531,7 +1531,7 @@ Admin_alter_configs(Handle *self, PyObject *args, PyObject *kwargs) { #ifdef Py_GIL_DISABLED Py_XDECREF(owned_resources); #endif - Py_XDECREF(ConfigResource_type); /* from lookup() */ + Py_XDECREF(ConfigResource_type); /* from lookup() */ /* Release our extra ref only on failure; on success the opaque keeps * it (see options_to_c()). */ if (future_incremented && !result) @@ -2150,9 +2150,9 @@ static PyObject *Admin_describe_user_scram_credentials(Handle *self, #ifdef Py_GIL_DISABLED PyObject *owned_users = NULL; #endif - rd_kafka_queue_t *rkqu = NULL; - PyObject *result = NULL; - int future_incremented = 0; + rd_kafka_queue_t *rkqu = NULL; + PyObject *result = NULL; + int future_incremented = 0; CallState cs; /* users is a list of strings. */ @@ -2280,11 +2280,11 @@ static PyObject *Admin_alter_user_scram_credentials(Handle *self, #ifdef Py_GIL_DISABLED PyObject *owned_alterations = NULL; #endif - PyObject *UserScramCredentialAlteration_type = NULL; - PyObject *UserScramCredentialUpsertion_type = NULL; - PyObject *UserScramCredentialDeletion_type = NULL; - PyObject *ScramCredentialInfo_type = NULL; - PyObject *ScramMechanism_type = NULL; + PyObject *UserScramCredentialAlteration_type = NULL; + PyObject *UserScramCredentialUpsertion_type = NULL; + PyObject *UserScramCredentialDeletion_type = NULL; + PyObject *ScramCredentialInfo_type = NULL; + PyObject *ScramMechanism_type = NULL; rd_kafka_queue_t *rkqu; CallState cs; @@ -2621,8 +2621,8 @@ Admin_describe_consumer_groups(Handle *self, PyObject *args, PyObject *kwargs) { rd_kafka_AdminOptions_t *c_options = NULL; CallState cs; rd_kafka_queue_t *rkqu; - int groups_cnt = 0; - int i = 0; + int groups_cnt = 0; + int i = 0; int entered_rk_use = 0; static char *kws[] = {"future", "group_ids", @@ -3853,7 +3853,8 @@ static PyObject *Admin_exit(Handle *self, PyObject *args) { * flushing and destroying it. */ if (!atomic_int_cas(&self->closing, 0, 1)) { - while (atomic_ptr_get(&self->rk) && atomic_int_get(&self->closing)) { + while (atomic_ptr_get(&self->rk) && + atomic_int_get(&self->closing)) { if (!Handle_sleep(self, 100)) return NULL; } diff --git a/src/confluent_kafka/src/Consumer.c b/src/confluent_kafka/src/Consumer.c index 4f4f56983..38b7eb0ba 100644 --- a/src/confluent_kafka/src/Consumer.c +++ b/src/confluent_kafka/src/Consumer.c @@ -112,15 +112,14 @@ Consumer_subscribe(Handle *self, PyObject *args, PyObject *kwargs) { NULL}; PyObject *tlist, *on_assign = NULL, *on_revoke = NULL, *on_lost = NULL; PyObject *result = NULL; - Py_ssize_t pos = 0; + Py_ssize_t pos = 0; rd_kafka_resp_err_t err; #ifdef Py_GIL_DISABLED PyObject *owned_tlist = NULL; #endif - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|OOO", kws, - &tlist, &on_assign, - &on_revoke, &on_lost)) + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|OOO", kws, &tlist, + &on_assign, &on_revoke, &on_lost)) return NULL; if (!Handle_serialize_enter(self)) @@ -255,8 +254,7 @@ static PyObject *Consumer_unsubscribe(Handle *self, PyObject *ignore) { } -static PyObject * -Consumer_incremental_assign(Handle *self, PyObject *tlist) { +static PyObject *Consumer_incremental_assign(Handle *self, PyObject *tlist) { PyObject *result = NULL; rd_kafka_topic_partition_list_t *c_parts; rd_kafka_error_t *error; @@ -356,8 +354,7 @@ static PyObject *Consumer_unassign(Handle *self, PyObject *ignore) { return result; } -static PyObject * -Consumer_incremental_unassign(Handle *self, PyObject *tlist) { +static PyObject *Consumer_incremental_unassign(Handle *self, PyObject *tlist) { PyObject *result = NULL; rd_kafka_topic_partition_list_t *c_parts; rd_kafka_error_t *error; @@ -393,8 +390,7 @@ Consumer_incremental_unassign(Handle *self, PyObject *tlist) { } static PyObject * -Consumer_assignment(Handle *self, PyObject *args, - PyObject *kwargs) { +Consumer_assignment(Handle *self, PyObject *args, PyObject *kwargs) { PyObject *result = NULL; rd_kafka_topic_partition_list_t *c_parts; @@ -529,9 +525,8 @@ Consumer_commit(Handle *self, PyObject *args, PyObject *kwargs) { return NULL; } - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "|OOOO", kws, - &msg, &offsets, &async_o, - &async_o)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "|OOOO", kws, &msg, + &offsets, &async_o, &async_o)) { Handle_serialize_exit(self); return NULL; } @@ -571,7 +566,7 @@ Consumer_commit(Handle *self, PyObject *args, PyObject *kwargs) { PyObject *error; PyObject *topic; - m = (Message *)msg; + m = (Message *)msg; error = Message_error(m, NULL); if (error != Py_None) { PyObject *errstr = @@ -587,7 +582,7 @@ Consumer_commit(Handle *self, PyObject *args, PyObject *kwargs) { } Py_DECREF(error); - topic = Message_topic(m, NULL); + topic = Message_topic(m, NULL); c_offsets = rd_kafka_topic_partition_list_new(1); rktpar = rd_kafka_topic_partition_list_add( c_offsets, cfl_PyUnistr_AsUTF8(topic, &uo8), m->partition); @@ -663,8 +658,7 @@ Consumer_commit(Handle *self, PyObject *args, PyObject *kwargs) { } static PyObject * -Consumer_store_offsets(Handle *self, PyObject *args, - PyObject *kwargs) { +Consumer_store_offsets(Handle *self, PyObject *args, PyObject *kwargs) { #if RD_KAFKA_VERSION < 0x000b0000 PyErr_Format(PyExc_NotImplementedError, "Consumer store_offsets require " @@ -680,8 +674,8 @@ Consumer_store_offsets(Handle *self, PyObject *args, rd_kafka_topic_partition_list_t *c_offsets; static char *kws[] = {"message", "offsets", NULL}; - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "|OO", kws, - &msg, &offsets)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "|OO", kws, &msg, + &offsets)) { return NULL; } @@ -726,7 +720,7 @@ Consumer_store_offsets(Handle *self, PyObject *args, goto done; } - m = (Message *)msg; + m = (Message *)msg; error = Message_error(m, NULL); if (error != Py_None) { PyObject *errstr = @@ -741,7 +735,7 @@ Consumer_store_offsets(Handle *self, PyObject *args, } Py_DECREF(error); - topic = Message_topic(m, NULL); + topic = Message_topic(m, NULL); c_offsets = rd_kafka_topic_partition_list_new(1); rktpar = rd_kafka_topic_partition_list_add( c_offsets, cfl_PyUnistr_AsUTF8(topic, &uo8), m->partition); @@ -775,8 +769,7 @@ Consumer_store_offsets(Handle *self, PyObject *args, static PyObject * -Consumer_committed(Handle *self, PyObject *args, - PyObject *kwargs) { +Consumer_committed(Handle *self, PyObject *args, PyObject *kwargs) { PyObject *plist; PyObject *result = NULL; @@ -785,8 +778,8 @@ Consumer_committed(Handle *self, PyObject *args, double tmout = -1.0f; static char *kws[] = {"partitions", "timeout", NULL}; - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|d", kws, - &plist, &tmout)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|d", kws, &plist, + &tmout)) { return NULL; } @@ -823,8 +816,7 @@ Consumer_committed(Handle *self, PyObject *args, } static PyObject * -Consumer_position(Handle *self, PyObject *args, - PyObject *kwargs) { +Consumer_position(Handle *self, PyObject *args, PyObject *kwargs) { PyObject *plist; PyObject *result = NULL; @@ -832,8 +824,7 @@ Consumer_position(Handle *self, PyObject *args, rd_kafka_resp_err_t err; static char *kws[] = {"partitions", NULL}; - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O", kws, - &plist)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O", kws, &plist)) { return NULL; } @@ -876,8 +867,7 @@ Consumer_pause(Handle *self, PyObject *args, PyObject *kwargs) { rd_kafka_resp_err_t err; static char *kws[] = {"partitions", NULL}; - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O", kws, - &plist)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O", kws, &plist)) { return NULL; } @@ -917,8 +907,7 @@ Consumer_resume(Handle *self, PyObject *args, PyObject *kwargs) { rd_kafka_resp_err_t err; static char *kws[] = {"partitions", NULL}; - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O", kws, - &plist)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O", kws, &plist)) { return NULL; } @@ -1025,8 +1014,7 @@ Consumer_get_watermark_offsets(Handle *self, PyObject *args, PyObject *kwargs) { double tmout = -1.0f; int cached = 0; int64_t low = RD_KAFKA_OFFSET_INVALID, high = RD_KAFKA_OFFSET_INVALID; - static char *kws[] = {"partition", "timeout", "cached", - NULL}; + static char *kws[] = {"partition", "timeout", "cached", NULL}; if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|db", kws, (PyObject **)&tp, &tmout, &cached)) { @@ -1075,8 +1063,7 @@ Consumer_get_watermark_offsets(Handle *self, PyObject *args, PyObject *kwargs) { } static PyObject * -Consumer_offsets_for_times(Handle *self, PyObject *args, - PyObject *kwargs) { +Consumer_offsets_for_times(Handle *self, PyObject *args, PyObject *kwargs) { #if RD_KAFKA_VERSION < 0x000b0000 PyErr_Format(PyExc_NotImplementedError, "Consumer offsets_for_times require " @@ -1089,13 +1076,13 @@ Consumer_offsets_for_times(Handle *self, PyObject *args, PyObject *plist; PyObject *result = NULL; - double tmout = -1.0f; + double tmout = -1.0f; rd_kafka_topic_partition_list_t *c_parts; rd_kafka_resp_err_t err; static char *kws[] = {"partitions", "timeout", NULL}; - if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|d", kws, - &plist, &tmout)) { + if (!PyArg_ParseTupleAndKeywords(args, kwargs, "O|d", kws, &plist, + &tmout)) { return NULL; } @@ -1151,12 +1138,11 @@ Consumer_offsets_for_times(Handle *self, PyObject *args, * @return PyObject* Message object, None if timeout, or NULL on error * (raises KeyboardInterrupt if signal detected) */ -static PyObject * -Consumer_poll(Handle *self, PyObject *args, PyObject *kwargs) { +static PyObject *Consumer_poll(Handle *self, PyObject *args, PyObject *kwargs) { double tmout = -1.0f; static char *kws[] = {"timeout", NULL}; rd_kafka_message_t *rkm = NULL; - PyObject *result = NULL; + PyObject *result = NULL; CallState cs; const int CHUNK_TIMEOUT_MS = 200; /* 200ms chunks for signal checking */ int total_timeout_ms; @@ -1240,8 +1226,7 @@ Consumer_poll(Handle *self, PyObject *args, PyObject *kwargs) { return result; } -static PyObject * -Consumer_memberid(Handle *self, PyObject *ignore) { +static PyObject *Consumer_memberid(Handle *self, PyObject *ignore) { char *memberid; PyObject *result = NULL; @@ -1301,8 +1286,7 @@ static PyObject * Consumer_consume(Handle *self, PyObject *args, PyObject *kwargs) { unsigned int num_messages = 1; double tmout = -1.0f; - static char *kws[] = {"num_messages", "timeout", - NULL}; + static char *kws[] = {"num_messages", "timeout", NULL}; rd_kafka_message_t **rkmessages; PyObject *msglist; rd_kafka_queue_t *rkqu; @@ -1455,8 +1439,8 @@ static PyObject *Consumer_exit(Handle *self, PyObject *args) { return result; } -static PyObject * -Consumer_consumer_group_metadata(Handle *self, PyObject *ignore) { +static PyObject *Consumer_consumer_group_metadata(Handle *self, + PyObject *ignore) { rd_kafka_consumer_group_metadata_t *cgmd; PyObject *result = NULL; diff --git a/src/confluent_kafka/src/Metadata.c b/src/confluent_kafka/src/Metadata.c index 8cbc6ef40..c18c2eb0b 100644 --- a/src/confluent_kafka/src/Metadata.c +++ b/src/confluent_kafka/src/Metadata.c @@ -372,10 +372,10 @@ PyObject *list_topics(Handle *self, PyObject *args, PyObject *kwargs) { if (topic != NULL) { if (!(only_rkt = rd_kafka_topic_new(self->rk, topic, NULL))) { PyErr_Format(PyExc_RuntimeError, - "Unable to create topic object " - "for \"%s\": %s", - topic, - rd_kafka_err2str(rd_kafka_last_error())); + "Unable to create topic object " + "for \"%s\": %s", + topic, + rd_kafka_err2str(rd_kafka_last_error())); goto end; /* result and only_rkt are NULL */ } } @@ -609,8 +609,7 @@ PyObject *list_groups(Handle *self, PyObject *args, PyObject *kwargs) { const struct rd_kafka_group_list *group_list = NULL; const char *group = NULL; double tmout = -1.0f; - static char *kws[] = {"group", "timeout", - NULL}; + static char *kws[] = {"group", "timeout", NULL}; PyErr_WarnEx(PyExc_DeprecationWarning, "list_groups() is deprecated, use list_consumer_groups() " diff --git a/src/confluent_kafka/src/Producer.c b/src/confluent_kafka/src/Producer.c index 715f38aba..c286a549f 100644 --- a/src/confluent_kafka/src/Producer.c +++ b/src/confluent_kafka/src/Producer.c @@ -486,7 +486,7 @@ Producer_flush(Handle *self, PyObject *args, PyObject *kwargs) { const int CHUNK_TIMEOUT_MS = 200; /* 200ms chunks for signal checking */ int total_timeout_ms; int chunk_timeout_ms; - int chunk_count = 0; + int chunk_count = 0; PyObject *result = NULL; /* NULL means an exception is already set */ if (!PyArg_ParseTupleAndKeywords(args, kwargs, "|d", kws, &tmout)) @@ -539,7 +539,8 @@ Producer_flush(Handle *self, PyObject *args, PyObject *kwargs) { /* Always check for signals between chunks (critical for * interruptibility) */ if (check_signals_between_chunks(self, &cs)) { - goto exit; /* Signal detected, result stays NULL */ + goto exit; /* Signal detected, result stays NULL + */ } /* If timeout error, continue to next chunk */ @@ -587,7 +588,8 @@ Producer_close(Handle *self, PyObject *args, PyObject *kwargs) { * flushing and destroying it. */ if (!atomic_int_cas(&self->closing, 0, 1)) { - while (atomic_ptr_get(&self->rk) && atomic_int_get(&self->closing)) { + while (atomic_ptr_get(&self->rk) && + atomic_int_get(&self->closing)) { if (!Handle_sleep(self, 100)) return NULL; } @@ -609,7 +611,8 @@ Producer_close(Handle *self, PyObject *args, PyObject *kwargs) { /* Signal in-flight calls to stop, and wait for them to finish * using self->rk before destroying it -- see Handle_rk_use_begin(). - * New calls will see `closing` and fail with ERR_MSG_PRODUCER_CLOSED. */ + * New calls will see `closing` and fail with ERR_MSG_PRODUCER_CLOSED. + */ /* TODO NOGIL: replace this poll loop with a mutex/condvar wait so * close() unblocks immediately instead of up to 100ms late. */ while (atomic_int_get(&self->active_calls) > 0) { @@ -654,7 +657,8 @@ Producer_close(Handle *self, PyObject *args, PyObject *kwargs) { if (txn_errstr[0]) { PyErr_WarnFormat(PyExc_RuntimeWarning, 1, "Producer abort_transaction failed during " - "close: %s", txn_errstr); + "close: %s", + txn_errstr); } /* If flush failed, warn but don't suppress original exception */ @@ -987,8 +991,8 @@ Producer_produce_batch(Handle *self, PyObject *args, PyObject *kwargs) { static PyObject *Producer_init_transactions(Handle *self, PyObject *args) { CallState cs; rd_kafka_error_t *error; - double tmout = -1.0; - PyObject *result = NULL; /* NULL means an exception is already set */ + double tmout = -1.0; + PyObject *result = NULL; /* NULL means an exception is already set */ if (!PyArg_ParseTuple(args, "|d", &tmout)) return NULL; @@ -1045,7 +1049,7 @@ static PyObject *Producer_send_offsets_to_transaction(Handle *self, PyObject *metadata = NULL, *offsets = NULL; rd_kafka_topic_partition_list_t *c_offsets = NULL; rd_kafka_consumer_group_metadata_t *cgmd = NULL; - double tmout = -1.0; + double tmout = -1.0; PyObject *result = NULL; /* NULL means an exception is already set */ if (!PyArg_ParseTuple(args, "OO|d", &offsets, &metadata, &tmout)) @@ -1092,8 +1096,8 @@ static PyObject *Producer_send_offsets_to_transaction(Handle *self, static PyObject *Producer_commit_transaction(Handle *self, PyObject *args) { CallState cs; rd_kafka_error_t *error; - double tmout = -1.0; - PyObject *result = NULL; /* NULL means an exception is already set */ + double tmout = -1.0; + PyObject *result = NULL; /* NULL means an exception is already set */ if (!PyArg_ParseTuple(args, "|d", &tmout)) return NULL; @@ -1128,8 +1132,8 @@ static PyObject *Producer_commit_transaction(Handle *self, PyObject *args) { static PyObject *Producer_abort_transaction(Handle *self, PyObject *args) { CallState cs; rd_kafka_error_t *error; - double tmout = -1.0; - PyObject *result = NULL; /* NULL means an exception is already set */ + double tmout = -1.0; + PyObject *result = NULL; /* NULL means an exception is already set */ if (!PyArg_ParseTuple(args, "|d", &tmout)) return NULL; diff --git a/src/confluent_kafka/src/ShareConsumer.c b/src/confluent_kafka/src/ShareConsumer.c index d6761af43..607be85bc 100644 --- a/src/confluent_kafka/src/ShareConsumer.c +++ b/src/confluent_kafka/src/ShareConsumer.c @@ -1099,7 +1099,8 @@ static PyMethodDef ShareConsumer_methods[] = { " broker responds to an acknowledgement commit. It is always dispatched\n" " on the application thread, from within whichever consumer call is\n" " serving the response queue (:py:func:`poll`, :py:func:`commit_sync`,\n" - " :py:func:`commit_async`, or :py:func:`close`), never from a background\n" + " :py:func:`commit_async`, or :py:func:`close`), never from a " + "background\n" " thread.\n" "\n" " :param callback: A callable\n" diff --git a/src/confluent_kafka/src/confluent_kafka.c b/src/confluent_kafka/src/confluent_kafka.c index d73c3762d..8a337fb70 100644 --- a/src/confluent_kafka/src/confluent_kafka.c +++ b/src/confluent_kafka/src/confluent_kafka.c @@ -101,7 +101,8 @@ static PyObject *KafkaError_str(KafkaError *self, PyObject *ignore) { if (self->str) return cfl_PyUnistr_FromStringSafe(self->str); else - return cfl_PyUnistr_FromStringSafe(rd_kafka_err2str(self->code)); + return cfl_PyUnistr_FromStringSafe( + rd_kafka_err2str(self->code)); } static PyObject *KafkaError_name(KafkaError *self, PyObject *ignore) { @@ -485,8 +486,7 @@ static void cfl_PyErr_Fatal(rd_kafka_resp_err_t err, const char *reason) { * (free-threaded builds) cannot drop the last reference between our * read and our INCREF. No-op on GIL builds. */ -static PyObject * -Message_get_field(Message *self, PyObject **field) { +static PyObject *Message_get_field(Message *self, PyObject **field) { PyObject *obj; Py_BEGIN_CRITICAL_SECTION(self); @@ -2992,12 +2992,12 @@ static void common_conf_set_software(rd_kafka_conf_t *conf) { static int resolve_aws_oauthbearer_marker(PyObject *confdict) { static const char MARKER_KEY[] = "sasl.oauthbearer.metadata.authentication.type"; - static const char MARKER_VALUE[] = "aws_iam"; - static const char METHOD_KEY[] = "sasl.oauthbearer.method"; + static const char MARKER_VALUE[] = "aws_iam"; + static const char METHOD_KEY[] = "sasl.oauthbearer.method"; static const char METHOD_OIDC_VALUE[] = "oidc"; - static const char CONFIG_KEY[] = "sasl.oauthbearer.config"; - static const char EXTENSIONS_KEY[] = "sasl.oauthbearer.extensions"; - static const char OAUTH_CB_KEY[] = "oauth_cb"; + static const char CONFIG_KEY[] = "sasl.oauthbearer.config"; + static const char EXTENSIONS_KEY[] = "sasl.oauthbearer.extensions"; + static const char OAUTH_CB_KEY[] = "oauth_cb"; static const char AUTOWIRE_MODULE[] = "confluent_kafka._oauthbearer.aws.aws_autowire"; static const char CREATE_HANDLER[] = "create_handler"; @@ -3027,7 +3027,8 @@ static int resolve_aws_oauthbearer_marker(PyObject *confdict) { const char *marker_c; const char *method_c; - /* Explicit oauth_cb wins: nothing to autowire, regardless of the marker. */ + /* Explicit oauth_cb wins: nothing to autowire, regardless of the + * marker. */ cb = PyDict_GetItemString(confdict, OAUTH_CB_KEY); if (cb && cb != Py_None) { return 0; @@ -3047,7 +3048,7 @@ static int resolve_aws_oauthbearer_marker(PyObject *confdict) { return 0; } - method = PyDict_GetItemString(confdict, METHOD_KEY); + method = PyDict_GetItemString(confdict, METHOD_KEY); method_c = (method && PyUnicode_Check(method)) ? PyUnicode_AsUTF8(method) : NULL; @@ -3094,8 +3095,8 @@ static int resolve_aws_oauthbearer_marker(PyObject *confdict) { if (!func) { return -1; } - callback = PyObject_CallFunction( - func, "OO", cfg_str, ext_str ? ext_str : Py_None); + callback = PyObject_CallFunction(func, "OO", cfg_str, + ext_str ? ext_str : Py_None); Py_DECREF(func); if (!callback) { return -1; @@ -3634,8 +3635,7 @@ int Handle_serialize_enter(Handle *h) { unsigned long identity = 0; PyObject *value = NULL; - if (PyContextVar_Get(Consumer_reentry_identity_var, NULL, &value) == - -1) + if (PyContextVar_Get(Consumer_reentry_identity_var, NULL, &value) == -1) return 0; if (value && PyLong_Check(value)) @@ -3661,8 +3661,7 @@ int Handle_serialize_enter(Handle *h) { } if (owner == 0 && - atomic_ulong_cas(&h->u.Consumer.gate_owner, 0, - identity)) { + atomic_ulong_cas(&h->u.Consumer.gate_owner, 0, identity)) { /* Gate looked unowned and we won the race to take * it. */ atomic_int_set(&h->u.Consumer.gate_depth, 1); diff --git a/src/confluent_kafka/src/confluent_kafka.h b/src/confluent_kafka/src/confluent_kafka.h index 52d2d69da..fea6876e3 100644 --- a/src/confluent_kafka/src/confluent_kafka.h +++ b/src/confluent_kafka/src/confluent_kafka.h @@ -48,26 +48,26 @@ #if defined(_MSC_VER) typedef volatile LONG atomic_int_t; -#define atomic_int_inc(p) InterlockedIncrement((p)) -#define atomic_int_dec(p) InterlockedDecrement((p)) -#define atomic_int_get(p) InterlockedCompareExchange((p), 0, 0) +#define atomic_int_inc(p) InterlockedIncrement((p)) +#define atomic_int_dec(p) InterlockedDecrement((p)) +#define atomic_int_get(p) InterlockedCompareExchange((p), 0, 0) #define atomic_int_set(p, v) InterlockedExchange((p), (v)) /** * @brief Atomic compare-and-swap: if *p == expected, set *p = desired and * return 1; otherwise leave *p unchanged and return 0. */ -static __inline int atomic_int_cas(atomic_int_t *p, LONG expected, - LONG desired) { +static __inline int +atomic_int_cas(atomic_int_t *p, LONG expected, LONG desired) { return InterlockedCompareExchange(p, desired, expected) == expected; } #else /* gcc / clang */ typedef int atomic_int_t; -#define atomic_int_inc(p) __atomic_add_fetch((p), 1, __ATOMIC_SEQ_CST) -#define atomic_int_dec(p) __atomic_sub_fetch((p), 1, __ATOMIC_SEQ_CST) -#define atomic_int_get(p) __atomic_load_n((p), __ATOMIC_SEQ_CST) +#define atomic_int_inc(p) __atomic_add_fetch((p), 1, __ATOMIC_SEQ_CST) +#define atomic_int_dec(p) __atomic_sub_fetch((p), 1, __ATOMIC_SEQ_CST) +#define atomic_int_get(p) __atomic_load_n((p), __ATOMIC_SEQ_CST) #define atomic_int_set(p, v) __atomic_store_n((p), (v), __ATOMIC_SEQ_CST) /** @@ -90,13 +90,14 @@ static inline int atomic_int_cas(atomic_int_t *p, int expected, int desired) { #if defined(_MSC_VER) typedef volatile LONG_PTR atomic_ulong_t; -#define atomic_ulong_get(p) \ - ((unsigned long)InterlockedCompareExchangePointer( \ +#define atomic_ulong_get(p) \ + ((unsigned long)InterlockedCompareExchangePointer( \ (PVOID volatile *)(p), 0, 0)) -#define atomic_ulong_set(p, v) \ +#define atomic_ulong_set(p, v) \ InterlockedExchangePointer((PVOID volatile *)(p), (PVOID)(v)) -static __inline int atomic_ulong_cas(atomic_ulong_t *p, unsigned long expected, +static __inline int atomic_ulong_cas(atomic_ulong_t *p, + unsigned long expected, unsigned long desired) { return InterlockedCompareExchangePointer( (PVOID volatile *)p, (PVOID)desired, (PVOID)expected) == @@ -106,10 +107,11 @@ static __inline int atomic_ulong_cas(atomic_ulong_t *p, unsigned long expected, #else /* gcc / clang */ typedef unsigned long atomic_ulong_t; -#define atomic_ulong_get(p) __atomic_load_n((p), __ATOMIC_SEQ_CST) +#define atomic_ulong_get(p) __atomic_load_n((p), __ATOMIC_SEQ_CST) #define atomic_ulong_set(p, v) __atomic_store_n((p), (v), __ATOMIC_SEQ_CST) -static inline int atomic_ulong_cas(atomic_ulong_t *p, unsigned long expected, +static inline int atomic_ulong_cas(atomic_ulong_t *p, + unsigned long expected, unsigned long desired) { return __atomic_compare_exchange_n(p, &expected, desired, 0 /* strong */, __ATOMIC_SEQ_CST, @@ -121,12 +123,12 @@ static inline int atomic_ulong_cas(atomic_ulong_t *p, unsigned long expected, * @brief Atomic accessors for Handle.rk itself. */ #if defined(_MSC_VER) -#define atomic_ptr_get(p) \ +#define atomic_ptr_get(p) \ InterlockedCompareExchangePointer((PVOID volatile *)(p), NULL, NULL) -#define atomic_ptr_set(p, v) \ +#define atomic_ptr_set(p, v) \ InterlockedExchangePointer((PVOID volatile *)(p), (PVOID)(v)) #else /* gcc / clang */ -#define atomic_ptr_get(p) __atomic_load_n((p), __ATOMIC_SEQ_CST) +#define atomic_ptr_get(p) __atomic_load_n((p), __ATOMIC_SEQ_CST) #define atomic_ptr_set(p, v) __atomic_store_n((p), (v), __ATOMIC_SEQ_CST) #endif @@ -180,7 +182,7 @@ static inline int atomic_ulong_cas(atomic_ulong_t *p, unsigned long expected, * no version guards. */ #ifndef Py_BEGIN_CRITICAL_SECTION #define Py_BEGIN_CRITICAL_SECTION(op) { -#define Py_END_CRITICAL_SECTION() } +#define Py_END_CRITICAL_SECTION() } #endif /** diff --git a/tests/schema_registry/_async/test_avro_serdes.py b/tests/schema_registry/_async/test_avro_serdes.py index f9f707c31..93e3264fd 100644 --- a/tests/schema_registry/_async/test_avro_serdes.py +++ b/tests/schema_registry/_async/test_avro_serdes.py @@ -16,11 +16,17 @@ # limitations under the License. # import json +import time +from datetime import datetime, timedelta, timezone +from decimal import Decimal import pytest +from fastavro._logical_readers import UUID from confluent_kafka.schema_registry import ( AsyncSchemaRegistryClient, + Metadata, + MetadataProperties, Schema, header_schema_id_serializer, ) @@ -34,9 +40,41 @@ AssociationCreateOrUpdateRequest, ) from confluent_kafka.schema_registry.common.serde import SubjectNameStrategyType -from confluent_kafka.schema_registry.schema_registry_client import SchemaReference +from confluent_kafka.schema_registry.rule_registry import RuleOverride, RuleRegistry +from confluent_kafka.schema_registry.rules.cel.cel_executor import CelExecutor +from confluent_kafka.schema_registry.rules.cel.cel_field_executor import CelFieldExecutor +from confluent_kafka.schema_registry.rules.encryption.dek_registry.dek_registry_client import ( + DekAlgorithm, + DekRegistryClient, +) +from confluent_kafka.schema_registry.rules.encryption.encrypt_executor import ( + Clock, + EncryptionExecutor, + FieldEncryptionExecutor, +) +from confluent_kafka.schema_registry.rules.jsonata.jsonata_executor import JsonataExecutor +from confluent_kafka.schema_registry.schema_registry_client import ( + Rule, + RuleKind, + RuleMode, + RuleParams, + RuleSet, + SchemaReference, + ServerConfig, +) +from confluent_kafka.schema_registry.serde import RuleConditionError from confluent_kafka.serialization import MessageField, SerializationContext, SerializationError + +class FakeClock(Clock): + + def __init__(self): + self.fixed_now = int(round(time.time() * 1000)) + + def now(self) -> int: + return self.fixed_now + + _BASE_URL = "mock://" # _BASE_URL = "http://localhost:8081" _TOPIC = "topic1" @@ -46,6 +84,12 @@ @pytest.fixture(autouse=True) async def run_before_and_after_tests(tmpdir): """Fixture to execute asserts before and after a test is run""" + # Setup: fill with any logic you want + + CelExecutor.register() + CelFieldExecutor.register() + JsonataExecutor.register() + yield # this is where the testing happens # Teardown : fill with any logic you want @@ -524,6 +568,2268 @@ async def test_avro_schema_evolution(): assert obj2.get('newOptionalField') == 'optional' +async def test_avro_cel_condition(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField == 'hi'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_condition_logical_type(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': {'type': 'string', 'logicalType': 'uuid'}}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + uuid = "550e8400-e29b-41d4-a716-446655440000" + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField == '" + uuid + "'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': UUID(uuid), + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_condition_fail(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField != 'hi'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_avro_cel_condition_ignore_fail(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField != 'hi'", + None, + "NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_field_transform(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi-suffix', + 'booleanField': True, + 'bytesField': b'foobar', + } + deser = await AsyncAvroDeserializer(client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_avro_cel_field_transform_missing_prop(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + {'name': 'missing', 'type': ['null', 'string'], 'default': None}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi-suffix-suffix', + 'booleanField': True, + 'bytesField': b'foobar', + 'missing': None, + } + deser = await AsyncAvroDeserializer(client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_avro_cel_field_transform_disable(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + + registry = RuleRegistry() + registry.register_rule_executor(CelFieldExecutor()) + registry.register_override(RuleOverride("CEL_FIELD", None, None, True)) + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_registry=registry) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + newobj = await deser(obj_bytes, ser_ctx) + assert "hi" == newobj['stringField'] + + +async def test_avro_cel_field_transform_complex(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'arrayField', 'type': {'type': 'array', 'items': 'string'}}, + {'name': 'mapField', 'type': {'type': 'map', 'values': 'string'}}, + {'name': 'unionField', 'type': ['null', 'string'], 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'arrayField': ['hello'], + 'mapField': {'key': 'world'}, + 'unionField': 'bye', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = { + 'arrayField': ['hello-suffix'], + 'mapField': {'key': 'world-suffix'}, + 'unionField': 'bye-suffix', + } + deser = await AsyncAvroDeserializer(client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_avro_cel_field_transform_complex_with_none(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'arrayField', 'type': {'type': 'array', 'items': 'string'}}, + {'name': 'mapField', 'type': {'type': 'map', 'values': 'string'}}, + {'name': 'unionField', 'type': ['null', 'string'], 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'arrayField': ['hello'], + 'mapField': {'key': 'world'}, + 'unionField': None, + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = { + 'arrayField': ['hello-suffix'], + 'mapField': {'key': 'world-suffix'}, + 'unionField': None, + } + deser = await AsyncAvroDeserializer(client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_avro_cel_field_transform_complex_nested(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'UnionTest', + 'namespace': 'test', + 'fields': [ + { + 'name': 'emails', + 'type': [ + 'null', + { + 'type': 'array', + 'items': { + 'type': 'record', + 'name': 'Email', + 'fields': [ + { + 'name': 'email', + 'type': ['null', 'string'], + 'doc': 'Email address', + 'confluent:tags': ['PII'], + } + ], + }, + }, + ], + 'doc': 'Communication Email', + } + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'emails': [{'email': 'john@acme.com'}]} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = {'emails': [{'email': 'john@acme.com-suffix'}]} + deser = await AsyncAvroDeserializer(client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_avro_cel_field_condition(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value == 'hi'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_field_condition_fail(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value == 'bye'", + None, + None, + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +_AVRO_DECIMAL_SCHEMA = { + 'type': 'record', + 'name': 'test', + 'fields': [ + { + 'name': 'decField', + 'type': { + 'type': 'bytes', + 'logicalType': 'decimal', + 'precision': 10, + 'scale': 2, + }, + }, + ], +} + + +async def test_avro_cel_decimal_passes(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.gt(decimal(message.decField), decimal("10.00"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_decimal_fails(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.lt(decimal(message.decField), decimal("10.00"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_avro_cel_decimal_needs_no_constructor(): + """Cross-client parity: an Avro ``decimal`` logical type is usable as a Decimal with **no + ``decimal(...)`` call**, and the wrapped form keeps working alongside it. fastavro decodes it + to a Python ``Decimal`` at the schema's scale, which is this client's in-CEL decimal + representation, so ``decimals.*`` accept it directly and ``==`` is numeric (Python's own + ``Decimal.__eq__``). + """ + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + + async def serialize(expr): + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + expr, + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + return await ser({'decField': Decimal('12.34')}, ser_ctx) + + # Bare: no constructor call on the field. + assert await serialize('decimals.eq(message.decField, decimal("12.34"))') is not None + assert await serialize('decimals.gt(message.decField, decimal("10.00"))') is not None + # The wrapped form must keep working (decimal(...) re-entry). + assert await serialize('decimals.eq(decimal(message.decField), decimal("12.34"))') is not None + # `==` is numeric on it: 12.34 equals 12.340 despite the differing scale. + assert await serialize('message.decField == decimal("12.340")') is not None + # The schema's scale is applied, not guessed: as scale 0 this would be 1234. + assert await serialize('decimals.lt(message.decField, decimal("100"))') is not None + # Negative control: a false comparison must fail. + with pytest.raises(SerializationError) as e: + await serialize('decimals.gt(message.decField, decimal("100"))') + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_avro_cel_decimal_arithmetic(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.add(decimal(message.decField), decimal("1.66")), decimal("14.00"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_decimal_string(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'string(decimal(message.decField)) == "12.34"', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +_AVRO_TIMESTAMP_SCHEMA = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'tsField', 'type': {'type': 'long', 'logicalType': 'timestamp-millis'}}, + ], +} + + +async def test_avro_cel_timestamp_millis_passes(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.tsField) < now', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_TIMESTAMP_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'tsField': datetime(2020, 1, 1, tzinfo=timezone.utc)} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_cel_timestamp_millis_fails(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.tsField) > now', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_TIMESTAMP_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'tsField': datetime(2020, 1, 1, tzinfo=timezone.utc)} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_avro_cel_timestamp_millis_needs_no_constructor(): + """Cross-client parity: an Avro timestamp logical type is usable as a timestamp with **no + constructor call at all**. fastavro decodes it to an aware datetime, which the boundary + binds as a CEL timestamp, so it is comparable against ``now`` and carries the timestamp + accessors. Every one of the seven clients has this test; the constructor is only needed for + a plain numeric field whose unit the schema cannot supply. + """ + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + + async def serialize(expr, value): + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + expr, + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_TIMESTAMP_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + return await ser({'tsField': value}, ser_ctx) + + past = datetime(2020, 1, 1, tzinfo=timezone.utc) + exact = datetime(2023, 11, 14, 22, 13, 20, 123000, tzinfo=timezone.utc) + + # Bare comparison against `now`. + assert await serialize('message.tsField < now', past) is not None + # The schema's millis unit is applied, not guessed, and the accessors work directly. + assert await serialize('message.tsField == timestamp("2023-11-14T22:13:20.123Z")', exact) is not None + assert await serialize('message.tsField.getFullYear() == 2023', exact) is not None + # Negative control: a future value must fail, so the comparison really happens. + with pytest.raises(SerializationError) as e: + await serialize('message.tsField < now', datetime(2100, 1, 1, tzinfo=timezone.utc)) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_avro_encryption(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + obj['bytesField'] = b'foobar' + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_complex_schema(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + { + 'name': 'complexField1', + 'type': { + 'fields': [ + {'name': 'stringValue', 'type': 'string', 'confluent:tags': ['PII']}, + ], + 'name': 'ComplexFieldType', + 'type': 'record', + }, + }, + {'name': 'complexField2', 'type': 'ComplexFieldType'}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'complexField1': {'stringValue': 'test1'}, + 'complexField2': {'stringValue': 'test2'}, + } + + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + assert 'test1' not in str(obj_bytes) + assert 'test2' not in str(obj_bytes) + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + actual = await deser(obj_bytes, ser_ctx) + assert actual['complexField1']['stringValue'] == 'test1' + assert actual['complexField2']['stringValue'] == 'test2' + + +async def test_avro_encryption_complex_schema_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + { + 'name': 'complexField1', + 'type': { + 'fields': [ + { + 'name': 'complexSubType1', + 'type': { + 'fields': [{'name': 'stringValue', 'type': 'string', 'confluent:tags': ['PII']}], + 'name': 'ComplexSubType', + 'type': 'record', + }, + }, + {'name': 'complexSubType2', 'type': 'ComplexSubType'}, + ], + 'name': 'ComplexFieldType', + 'type': 'record', + }, + }, + {'name': 'complexField2', 'type': ['null', 'ComplexFieldType']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'complexField1': {'complexSubType1': {'stringValue': 'test1'}, 'complexSubType2': {'stringValue': 'test2'}}, + 'complexField2': {'complexSubType1': {'stringValue': 'test3'}, 'complexSubType2': {'stringValue': 'test4'}}, + } + + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + assert 'test1' not in str(obj_bytes) + assert 'test2' not in str(obj_bytes) + assert 'test3' not in str(obj_bytes) + assert 'test4' not in str(obj_bytes) + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + actual = await deser(obj_bytes, ser_ctx) + assert actual['complexField1']['complexSubType1']['stringValue'] == 'test1' + assert actual['complexField1']['complexSubType2']['stringValue'] == 'test2' + assert actual['complexField2']['complexSubType1']['stringValue'] == 'test3' + assert actual['complexField2']['complexSubType2']['stringValue'] == 'test4' + + +async def test_avro_payload_encryption(): + executor = EncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT_PAYLOAD", + None, + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_alternate_keks(): + executor = EncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret', 'encrypt.alternate.kms.key.ids': 'mykey2,mykey3'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT_PAYLOAD", + None, + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_deterministic(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.algorithm": "AES256_SIV", + } + ), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + obj['bytesField'] = b'foobar' + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_wrapped_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + "fields": [ + {"name": "id", "type": "int"}, + { + "name": "result", + "type": [ + "null", + { + "fields": [ + {"name": "code", "type": "int"}, + {"confluent:tags": ["PII"], "name": "secret", "type": ["null", "string"]}, + ], + "name": "Data", + "type": "record", + }, + { + "fields": [{"name": "code", "type": "int"}, {"name": "reason", "type": ["null", "string"]}], + "name": "Error", + "type": "record", + }, + ], + }, + ], + "name": "Result", + "namespace": "com.acme", + "type": "record", + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'id': 123, 'result': ('com.acme.Data', {'code': 456, 'secret': 'mypii'})} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['result'][1]['secret'] != 'mypii' + # remove union wrapper + obj['result'] = {'code': 456, 'secret': 'mypii'} + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_typed_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + "fields": [ + {"name": "id", "type": "int"}, + { + "name": "result", + "type": [ + "null", + { + "fields": [ + {"name": "code", "type": "int"}, + {"confluent:tags": ["PII"], "name": "secret", "type": ["null", "string"]}, + ], + "name": "Data", + "type": "record", + }, + { + "fields": [{"name": "code", "type": "int"}, {"name": "reason", "type": ["null", "string"]}], + "name": "Error", + "type": "record", + }, + ], + }, + ], + "name": "Result", + "namespace": "com.acme", + "type": "record", + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'id': 123, 'result': {'-type': 'com.acme.Data', 'code': 456, 'secret': 'mypii'}} + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['result']['secret'] != 'mypii' + # remove union wrapper + obj['result'] = {'code': 456, 'secret': 'mypii'} + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_cel(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule1 = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + rule2 = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule1, rule2]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi-suffix' + obj['stringField'] = 'hi-suffix' + obj['bytesField'] = b'foobar' + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_dek_rotation(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1-rot", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.expiry.days": "1", + } + ), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client: DekRegistryClient = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + dek_client = executor.executor.client + dek = dek_client.get_dek("kek1-rot", _SUBJECT, version=-1) + assert dek.version == 1 + + # advance 2 days + now = datetime.now() + timedelta(days=2) + executor.executor.clock.fixed_now = int(round(now.timestamp() * 1000)) + + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + dek = dek_client.get_dek("kek1-rot", _SUBJECT, version=-1) + assert dek.version == 2 + + # advance 2 days + now = datetime.now() + timedelta(days=2) + executor.executor.clock.fixed_now = int(round(now.timestamp() * 1000)) + + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + dek = dek_client.get_dek("kek1-rot", _SUBJECT, version=-1) + assert dek.version == 3 + + +async def test_avro_encryption_f1_preserialized(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'f1Schema', + 'fields': [{'name': 'f1', 'type': 'string', 'confluent:tags': ['PII']}], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1-f1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,ERROR", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'f1': 'hello world'} + + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + + dek_client: DekRegistryClient = executor.executor.client + dek_client.register_kek("kek1-f1", "local-kms", "mykey") + + encrypted_dek = "07V2ndh02DA73p+dTybwZFm7DKQSZN1tEwQh+FoX1DZLk4Yj2LLu4omYjp/84tAg3BYlkfGSz+zZacJHIE4=" + dek_client.register_dek("kek1-f1", _SUBJECT, encrypted_dek) + + obj_bytes = bytes( + [ + 0, + 0, + 0, + 0, + 1, + 104, + 122, + 103, + 121, + 47, + 106, + 70, + 78, + 77, + 86, + 47, + 101, + 70, + 105, + 108, + 97, + 72, + 114, + 77, + 121, + 101, + 66, + 103, + 100, + 97, + 86, + 122, + 114, + 82, + 48, + 117, + 100, + 71, + 101, + 111, + 116, + 87, + 56, + 99, + 65, + 47, + 74, + 97, + 108, + 55, + 117, + 107, + 114, + 43, + 77, + 47, + 121, + 122, + ] + ) + + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_deterministic_f1_preserialized(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'f1Schema', + 'fields': [{'name': 'f1', 'type': 'string', 'confluent:tags': ['PII']}], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1-det-f1", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.algorithm": "AES256_SIV", + } + ), + None, + None, + "ERROR,ERROR", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'f1': 'hello world'} + + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + + dek_client: DekRegistryClient = executor.executor.client + dek_client.register_kek("kek1-det-f1", "local-kms", "mykey") + + encrypted_dek = ( + "YSx3DTlAHrmpoDChquJMifmPntBzxgRVdMzgYL82rgWBKn7aUSnG+WIu9oz" + "BNS3y2vXd++mBtK07w4/W/G6w0da39X9hfOVZsGnkSvry/QRht84V8yz3dqKxGMOK5A==" + ) + dek_client.register_dek("kek1-det-f1", _SUBJECT, encrypted_dek, algorithm=DekAlgorithm.AES256_SIV) + + obj_bytes = bytes( + [ + 0, + 0, + 0, + 0, + 1, + 72, + 68, + 54, + 89, + 116, + 120, + 114, + 108, + 66, + 110, + 107, + 84, + 87, + 87, + 57, + 78, + 54, + 86, + 98, + 107, + 51, + 73, + 73, + 110, + 106, + 87, + 72, + 56, + 49, + 120, + 109, + 89, + 104, + 51, + 107, + 52, + 100, + ] + ) + + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_dek_rotation_f1_preserialized(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'f1Schema', + 'fields': [{'name': 'f1', 'type': 'string', 'confluent:tags': ['PII']}], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1-rot-f1", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.expiry.days": "1", + } + ), + None, + None, + "ERROR,ERROR", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'f1': 'hello world'} + + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + + dek_client: DekRegistryClient = executor.executor.client + dek_client.register_kek("kek1-rot-f1", "local-kms", "mykey") + + encrypted_dek = "W/v6hOQYq1idVAcs1pPWz9UUONMVZW4IrglTnG88TsWjeCjxmtRQ4VaNe/I5dCfm2zyY9Cu0nqdvqImtUk4=" + dek_client.register_dek("kek1-rot-f1", _SUBJECT, encrypted_dek, algorithm=DekAlgorithm.AES256_GCM) + + obj_bytes = bytes( + [ + 0, + 0, + 0, + 0, + 1, + 120, + 65, + 65, + 65, + 65, + 65, + 65, + 71, + 52, + 72, + 73, + 54, + 98, + 49, + 110, + 88, + 80, + 88, + 113, + 76, + 121, + 71, + 56, + 99, + 73, + 73, + 51, + 53, + 78, + 72, + 81, + 115, + 101, + 113, + 113, + 85, + 67, + 100, + 43, + 73, + 101, + 76, + 101, + 70, + 86, + 65, + 101, + 78, + 112, + 83, + 83, + 51, + 102, + 120, + 80, + 110, + 74, + 51, + 50, + 65, + 61, + ] + ) + + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_references(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + + referenced = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + obj = {'refField': referenced} + ref_schema = { + 'type': 'record', + 'name': 'ref', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + await client.register_schema('ref', Schema(json.dumps(ref_schema))) + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'refField', 'type': 'ref'}, + ], + } + refs = [SchemaReference('ref', 'ref', 1)] + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1-ref", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", refs, None, RuleSet(None, [rule]))) + + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['refField']['stringField'] != 'hi' + obj['refField']['stringField'] = 'hi' + obj['refField']['bytesField'] = b'foobar' + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_encryption_with_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': ['null', 'string'], 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': ['null', 'bytes'], 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1-union", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + obj['bytesField'] = b'foobar' + + deser = await AsyncAvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_avro_jsonata_with_cel(): + rule1_to_2 = "$merge([$sift($, function($v, $k) {$k != 'size'}), {'height': $.'size'}])" + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + + await client.set_config(_SUBJECT, ServerConfig(compatibility_group='application.version')) + + schema = { + 'type': 'record', + 'name': 'old', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'size', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + await client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v1"}), None), + None, + ), + ) + + schema = { + 'type': 'record', + 'name': 'new', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'height', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + + rule1 = Rule( + "test-jsonata", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule1_to_2, None, None, False + ) + rule2 = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.READ, + "CEL_FIELD", + None, + None, + "name == 'name' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v2"}), None), + RuleSet([rule1], [rule2]), + ), + ) + + obj = { + 'name': 'alice', + 'size': 123, + 'version': 1, + } + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v1'}, + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = { + 'name': 'alice-suffix', + 'height': 123, + 'version': 1, + } + deser_conf = {'use.latest.with.metadata': {'application.version': 'v2'}} + deser = await AsyncAvroDeserializer(client, conf=deser_conf) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_avro_jsonata_fully_compatible(): + rule1_to_2 = "$merge([$sift($, function($v, $k) {$k != 'size'}), {'height': $.'size'}])" + rule2_to_1 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'size': $.'height'}])" + rule2_to_3 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'length': $.'height'}])" + rule3_to_2 = "$merge([$sift($, function($v, $k) {$k != 'length'}), {'height': $.'length'}])" + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + + await client.set_config(_SUBJECT, ServerConfig(compatibility_group='application.version')) + + schema = { + 'type': 'record', + 'name': 'old', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'size', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + await client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v1"}), None), + None, + ), + ) + + schema = { + 'type': 'record', + 'name': 'new', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'height', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + + rule1 = Rule( + "rule1", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule1_to_2, None, None, False + ) + rule2 = Rule( + "rule2", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule2_to_1, None, None, False + ) + await client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v2"}), None), + RuleSet([rule1, rule2], None), + ), + ) + + schema = { + 'type': 'record', + 'name': 'newer', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'length', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + + rule3 = Rule( + "rule3", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule2_to_3, None, None, False + ) + rule4 = Rule( + "rule4", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule3_to_2, None, None, False + ) + await client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v3"}), None), + RuleSet([rule3, rule4], None), + ), + ) + + obj = { + 'name': 'alice', + 'size': 123, + 'version': 1, + } + obj2 = { + 'name': 'alice', + 'height': 123, + 'version': 1, + } + obj3 = { + 'name': 'alice', + 'length': 123, + 'version': 1, + } + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v1'}, + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + await deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v2'}, + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj2, ser_ctx) + + await deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v3'}, + } + ser = await AsyncAvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj3, ser_ctx) + + await deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + +async def deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3): + deser_conf = {'use.latest.with.metadata': {'application.version': 'v1'}} + deser = await AsyncAvroDeserializer(client, conf=deser_conf) + newobj = await deser(obj_bytes, ser_ctx) + assert obj == newobj + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v2'}} + deser = await AsyncAvroDeserializer(client, conf=deser_conf) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v3'}} + deser = await AsyncAvroDeserializer(client, conf=deser_conf) + newobj = await deser(obj_bytes, ser_ctx) + assert obj3 == newobj + + async def test_avro_reference(): conf = {'url': _BASE_URL} client = AsyncSchemaRegistryClient.new_client(conf) diff --git a/tests/schema_registry/_async/test_cel_avro_message_transform.py b/tests/schema_registry/_async/test_cel_avro_message_transform.py new file mode 100644 index 000000000..cc1717b0a --- /dev/null +++ b/tests/schema_registry/_async/test_cel_avro_message_transform.py @@ -0,0 +1,175 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +Message-level ``CEL`` transforms over Avro, and specifically that they **replace** rather than +merge: the rule's map is the whole new record, so a field the rule does not name takes the +schema's declared default rather than the value it had on the way in. + +This case existed only on the protobuf side, and its absence hid a real defect elsewhere - the +C++ client seeded its result record from the input before applying the map, so it merged. Every +other C6/C7 case names *all* of the record's fields, which makes merge and replace +indistinguishable, the same way a condition fixture cannot tell a passing rule from an unfired +one without a must-fail twin. + +Driven end to end through the serializer rather than through the executor alone, which is not +incidental: this client's Avro write-back hands fastavro the rule's result more or less +unchanged, so whether fastavro can fill an omitted field is the whole question. An +executor-level test cannot see it. +""" + +import json + +from confluent_kafka.schema_registry import Schema +from confluent_kafka.schema_registry._async.schema_registry_client import AsyncSchemaRegistryClient +from confluent_kafka.schema_registry.avro import AsyncAvroDeserializer, AsyncAvroSerializer +from confluent_kafka.schema_registry.rules.cel.cel_executor import CelExecutor +from confluent_kafka.schema_registry.rules.cel.cel_field_executor import CelFieldExecutor +from confluent_kafka.schema_registry.schema_registry_client import Rule, RuleKind, RuleMode, RuleSet +from confluent_kafka.serialization import MessageField, SerializationContext + +CelExecutor.register() +CelFieldExecutor.register() + +_TOPIC = "cel-avro-message-transform" + +_SCHEMA = { + "type": "record", + "name": "Defaults", + "fields": [ + {"name": "kept", "type": "string"}, + {"name": "withDefault", "type": "string", "default": "fallback"}, + {"name": "nullable", "type": ["null", "string"], "default": None}, + ], +} + +_RECORD = { + "kept": "original-kept", + "withDefault": "original-withDefault", + "nullable": "original-nullable", +} + + +async def _round_trip(subject_suffix, expr): + """Serializes the fixture under one message-level CEL transform and reads it back.""" + topic = _TOPIC + "-" + subject_suffix + client = AsyncSchemaRegistryClient.new_client({"url": "mock://"}) + rule = Rule("r", "", RuleKind.TRANSFORM, RuleMode.WRITE, "CEL", None, None, expr, None, None, False) + await client.register_schema(topic + "-value", Schema(json.dumps(_SCHEMA), "AVRO", [], None, RuleSet(None, [rule]))) + ser = await AsyncAvroSerializer( + client, schema_str=None, conf={"auto.register.schemas": False, "use.latest.version": True} + ) + ctx = SerializationContext(topic, MessageField.VALUE) + # Each await is on its own call: tools/unasync.py strips "await " by word boundary, so + # "await (" survives the rewrite and the generated sync file will not parse. + payload = await ser(_RECORD, ctx) + deser = await AsyncAvroDeserializer(client) + return await deser(payload, ctx) + + +async def test_a_field_the_rule_does_not_name_takes_its_declared_default(): + """The case this file exists for. Under merge, `withDefault` would still read + "original-withDefault"; under replace it takes the schema's declared default. + + `nullable` is the half that used to fail. fastavro fills an omitted field with + ``datum.get(name, field.get("default"))``, and celpy's MapType.get *raises* KeyError when + the key is absent and the default is None - so a null default blew up where a non-null one + worked. The executor now hands fastavro a plain dict. + """ + out = await _round_trip("drop", '{"kept": message.kept}') + + assert out["kept"] == "original-kept" + assert out["withDefault"] == "fallback" + assert out["withDefault"] != "original-withDefault", "merged instead of replacing" + assert out["nullable"] is None + + +async def test_naming_every_field_round_trips(): + """The must-fail twin. Without it, "the other fields took their defaults" is equally + consistent with the transform having stopped working altogether.""" + out = await _round_trip( + "all", '{"kept": message.kept, "withDefault": message.withDefault, ' '"nullable": message.nullable}' + ) + + assert out == _RECORD + + +# fastavro selects a union branch from a `(record_name, value)` pair, which is the only way to +# disambiguate two branches of the same shape, and `common/avro.py` preserves that pair through +# the field-level walk for exactly that reason. `_value_to_cel` has no tuple arm, so such a pair +# reaches a rule unconverted and comes back out of an identity transform unchanged - but +# `_to_plain_containers` flattened it to a list, and fastavro then refused the value outright: +# +# ValueError: ['B', {'x': 5}] (type ) do not match [{'type': 'record', ...}] +# +# The JVM has no tuple notation - a GenericRecord carries its own schema, so the branch is never +# ambiguous there - so the reference behaviour is simply that an identity transform preserves the +# branch selection, which is what this asserts. +_AMBIGUOUS_UNION_SCHEMA = { + "type": "record", + "name": "Outer", + "fields": [ + { + "name": "u", + "type": [ + {"type": "record", "name": "A", "fields": [{"name": "x", "type": "int"}]}, + {"type": "record", "name": "B", "fields": [{"name": "x", "type": "int"}]}, + ], + } + ], +} + + +async def _round_trip_union(subject_suffix, expr, record): + topic = _TOPIC + "-" + subject_suffix + client = AsyncSchemaRegistryClient.new_client({"url": "mock://"}) + rule = Rule("r", "", RuleKind.TRANSFORM, RuleMode.WRITE, "CEL", None, None, expr, None, None, False) + schema = Schema(json.dumps(_AMBIGUOUS_UNION_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])) + await client.register_schema(topic + "-value", schema) + ser = await AsyncAvroSerializer( + client, schema_str=None, conf={"auto.register.schemas": False, "use.latest.version": True} + ) + ctx = SerializationContext(topic, MessageField.VALUE) + payload = await ser(record, ctx) + deser = await AsyncAvroDeserializer(client) + return await deser(payload, ctx) + + +async def test_a_union_branch_selected_by_tuple_survives_the_transform(): + """The two branches have identical field shapes, so the tuple is load-bearing: without it + fastavro cannot tell A from B, and flattening it made the value match neither.""" + out = await _round_trip_union("union-tuple", '{"u": message.u}', {"u": ("B", {"x": 5})}) + + assert out["u"] == {"x": 5} + + +async def test_the_tuple_contents_are_still_normalised(): + """Preserving the tuple must not stop the recursion: the value inside it is a dict that + still has to reach fastavro as a plain one, which is what the whole function is for.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.cel_executor import _to_plain_containers + + inner = celtypes.MapType() + inner[celtypes.StringType("x")] = celtypes.IntType(5) + out = _to_plain_containers({"u": (celtypes.StringType("B"), inner)}) + + assert isinstance(out["u"], tuple) + # The branch name stays a celpy StringType, which is a str subclass, so fastavro's + # comparison against the record name works: only dict *keys* are normalised, as the + # function's own docstring says. + assert out["u"][0] == "B" + assert out["u"][1] == {"x": 5} and type(out["u"][1]) is dict diff --git a/tests/schema_registry/_async/test_proto.py b/tests/schema_registry/_async/test_proto.py index 3a7fe175c..50a6f2e0b 100644 --- a/tests/schema_registry/_async/test_proto.py +++ b/tests/schema_registry/_async/test_proto.py @@ -22,6 +22,7 @@ import pytest from google.protobuf import descriptor_pb2 +from confluent_kafka.schema_registry.confluent.type import decimal_pb2 from confluent_kafka.schema_registry.protobuf import ( AsyncProtobufDeserializer, AsyncProtobufSerializer, @@ -158,3 +159,245 @@ def test_proto_decimal(decimal, scale): converted = decimal_to_protobuf(input, scale) result = protobuf_to_decimal(converted) assert result == input + + +# BigDecimal.setScale(scale) narrows a scale whenever no rounding is needed -- only the digits +# being dropped must be zeros. decimal_to_protobuf used to refuse every reduction (`delta < 0`), +# which rejected exact conversions: Decimal("1.50") at scale 1, and the negative scale that +# protobuf_to_decimal itself produces for a value like 1E+3. Values requiring real rounding are +# still refused, as setScale does without a rounding mode. +@pytest.mark.parametrize( + "decimal, scale, unscaled, out_scale", + [ + ("12.3400", 2, 1234, 2), # trailing zeros dropped, exact + ("1.50", 1, 15, 1), + ("-1.50", 1, -15, 1), + ("1000", -3, 1, -3), # negative scale, exact + ("-1000", -3, -1, -3), + ("0.00", 0, 0, 0), + ("12.34", 4, 123400, 4), # widening still works + ("12.34", 2, 1234, 2), # exact match still works + ], +) +def test_proto_decimal_narrows_scale_losslessly(decimal, scale, unscaled, out_scale): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert int.from_bytes(msg.value, byteorder="big", signed=True) == unscaled + assert msg.scale == out_scale + + +@pytest.mark.parametrize("decimal, scale", [("12.345", 2), ("1.01", 1), ("999", -1)]) +def test_proto_decimal_rejects_lossy_scale(decimal, scale): + with pytest.raises(ValueError, match="Scale provided does not match the decimal"): + decimal_to_protobuf(Decimal(decimal), scale) + + +# Both rescaling directions used to build a power of ten before deciding anything, which the +# requested scale sizes: `10**-delta` for the exactness check when narrowing, `10**delta` for +# the coefficient when widening. Measured against this function before the guards: +# +# scale -1e7 -> 5.2s to reach a ValueError +# scale -1e8 -> 178s to reach the same ValueError +# scale 1e7 -> 5.3s and a 4.1 MB field written +# scale 1e8 -> did not finish inside 240s +# +# BigDecimal.setScale(scale) is the reference for the whole function, and it answers these +# without the arithmetic. Measured against the JDK: +# +# setScale(1, -1e7) THROW ArithmeticException: Rounding necessary +# setScale(0, -1e9) OK, scale=-1000000000, instant - a zero has no digits to lose +# setScale(0.00, -1e9) OK, same +# setScale(1, 1e7) OK, precision 10000001 (1.4s - the JVM pays here too) +# setScale(1, 1e9) THROW ArithmeticException: BigInteger would overflow ... +# setScale(1E+1000000000, 0) THROW, same +@pytest.mark.parametrize( + "decimal, scale", + [ + # Narrowing a non-zero value past its trailing zeros: the JVM's "Rounding necessary". + ("1", -10000000), + ("1", -1000000000), + ("1.23", -1000000000), + # Widening past what a BigInteger coefficient can hold. + ("1", 1000000000), + ("1E+1000000000", 0), + ("12.34", 2000000000), + ], +) +def test_proto_decimal_rejects_the_scales_java_rejects(decimal, scale): + with pytest.raises(ValueError): + decimal_to_protobuf(Decimal(decimal), scale) + + +@pytest.mark.parametrize( + "decimal, scale, unscaled", + [ + # A zero narrows to any scale, which is what setScale does with a zero coefficient. + ("0", -1000000000, 0), + ("0.00", -1000000000, 0), + ("0E+10", -1000000000, 0), + # And the ordinary cases keep working. + ("1000", -3, 1), + ("1.50", 1, 15), + ], +) +def test_proto_decimal_accepts_the_scales_java_accepts(decimal, scale, unscaled): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert int.from_bytes(msg.value, byteorder="big", signed=True) == unscaled + assert msg.scale == scale + + +# A timing bound, because the cost *is* the defect: the answer was already right, it just took +# 5.2s to give at this scale and 178s one power of ten further out. The fixed path measures +# 0.000s, so a one-second budget separates them by three orders of magnitude and cannot flake. +def test_proto_decimal_rejects_a_wide_scale_without_the_arithmetic(): + import time + + start = time.monotonic() + with pytest.raises(ValueError): + decimal_to_protobuf(Decimal("1"), -10000000) + with pytest.raises(ValueError): + decimal_to_protobuf(Decimal("1"), 10000000000) + assert time.monotonic() - start < 1.0 + + +# `protobuf_to_decimal` deliberately does **not** apply the message's precision, and these are +# the cases that tell the two policies apart. +# +# Java reads it as `new BigDecimal(unscaled, scale, new MathContext(precision))`, so a precision +# narrower than the coefficient's digit count *rounds the value on read*. Every client - this one +# included - writes precision as the unscaled value's own digit count, which makes that +# MathContext a guaranteed no-op, so it only ever has an effect on a message from a foreign +# producer carrying a declared/column precision. There its effect is to silently round data the +# producer sent exactly, which is why six of the seven clients ignore it and this path now does +# too. See ~/Documents/decimals.md section 2. +# +# The values below are what the reference would have returned for the same inputs, kept in the +# comments so the divergence is legible: unscaled 125 at precision 2 is 1.3E+2 on the JVM. +@pytest.mark.parametrize( + "unscaled, scale, precision, expected", + [ + ("12325", 0, 4, "12325"), # JVM, rounding to 4 digits: 1.233E+4 + ("125", 0, 2, "125"), # JVM: 1.3E+2 + ("-125", 0, 2, "-125"), # JVM: -1.3E+2 + ("12315", 0, 4, "12315"), # JVM: 1.232E+4 + ("135", 0, 2, "135"), # JVM: 1.4E+2 + ("12345", 2, 3, "123.45"), # JVM: 123 + # Where precision is at least the digit count - i.e. everything this client family + # writes - the two policies agree, and always did. + ("1234", 2, 4, "12.34"), + ("0", 2, 1, "0.00"), + ("1", -3, 1, "1E+3"), # negative scale survives + # And precision 0, which the reference cannot produce but three of our write paths did + # until now: MathContext(0) is UNLIMITED, so this agreed too. + ("12345", 2, 0, "123.45"), + ], +) +def test_protobuf_to_decimal_ignores_precision(unscaled, scale, precision, expected): + msg = decimal_pb2.Decimal( + value=int(unscaled).to_bytes(16, byteorder="big", signed=True), + scale=scale, + precision=precision, + ) + assert str(protobuf_to_decimal(msg)) == expected + + +# The other half of section 2: the same message read through the CEL binding and through the +# serde must give the same value. It did not - this path applied precision and that one never +# has - so unscaled 125 at precision 2 was 1.3E+2 here and 125 there. +def test_both_read_paths_agree_on_precision(): + from confluent_kafka.schema_registry.confluent.type.decimal_utils import from_proto_decimal + + msg = decimal_pb2.Decimal(value=(125).to_bytes(2, byteorder="big", signed=True), scale=0, precision=2) + assert protobuf_to_decimal(msg) == from_proto_decimal(msg) + + +# decimal_to_protobuf left `precision` at 0, which the reference cannot produce - +# `BigDecimal.precision()` is never less than 1, zero's precision being 1 - so a JVM consumer +# rewrites such a message on its next touch. It now carries the unscaled value's digit count, +# derived from the integer actually written so a rescale cannot leave it stale. +@pytest.mark.parametrize( + "decimal, scale, precision", + [ + ("12.34", 2, 4), + ("1000", -3, 1), # unscaled 1 + ("0", 0, 1), # BigDecimal.ZERO.precision() == 1 + ("0.00", 2, 1), + ("-1.50", 1, 2), # the sign is not a digit + ("12.3400", 2, 4), # after the exact narrowing, not before + ], +) +def test_decimal_to_protobuf_writes_the_digit_count(decimal, scale, precision): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert msg.precision == precision + assert msg.scale == scale + + +# The digit count is counted arithmetically, not through `len(str(abs(unscaled)))`. CPython caps +# str <-> int conversion at 4300 digits, so the string form raised `ValueError: Exceeds the limit +# (4300 digits) for integer string conversion` for a coefficient this function otherwise accepts +# (`_MAX_COEFFICIENT_DIGITS` is 646456993 here) - and raised it *after* `result.value` had been +# assigned. Every expected value below is measured on the JDK, which is what the count mirrors: +# BigDecimal("1").setScale(4300) -> precision 4301 +# BigDecimal("1").setScale(5000) -> precision 5001 +# BigDecimal("1").setScale(1000000) -> precision 1000001 +# BigDecimal("0").setScale(5000) -> precision 1 +@pytest.mark.parametrize( + "decimal, scale, precision", + [ + ("1", 4299, 4300), # the widest the string form managed + ("1", 4300, 4301), # the first it refused + ("1", 5000, 5001), + ("1", 100000, 100001), + ("12.34", 5000, 5002), # digits and delta both contribute + # Zero has the digit tuple (0,) at every scale, so it needs the special case; the + # reference agrees that it is 1 and not 1 + delta. + ("0", 5000, 1), + ("0", 100000, 1), + ], +) +def test_decimal_to_protobuf_counts_wide_precision_without_str(decimal, scale, precision): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert msg.precision == precision + assert msg.scale == scale + # And the value itself still round-trips, so the count describes what was written. + assert protobuf_to_decimal(msg) == Decimal(decimal) + + +# The arithmetic count has to agree with the string form everywhere the string form is legal, +# which is what makes it a refactor rather than a second implementation. It is exact in both +# rescale directions only because this function refuses an inexact narrowing: widening appends +# `delta` zeros with no carry, and narrowing only drops digits already proven to be zeros. A +# rounding rescale could carry - 9.9 to scale 0 is 10, one digit becoming two - and would need +# the count taken after the fact. +def test_precision_count_matches_the_string_form(): + checked = 0 + for text in [ + "0", + "0.00", + "-0.00", + "1", + "12.34", + "1.50", + "1000", + "0.001", + "-999.5", + "9.9", + "99.99", + "1E+3", + "1E-3", + "100", + "123456789012345678901234567890", + "0.0000000001", + "-1.50", + "1.000000", + "9" * 100, + "1" + "0" * 200, + ]: + for scale in range(-8, 12): + try: + msg = decimal_to_protobuf(Decimal(text), scale) + except ValueError: + continue # an inexact narrowing, which this function refuses + unscaled = int.from_bytes(msg.value, byteorder="big", signed=True) + assert msg.precision == len(str(abs(unscaled))), (text, scale) + checked += 1 + assert checked > 200 diff --git a/tests/schema_registry/_async/test_proto_serdes.py b/tests/schema_registry/_async/test_proto_serdes.py index 90ac644d3..1a49424c4 100644 --- a/tests/schema_registry/_async/test_proto_serdes.py +++ b/tests/schema_registry/_async/test_proto_serdes.py @@ -17,10 +17,11 @@ # import os import sys +import time import pytest -from confluent_kafka.schema_registry import Schema, header_schema_id_serializer +from confluent_kafka.schema_registry import Metadata, MetadataProperties, Schema, header_schema_id_serializer from confluent_kafka.schema_registry._async.protobuf import AsyncProtobufDeserializer, AsyncProtobufSerializer from confluent_kafka.schema_registry._async.schema_registry_client import AsyncSchemaRegistryClient from confluent_kafka.schema_registry._async.serde import ( @@ -34,6 +35,20 @@ ) from confluent_kafka.schema_registry.common.serde import SubjectNameStrategyType from confluent_kafka.schema_registry.protobuf import _schema_to_str +from confluent_kafka.schema_registry.rules.encryption.encrypt_executor import ( + Clock, + EncryptionExecutor, + FieldEncryptionExecutor, +) +from confluent_kafka.schema_registry.schema_registry_client import ( + Rule, + RuleKind, + RuleMode, + RuleParams, + RuleSet, + ServerConfig, +) +from confluent_kafka.schema_registry.serde import RuleConditionError from confluent_kafka.serialization import MessageField, SerializationContext, SerializationError # Add proto directory to sys.path to resolve protobuf import dependencies @@ -47,9 +62,22 @@ example_pb2, map_widget_pb2, nested_pb2, + newerwidget_pb2, + newwidget_pb2, test_pb2, + widget_pb2, ) + +class FakeClock(Clock): + + def __init__(self): + self.fixed_now = int(round(time.time() * 1000)) + + def now(self) -> int: + return self.fixed_now + + _BASE_URL = "mock://" # _BASE_URL = "http://localhost:8081" _TOPIC = "topic1" @@ -263,6 +291,662 @@ async def test_proto_cycle(): assert obj == obj2 +async def test_proto_cel_condition(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.name == 'Kafka'", + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_condition_fail(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.name != 'Kafka'", + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_proto_cel_field_transform(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + obj2 = example_pb2.Author( + name='Kafka-suffix', + id=123, + picture=b'foobar', + works=['The Castle-suffix', 'TheTrial-suffix'], + oneof_string='oneof-suffix', + ) + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +async def test_proto_cel_field_condition(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'name' ; value == 'Kafka'", + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj == newobj + + +async def test_proto_cel_field_condition_fail(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'name' ; value != 'Kafka'", + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_proto_cel_decimal_passes(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.gt(decimal("12.34"), decimal("10.00"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_decimal_fails(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.lt(decimal("12.34"), decimal("10.00"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_proto_cel_decimal_arithmetic(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.add(decimal("12.34"), decimal("1.66")), decimal("14.00"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_decimal_mod(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.mod(decimal("10"), decimal("3")), decimal("1"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_decimal_greatest_least(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.greatest(decimal("2.5"), decimal("9.99")), decimal("9.99")) ' + '&& decimals.eq(decimals.least(decimal("2.5"), decimal("9.99")), decimal("2.5"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_decimal_sqrt(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.sqrt(decimal("144")), decimal("12"))', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_decimal_to_double(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'double(decimal("100.50")) == 100.5', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_timestamp_passes(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.updated_at) < now', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(nested_pb2.NestedMessage.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = nested_pb2.NestedMessage() + obj.user_id.kafka_user_id = 'u1' + obj.updated_at.seconds = 1577836800 # 2020-01-01 UTC + ser = await AsyncProtobufSerializer(nested_pb2.NestedMessage, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(nested_pb2.NestedMessage, deser_conf, client) + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_cel_timestamp_fails(): + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.updated_at) > now', + None, + None, + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(nested_pb2.NestedMessage.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = nested_pb2.NestedMessage() + obj.user_id.kafka_user_id = 'u1' + obj.updated_at.seconds = 1577836800 # 2020-01-01 UTC, before now + ser = await AsyncProtobufSerializer(nested_pb2.NestedMessage, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + await ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +async def test_proto_encryption(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule_conf = {'secret': 'mysecret'} + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + # reset encrypted fields + assert obj.name != 'Kafka' + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_payload_encryption(): + executor = EncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule_conf = {'secret': 'mysecret'} + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT_PAYLOAD", + None, + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + await client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = await AsyncProtobufSerializer(example_pb2.Author, client, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(example_pb2.Author, deser_conf, client, rule_conf=rule_conf) + executor.client = dek_client + obj2 = await deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +async def test_proto_jsonata_fully_compatible(): + rule1_to_2 = "$merge([$sift($, function($v, $k) {$k != 'size'}), {'height': $.'size'}])" + rule2_to_1 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'size': $.'height'}])" + rule2_to_3 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'length': $.'height'}])" + rule3_to_2 = "$merge([$sift($, function($v, $k) {$k != 'length'}), {'height': $.'length'}])" + + conf = {'url': _BASE_URL} + client = AsyncSchemaRegistryClient.new_client(conf) + + await client.set_config(_SUBJECT, ServerConfig(compatibility_group='application.version')) + + await client.register_schema( + _SUBJECT, + Schema( + _schema_to_str(widget_pb2.Widget.DESCRIPTOR.file), + "PROTOBUF", + [], + Metadata(None, MetadataProperties({"application.version": "v1"}), None), + None, + ), + ) + + rule1 = Rule( + "rule1", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule1_to_2, None, None, False + ) + rule2 = Rule( + "rule2", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule2_to_1, None, None, False + ) + await client.register_schema( + _SUBJECT, + Schema( + _schema_to_str(newwidget_pb2.NewWidget.DESCRIPTOR.file), + "PROTOBUF", + [], + Metadata(None, MetadataProperties({"application.version": "v2"}), None), + RuleSet([rule1, rule2], None), + ), + ) + + rule3 = Rule( + "rule3", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule2_to_3, None, None, False + ) + rule4 = Rule( + "rule4", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule3_to_2, None, None, False + ) + await client.register_schema( + _SUBJECT, + Schema( + _schema_to_str(newerwidget_pb2.NewerWidget.DESCRIPTOR.file), + "PROTOBUF", + [], + Metadata(None, MetadataProperties({"application.version": "v3"}), None), + RuleSet([rule3, rule4], None), + ), + ) + + obj = widget_pb2.Widget(name='alice', size=123, version=1) + obj2 = newwidget_pb2.NewWidget(name='alice', height=123, version=1) + obj3 = newerwidget_pb2.NewerWidget(name='alice', length=123, version=1) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v1'}, + 'use.deprecated.format': False, + } + ser = await AsyncProtobufSerializer(widget_pb2.Widget, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj, ser_ctx) + + await deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v2'}, + 'use.deprecated.format': False, + } + ser = await AsyncProtobufSerializer(newwidget_pb2.NewWidget, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj2, ser_ctx) + + await deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v3'}, + 'use.deprecated.format': False, + } + ser = await AsyncProtobufSerializer(newerwidget_pb2.NewerWidget, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = await ser(obj3, ser_ctx) + + await deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + +async def deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3): + deser_conf = {'use.latest.with.metadata': {'application.version': 'v1'}, 'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(widget_pb2.Widget, deser_conf, client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj.size == newobj.size + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v2'}, 'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(newwidget_pb2.NewWidget, deser_conf, client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj2.height == newobj.height + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v3'}, 'use.deprecated.format': False} + deser = await AsyncProtobufDeserializer(newerwidget_pb2.NewerWidget, deser_conf, client) + newobj = await deser(obj_bytes, ser_ctx) + assert obj3.length == newobj.length + + async def test_associated_name_strategy_with_association(): """Test that AssociatedNameStrategy returns subject from association""" conf = {'url': _BASE_URL} diff --git a/tests/schema_registry/_sync/test_avro_serdes.py b/tests/schema_registry/_sync/test_avro_serdes.py index aca76da03..0b2ee4fe7 100644 --- a/tests/schema_registry/_sync/test_avro_serdes.py +++ b/tests/schema_registry/_sync/test_avro_serdes.py @@ -16,10 +16,16 @@ # limitations under the License. # import json +import time +from datetime import datetime, timedelta, timezone +from decimal import Decimal import pytest +from fastavro._logical_readers import UUID from confluent_kafka.schema_registry import ( + Metadata, + MetadataProperties, Schema, SchemaRegistryClient, header_schema_id_serializer, @@ -34,9 +40,41 @@ AssociationCreateOrUpdateRequest, ) from confluent_kafka.schema_registry.common.serde import SubjectNameStrategyType -from confluent_kafka.schema_registry.schema_registry_client import SchemaReference +from confluent_kafka.schema_registry.rule_registry import RuleOverride, RuleRegistry +from confluent_kafka.schema_registry.rules.cel.cel_executor import CelExecutor +from confluent_kafka.schema_registry.rules.cel.cel_field_executor import CelFieldExecutor +from confluent_kafka.schema_registry.rules.encryption.dek_registry.dek_registry_client import ( + DekAlgorithm, + DekRegistryClient, +) +from confluent_kafka.schema_registry.rules.encryption.encrypt_executor import ( + Clock, + EncryptionExecutor, + FieldEncryptionExecutor, +) +from confluent_kafka.schema_registry.rules.jsonata.jsonata_executor import JsonataExecutor +from confluent_kafka.schema_registry.schema_registry_client import ( + Rule, + RuleKind, + RuleMode, + RuleParams, + RuleSet, + SchemaReference, + ServerConfig, +) +from confluent_kafka.schema_registry.serde import RuleConditionError from confluent_kafka.serialization import MessageField, SerializationContext, SerializationError + +class FakeClock(Clock): + + def __init__(self): + self.fixed_now = int(round(time.time() * 1000)) + + def now(self) -> int: + return self.fixed_now + + _BASE_URL = "mock://" # _BASE_URL = "http://localhost:8081" _TOPIC = "topic1" @@ -46,6 +84,12 @@ @pytest.fixture(autouse=True) def run_before_and_after_tests(tmpdir): """Fixture to execute asserts before and after a test is run""" + # Setup: fill with any logic you want + + CelExecutor.register() + CelFieldExecutor.register() + JsonataExecutor.register() + yield # this is where the testing happens # Teardown : fill with any logic you want @@ -524,6 +568,2268 @@ def test_avro_schema_evolution(): assert obj2.get('newOptionalField') == 'optional' +def test_avro_cel_condition(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField == 'hi'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_condition_logical_type(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': {'type': 'string', 'logicalType': 'uuid'}}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + uuid = "550e8400-e29b-41d4-a716-446655440000" + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField == '" + uuid + "'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': UUID(uuid), + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_condition_fail(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField != 'hi'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_avro_cel_condition_ignore_fail(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.stringField != 'hi'", + None, + "NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_field_transform(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi-suffix', + 'booleanField': True, + 'bytesField': b'foobar', + } + deser = AvroDeserializer(client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_avro_cel_field_transform_missing_prop(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + {'name': 'missing', 'type': ['null', 'string'], 'default': None}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi-suffix-suffix', + 'booleanField': True, + 'bytesField': b'foobar', + 'missing': None, + } + deser = AvroDeserializer(client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_avro_cel_field_transform_disable(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + + registry = RuleRegistry() + registry.register_rule_executor(CelFieldExecutor()) + registry.register_override(RuleOverride("CEL_FIELD", None, None, True)) + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_registry=registry) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + newobj = deser(obj_bytes, ser_ctx) + assert "hi" == newobj['stringField'] + + +def test_avro_cel_field_transform_complex(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'arrayField', 'type': {'type': 'array', 'items': 'string'}}, + {'name': 'mapField', 'type': {'type': 'map', 'values': 'string'}}, + {'name': 'unionField', 'type': ['null', 'string'], 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'arrayField': ['hello'], + 'mapField': {'key': 'world'}, + 'unionField': 'bye', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = { + 'arrayField': ['hello-suffix'], + 'mapField': {'key': 'world-suffix'}, + 'unionField': 'bye-suffix', + } + deser = AvroDeserializer(client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_avro_cel_field_transform_complex_with_none(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'arrayField', 'type': {'type': 'array', 'items': 'string'}}, + {'name': 'mapField', 'type': {'type': 'map', 'values': 'string'}}, + {'name': 'unionField', 'type': ['null', 'string'], 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'arrayField': ['hello'], + 'mapField': {'key': 'world'}, + 'unionField': None, + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = { + 'arrayField': ['hello-suffix'], + 'mapField': {'key': 'world-suffix'}, + 'unionField': None, + } + deser = AvroDeserializer(client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_avro_cel_field_transform_complex_nested(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'UnionTest', + 'namespace': 'test', + 'fields': [ + { + 'name': 'emails', + 'type': [ + 'null', + { + 'type': 'array', + 'items': { + 'type': 'record', + 'name': 'Email', + 'fields': [ + { + 'name': 'email', + 'type': ['null', 'string'], + 'doc': 'Email address', + 'confluent:tags': ['PII'], + } + ], + }, + }, + ], + 'doc': 'Communication Email', + } + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'emails': [{'email': 'john@acme.com'}]} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = {'emails': [{'email': 'john@acme.com-suffix'}]} + deser = AvroDeserializer(client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_avro_cel_field_condition(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value == 'hi'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_field_condition_fail(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string'}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value == 'bye'", + None, + None, + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +_AVRO_DECIMAL_SCHEMA = { + 'type': 'record', + 'name': 'test', + 'fields': [ + { + 'name': 'decField', + 'type': { + 'type': 'bytes', + 'logicalType': 'decimal', + 'precision': 10, + 'scale': 2, + }, + }, + ], +} + + +def test_avro_cel_decimal_passes(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.gt(decimal(message.decField), decimal("10.00"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_decimal_fails(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.lt(decimal(message.decField), decimal("10.00"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_avro_cel_decimal_needs_no_constructor(): + """Cross-client parity: an Avro ``decimal`` logical type is usable as a Decimal with **no + ``decimal(...)`` call**, and the wrapped form keeps working alongside it. fastavro decodes it + to a Python ``Decimal`` at the schema's scale, which is this client's in-CEL decimal + representation, so ``decimals.*`` accept it directly and ``==`` is numeric (Python's own + ``Decimal.__eq__``). + """ + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + + def serialize(expr): + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + expr, + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + return ser({'decField': Decimal('12.34')}, ser_ctx) + + # Bare: no constructor call on the field. + assert serialize('decimals.eq(message.decField, decimal("12.34"))') is not None + assert serialize('decimals.gt(message.decField, decimal("10.00"))') is not None + # The wrapped form must keep working (decimal(...) re-entry). + assert serialize('decimals.eq(decimal(message.decField), decimal("12.34"))') is not None + # `==` is numeric on it: 12.34 equals 12.340 despite the differing scale. + assert serialize('message.decField == decimal("12.340")') is not None + # The schema's scale is applied, not guessed: as scale 0 this would be 1234. + assert serialize('decimals.lt(message.decField, decimal("100"))') is not None + # Negative control: a false comparison must fail. + with pytest.raises(SerializationError) as e: + serialize('decimals.gt(message.decField, decimal("100"))') + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_avro_cel_decimal_arithmetic(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.add(decimal(message.decField), decimal("1.66")), decimal("14.00"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_decimal_string(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'string(decimal(message.decField)) == "12.34"', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_DECIMAL_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'decField': Decimal('12.34')} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +_AVRO_TIMESTAMP_SCHEMA = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'tsField', 'type': {'type': 'long', 'logicalType': 'timestamp-millis'}}, + ], +} + + +def test_avro_cel_timestamp_millis_passes(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.tsField) < now', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_TIMESTAMP_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'tsField': datetime(2020, 1, 1, tzinfo=timezone.utc)} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_cel_timestamp_millis_fails(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.tsField) > now', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_TIMESTAMP_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + + obj = {'tsField': datetime(2020, 1, 1, tzinfo=timezone.utc)} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_avro_cel_timestamp_millis_needs_no_constructor(): + """Cross-client parity: an Avro timestamp logical type is usable as a timestamp with **no + constructor call at all**. fastavro decodes it to an aware datetime, which the boundary + binds as a CEL timestamp, so it is comparable against ``now`` and carries the timestamp + accessors. Every one of the seven clients has this test; the constructor is only needed for + a plain numeric field whose unit the schema cannot supply. + """ + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + + def serialize(expr, value): + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + expr, + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(json.dumps(_AVRO_TIMESTAMP_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])), + ) + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + return ser({'tsField': value}, ser_ctx) + + past = datetime(2020, 1, 1, tzinfo=timezone.utc) + exact = datetime(2023, 11, 14, 22, 13, 20, 123000, tzinfo=timezone.utc) + + # Bare comparison against `now`. + assert serialize('message.tsField < now', past) is not None + # The schema's millis unit is applied, not guessed, and the accessors work directly. + assert serialize('message.tsField == timestamp("2023-11-14T22:13:20.123Z")', exact) is not None + assert serialize('message.tsField.getFullYear() == 2023', exact) is not None + # Negative control: a future value must fail, so the comparison really happens. + with pytest.raises(SerializationError) as e: + serialize('message.tsField < now', datetime(2100, 1, 1, tzinfo=timezone.utc)) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_avro_encryption(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + obj['bytesField'] = b'foobar' + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_complex_schema(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + { + 'name': 'complexField1', + 'type': { + 'fields': [ + {'name': 'stringValue', 'type': 'string', 'confluent:tags': ['PII']}, + ], + 'name': 'ComplexFieldType', + 'type': 'record', + }, + }, + {'name': 'complexField2', 'type': 'ComplexFieldType'}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'complexField1': {'stringValue': 'test1'}, + 'complexField2': {'stringValue': 'test2'}, + } + + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + assert 'test1' not in str(obj_bytes) + assert 'test2' not in str(obj_bytes) + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + actual = deser(obj_bytes, ser_ctx) + assert actual['complexField1']['stringValue'] == 'test1' + assert actual['complexField2']['stringValue'] == 'test2' + + +def test_avro_encryption_complex_schema_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + { + 'name': 'complexField1', + 'type': { + 'fields': [ + { + 'name': 'complexSubType1', + 'type': { + 'fields': [{'name': 'stringValue', 'type': 'string', 'confluent:tags': ['PII']}], + 'name': 'ComplexSubType', + 'type': 'record', + }, + }, + {'name': 'complexSubType2', 'type': 'ComplexSubType'}, + ], + 'name': 'ComplexFieldType', + 'type': 'record', + }, + }, + {'name': 'complexField2', 'type': ['null', 'ComplexFieldType']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'complexField1': {'complexSubType1': {'stringValue': 'test1'}, 'complexSubType2': {'stringValue': 'test2'}}, + 'complexField2': {'complexSubType1': {'stringValue': 'test3'}, 'complexSubType2': {'stringValue': 'test4'}}, + } + + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + assert 'test1' not in str(obj_bytes) + assert 'test2' not in str(obj_bytes) + assert 'test3' not in str(obj_bytes) + assert 'test4' not in str(obj_bytes) + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + actual = deser(obj_bytes, ser_ctx) + assert actual['complexField1']['complexSubType1']['stringValue'] == 'test1' + assert actual['complexField1']['complexSubType2']['stringValue'] == 'test2' + assert actual['complexField2']['complexSubType1']['stringValue'] == 'test3' + assert actual['complexField2']['complexSubType2']['stringValue'] == 'test4' + + +def test_avro_payload_encryption(): + executor = EncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT_PAYLOAD", + None, + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_alternate_keks(): + executor = EncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret', 'encrypt.alternate.kms.key.ids': 'mykey2,mykey3'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT_PAYLOAD", + None, + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_deterministic(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.algorithm": "AES256_SIV", + } + ), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + obj['bytesField'] = b'foobar' + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_wrapped_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + "fields": [ + {"name": "id", "type": "int"}, + { + "name": "result", + "type": [ + "null", + { + "fields": [ + {"name": "code", "type": "int"}, + {"confluent:tags": ["PII"], "name": "secret", "type": ["null", "string"]}, + ], + "name": "Data", + "type": "record", + }, + { + "fields": [{"name": "code", "type": "int"}, {"name": "reason", "type": ["null", "string"]}], + "name": "Error", + "type": "record", + }, + ], + }, + ], + "name": "Result", + "namespace": "com.acme", + "type": "record", + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'id': 123, 'result': ('com.acme.Data', {'code': 456, 'secret': 'mypii'})} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['result'][1]['secret'] != 'mypii' + # remove union wrapper + obj['result'] = {'code': 456, 'secret': 'mypii'} + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_typed_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + "fields": [ + {"name": "id", "type": "int"}, + { + "name": "result", + "type": [ + "null", + { + "fields": [ + {"name": "code", "type": "int"}, + {"confluent:tags": ["PII"], "name": "secret", "type": ["null", "string"]}, + ], + "name": "Data", + "type": "record", + }, + { + "fields": [{"name": "code", "type": "int"}, {"name": "reason", "type": ["null", "string"]}], + "name": "Error", + "type": "record", + }, + ], + }, + ], + "name": "Result", + "namespace": "com.acme", + "type": "record", + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'id': 123, 'result': {'-type': 'com.acme.Data', 'code': 456, 'secret': 'mypii'}} + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['result']['secret'] != 'mypii' + # remove union wrapper + obj['result'] = {'code': 456, 'secret': 'mypii'} + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_cel(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes', 'confluent:tags': ['PII']}, + ], + } + + rule1 = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'stringField' ; value + '-suffix'", + None, + None, + False, + ) + rule2 = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule1, rule2]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi-suffix' + obj['stringField'] = 'hi-suffix' + obj['bytesField'] = b'foobar' + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_dek_rotation(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1-rot", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.expiry.days": "1", + } + ), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client: DekRegistryClient = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + dek_client = executor.executor.client + dek = dek_client.get_dek("kek1-rot", _SUBJECT, version=-1) + assert dek.version == 1 + + # advance 2 days + now = datetime.now() + timedelta(days=2) + executor.executor.clock.fixed_now = int(round(now.timestamp() * 1000)) + + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + dek = dek_client.get_dek("kek1-rot", _SUBJECT, version=-1) + assert dek.version == 2 + + # advance 2 days + now = datetime.now() + timedelta(days=2) + executor.executor.clock.fixed_now = int(round(now.timestamp() * 1000)) + + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + dek = dek_client.get_dek("kek1-rot", _SUBJECT, version=-1) + assert dek.version == 3 + + +def test_avro_encryption_f1_preserialized(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'f1Schema', + 'fields': [{'name': 'f1', 'type': 'string', 'confluent:tags': ['PII']}], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1-f1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,ERROR", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'f1': 'hello world'} + + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + deser = AvroDeserializer(client, rule_conf=rule_conf) + + dek_client: DekRegistryClient = executor.executor.client + dek_client.register_kek("kek1-f1", "local-kms", "mykey") + + encrypted_dek = "07V2ndh02DA73p+dTybwZFm7DKQSZN1tEwQh+FoX1DZLk4Yj2LLu4omYjp/84tAg3BYlkfGSz+zZacJHIE4=" + dek_client.register_dek("kek1-f1", _SUBJECT, encrypted_dek) + + obj_bytes = bytes( + [ + 0, + 0, + 0, + 0, + 1, + 104, + 122, + 103, + 121, + 47, + 106, + 70, + 78, + 77, + 86, + 47, + 101, + 70, + 105, + 108, + 97, + 72, + 114, + 77, + 121, + 101, + 66, + 103, + 100, + 97, + 86, + 122, + 114, + 82, + 48, + 117, + 100, + 71, + 101, + 111, + 116, + 87, + 56, + 99, + 65, + 47, + 74, + 97, + 108, + 55, + 117, + 107, + 114, + 43, + 77, + 47, + 121, + 122, + ] + ) + + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_deterministic_f1_preserialized(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'f1Schema', + 'fields': [{'name': 'f1', 'type': 'string', 'confluent:tags': ['PII']}], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1-det-f1", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.algorithm": "AES256_SIV", + } + ), + None, + None, + "ERROR,ERROR", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'f1': 'hello world'} + + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + deser = AvroDeserializer(client, rule_conf=rule_conf) + + dek_client: DekRegistryClient = executor.executor.client + dek_client.register_kek("kek1-det-f1", "local-kms", "mykey") + + encrypted_dek = ( + "YSx3DTlAHrmpoDChquJMifmPntBzxgRVdMzgYL82rgWBKn7aUSnG+WIu9oz" + "BNS3y2vXd++mBtK07w4/W/G6w0da39X9hfOVZsGnkSvry/QRht84V8yz3dqKxGMOK5A==" + ) + dek_client.register_dek("kek1-det-f1", _SUBJECT, encrypted_dek, algorithm=DekAlgorithm.AES256_SIV) + + obj_bytes = bytes( + [ + 0, + 0, + 0, + 0, + 1, + 72, + 68, + 54, + 89, + 116, + 120, + 114, + 108, + 66, + 110, + 107, + 84, + 87, + 87, + 57, + 78, + 54, + 86, + 98, + 107, + 51, + 73, + 73, + 110, + 106, + 87, + 72, + 56, + 49, + 120, + 109, + 89, + 104, + 51, + 107, + 52, + 100, + ] + ) + + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_dek_rotation_f1_preserialized(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'f1Schema', + 'fields': [{'name': 'f1', 'type': 'string', 'confluent:tags': ['PII']}], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams( + { + "encrypt.kek.name": "kek1-rot-f1", + "encrypt.kms.type": "local-kms", + "encrypt.kms.key.id": "mykey", + "encrypt.dek.expiry.days": "1", + } + ), + None, + None, + "ERROR,ERROR", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = {'f1': 'hello world'} + + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + deser = AvroDeserializer(client, rule_conf=rule_conf) + + dek_client: DekRegistryClient = executor.executor.client + dek_client.register_kek("kek1-rot-f1", "local-kms", "mykey") + + encrypted_dek = "W/v6hOQYq1idVAcs1pPWz9UUONMVZW4IrglTnG88TsWjeCjxmtRQ4VaNe/I5dCfm2zyY9Cu0nqdvqImtUk4=" + dek_client.register_dek("kek1-rot-f1", _SUBJECT, encrypted_dek, algorithm=DekAlgorithm.AES256_GCM) + + obj_bytes = bytes( + [ + 0, + 0, + 0, + 0, + 1, + 120, + 65, + 65, + 65, + 65, + 65, + 65, + 71, + 52, + 72, + 73, + 54, + 98, + 49, + 110, + 88, + 80, + 88, + 113, + 76, + 121, + 71, + 56, + 99, + 73, + 73, + 51, + 53, + 78, + 72, + 81, + 115, + 101, + 113, + 113, + 85, + 67, + 100, + 43, + 73, + 101, + 76, + 101, + 70, + 86, + 65, + 101, + 78, + 112, + 83, + 83, + 51, + 102, + 120, + 80, + 110, + 74, + 51, + 50, + 65, + 61, + ] + ) + + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_references(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + + referenced = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + obj = {'refField': referenced} + ref_schema = { + 'type': 'record', + 'name': 'ref', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': 'string', 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': 'bytes'}, + ], + } + client.register_schema('ref', Schema(json.dumps(ref_schema))) + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'refField', 'type': 'ref'}, + ], + } + refs = [SchemaReference('ref', 'ref', 1)] + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1-ref", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", refs, None, RuleSet(None, [rule]))) + + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['refField']['stringField'] != 'hi' + obj['refField']['stringField'] = 'hi' + obj['refField']['bytesField'] = b'foobar' + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_encryption_with_union(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True} + rule_conf = {'secret': 'mysecret'} + schema = { + 'type': 'record', + 'name': 'test', + 'fields': [ + {'name': 'intField', 'type': 'int'}, + {'name': 'doubleField', 'type': 'double'}, + {'name': 'stringField', 'type': ['null', 'string'], 'confluent:tags': ['PII']}, + {'name': 'booleanField', 'type': 'boolean'}, + {'name': 'bytesField', 'type': ['null', 'bytes'], 'confluent:tags': ['PII']}, + ], + } + + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1-union", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema(_SUBJECT, Schema(json.dumps(schema), "AVRO", [], None, RuleSet(None, [rule]))) + + obj = { + 'intField': 123, + 'doubleField': 45.67, + 'stringField': 'hi', + 'booleanField': True, + 'bytesField': b'foobar', + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj['stringField'] != 'hi' + obj['stringField'] = 'hi' + obj['bytesField'] = b'foobar' + + deser = AvroDeserializer(client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_avro_jsonata_with_cel(): + rule1_to_2 = "$merge([$sift($, function($v, $k) {$k != 'size'}), {'height': $.'size'}])" + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + + client.set_config(_SUBJECT, ServerConfig(compatibility_group='application.version')) + + schema = { + 'type': 'record', + 'name': 'old', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'size', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v1"}), None), + None, + ), + ) + + schema = { + 'type': 'record', + 'name': 'new', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'height', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + + rule1 = Rule( + "test-jsonata", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule1_to_2, None, None, False + ) + rule2 = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.READ, + "CEL_FIELD", + None, + None, + "name == 'name' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v2"}), None), + RuleSet([rule1], [rule2]), + ), + ) + + obj = { + 'name': 'alice', + 'size': 123, + 'version': 1, + } + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v1'}, + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = { + 'name': 'alice-suffix', + 'height': 123, + 'version': 1, + } + deser_conf = {'use.latest.with.metadata': {'application.version': 'v2'}} + deser = AvroDeserializer(client, conf=deser_conf) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_avro_jsonata_fully_compatible(): + rule1_to_2 = "$merge([$sift($, function($v, $k) {$k != 'size'}), {'height': $.'size'}])" + rule2_to_1 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'size': $.'height'}])" + rule2_to_3 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'length': $.'height'}])" + rule3_to_2 = "$merge([$sift($, function($v, $k) {$k != 'length'}), {'height': $.'length'}])" + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + + client.set_config(_SUBJECT, ServerConfig(compatibility_group='application.version')) + + schema = { + 'type': 'record', + 'name': 'old', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'size', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v1"}), None), + None, + ), + ) + + schema = { + 'type': 'record', + 'name': 'new', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'height', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + + rule1 = Rule( + "rule1", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule1_to_2, None, None, False + ) + rule2 = Rule( + "rule2", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule2_to_1, None, None, False + ) + client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v2"}), None), + RuleSet([rule1, rule2], None), + ), + ) + + schema = { + 'type': 'record', + 'name': 'newer', + 'fields': [ + {'name': 'name', 'type': 'string'}, + {'name': 'length', 'type': 'int'}, + {'name': 'version', 'type': 'int'}, + ], + } + + rule3 = Rule( + "rule3", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule2_to_3, None, None, False + ) + rule4 = Rule( + "rule4", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule3_to_2, None, None, False + ) + client.register_schema( + _SUBJECT, + Schema( + json.dumps(schema), + "AVRO", + [], + Metadata(None, MetadataProperties({"application.version": "v3"}), None), + RuleSet([rule3, rule4], None), + ), + ) + + obj = { + 'name': 'alice', + 'size': 123, + 'version': 1, + } + obj2 = { + 'name': 'alice', + 'height': 123, + 'version': 1, + } + obj3 = { + 'name': 'alice', + 'length': 123, + 'version': 1, + } + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v1'}, + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v2'}, + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj2, ser_ctx) + + deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v3'}, + } + ser = AvroSerializer(client, schema_str=None, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj3, ser_ctx) + + deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + +def deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3): + deser_conf = {'use.latest.with.metadata': {'application.version': 'v1'}} + deser = AvroDeserializer(client, conf=deser_conf) + newobj = deser(obj_bytes, ser_ctx) + assert obj == newobj + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v2'}} + deser = AvroDeserializer(client, conf=deser_conf) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v3'}} + deser = AvroDeserializer(client, conf=deser_conf) + newobj = deser(obj_bytes, ser_ctx) + assert obj3 == newobj + + def test_avro_reference(): conf = {'url': _BASE_URL} client = SchemaRegistryClient.new_client(conf) diff --git a/tests/schema_registry/_sync/test_cel_avro_message_transform.py b/tests/schema_registry/_sync/test_cel_avro_message_transform.py new file mode 100644 index 000000000..65bd5468c --- /dev/null +++ b/tests/schema_registry/_sync/test_cel_avro_message_transform.py @@ -0,0 +1,171 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +Message-level ``CEL`` transforms over Avro, and specifically that they **replace** rather than +merge: the rule's map is the whole new record, so a field the rule does not name takes the +schema's declared default rather than the value it had on the way in. + +This case existed only on the protobuf side, and its absence hid a real defect elsewhere - the +C++ client seeded its result record from the input before applying the map, so it merged. Every +other C6/C7 case names *all* of the record's fields, which makes merge and replace +indistinguishable, the same way a condition fixture cannot tell a passing rule from an unfired +one without a must-fail twin. + +Driven end to end through the serializer rather than through the executor alone, which is not +incidental: this client's Avro write-back hands fastavro the rule's result more or less +unchanged, so whether fastavro can fill an omitted field is the whole question. An +executor-level test cannot see it. +""" + +import json + +from confluent_kafka.schema_registry import Schema +from confluent_kafka.schema_registry._sync.schema_registry_client import SchemaRegistryClient +from confluent_kafka.schema_registry.avro import AvroDeserializer, AvroSerializer +from confluent_kafka.schema_registry.rules.cel.cel_executor import CelExecutor +from confluent_kafka.schema_registry.rules.cel.cel_field_executor import CelFieldExecutor +from confluent_kafka.schema_registry.schema_registry_client import Rule, RuleKind, RuleMode, RuleSet +from confluent_kafka.serialization import MessageField, SerializationContext + +CelExecutor.register() +CelFieldExecutor.register() + +_TOPIC = "cel-avro-message-transform" + +_SCHEMA = { + "type": "record", + "name": "Defaults", + "fields": [ + {"name": "kept", "type": "string"}, + {"name": "withDefault", "type": "string", "default": "fallback"}, + {"name": "nullable", "type": ["null", "string"], "default": None}, + ], +} + +_RECORD = { + "kept": "original-kept", + "withDefault": "original-withDefault", + "nullable": "original-nullable", +} + + +def _round_trip(subject_suffix, expr): + """Serializes the fixture under one message-level CEL transform and reads it back.""" + topic = _TOPIC + "-" + subject_suffix + client = SchemaRegistryClient.new_client({"url": "mock://"}) + rule = Rule("r", "", RuleKind.TRANSFORM, RuleMode.WRITE, "CEL", None, None, expr, None, None, False) + client.register_schema(topic + "-value", Schema(json.dumps(_SCHEMA), "AVRO", [], None, RuleSet(None, [rule]))) + ser = AvroSerializer(client, schema_str=None, conf={"auto.register.schemas": False, "use.latest.version": True}) + ctx = SerializationContext(topic, MessageField.VALUE) + # Each is on its own call: tools/unasync.py strips "await " by word boundary, so + # "await (" survives the rewrite and the generated sync file will not parse. + payload = ser(_RECORD, ctx) + deser = AvroDeserializer(client) + return deser(payload, ctx) + + +def test_a_field_the_rule_does_not_name_takes_its_declared_default(): + """The case this file exists for. Under merge, `withDefault` would still read + "original-withDefault"; under replace it takes the schema's declared default. + + `nullable` is the half that used to fail. fastavro fills an omitted field with + ``datum.get(name, field.get("default"))``, and celpy's MapType.get *raises* KeyError when + the key is absent and the default is None - so a null default blew up where a non-null one + worked. The executor now hands fastavro a plain dict. + """ + out = _round_trip("drop", '{"kept": message.kept}') + + assert out["kept"] == "original-kept" + assert out["withDefault"] == "fallback" + assert out["withDefault"] != "original-withDefault", "merged instead of replacing" + assert out["nullable"] is None + + +def test_naming_every_field_round_trips(): + """The must-fail twin. Without it, "the other fields took their defaults" is equally + consistent with the transform having stopped working altogether.""" + out = _round_trip( + "all", '{"kept": message.kept, "withDefault": message.withDefault, ' '"nullable": message.nullable}' + ) + + assert out == _RECORD + + +# fastavro selects a union branch from a `(record_name, value)` pair, which is the only way to +# disambiguate two branches of the same shape, and `common/avro.py` preserves that pair through +# the field-level walk for exactly that reason. `_value_to_cel` has no tuple arm, so such a pair +# reaches a rule unconverted and comes back out of an identity transform unchanged - but +# `_to_plain_containers` flattened it to a list, and fastavro then refused the value outright: +# +# ValueError: ['B', {'x': 5}] (type ) do not match [{'type': 'record', ...}] +# +# The JVM has no tuple notation - a GenericRecord carries its own schema, so the branch is never +# ambiguous there - so the reference behaviour is simply that an identity transform preserves the +# branch selection, which is what this asserts. +_AMBIGUOUS_UNION_SCHEMA = { + "type": "record", + "name": "Outer", + "fields": [ + { + "name": "u", + "type": [ + {"type": "record", "name": "A", "fields": [{"name": "x", "type": "int"}]}, + {"type": "record", "name": "B", "fields": [{"name": "x", "type": "int"}]}, + ], + } + ], +} + + +def _round_trip_union(subject_suffix, expr, record): + topic = _TOPIC + "-" + subject_suffix + client = SchemaRegistryClient.new_client({"url": "mock://"}) + rule = Rule("r", "", RuleKind.TRANSFORM, RuleMode.WRITE, "CEL", None, None, expr, None, None, False) + schema = Schema(json.dumps(_AMBIGUOUS_UNION_SCHEMA), "AVRO", [], None, RuleSet(None, [rule])) + client.register_schema(topic + "-value", schema) + ser = AvroSerializer(client, schema_str=None, conf={"auto.register.schemas": False, "use.latest.version": True}) + ctx = SerializationContext(topic, MessageField.VALUE) + payload = ser(record, ctx) + deser = AvroDeserializer(client) + return deser(payload, ctx) + + +def test_a_union_branch_selected_by_tuple_survives_the_transform(): + """The two branches have identical field shapes, so the tuple is load-bearing: without it + fastavro cannot tell A from B, and flattening it made the value match neither.""" + out = _round_trip_union("union-tuple", '{"u": message.u}', {"u": ("B", {"x": 5})}) + + assert out["u"] == {"x": 5} + + +def test_the_tuple_contents_are_still_normalised(): + """Preserving the tuple must not stop the recursion: the value inside it is a dict that + still has to reach fastavro as a plain one, which is what the whole function is for.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.cel_executor import _to_plain_containers + + inner = celtypes.MapType() + inner[celtypes.StringType("x")] = celtypes.IntType(5) + out = _to_plain_containers({"u": (celtypes.StringType("B"), inner)}) + + assert isinstance(out["u"], tuple) + # The branch name stays a celpy StringType, which is a str subclass, so fastavro's + # comparison against the record name works: only dict *keys* are normalised, as the + # function's own docstring says. + assert out["u"][0] == "B" + assert out["u"][1] == {"x": 5} and type(out["u"][1]) is dict diff --git a/tests/schema_registry/_sync/test_proto.py b/tests/schema_registry/_sync/test_proto.py index c9b932987..3844b0899 100644 --- a/tests/schema_registry/_sync/test_proto.py +++ b/tests/schema_registry/_sync/test_proto.py @@ -22,6 +22,7 @@ import pytest from google.protobuf import descriptor_pb2 +from confluent_kafka.schema_registry.confluent.type import decimal_pb2 from confluent_kafka.schema_registry.protobuf import ( ProtobufDeserializer, ProtobufSerializer, @@ -158,3 +159,245 @@ def test_proto_decimal(decimal, scale): converted = decimal_to_protobuf(input, scale) result = protobuf_to_decimal(converted) assert result == input + + +# BigDecimal.setScale(scale) narrows a scale whenever no rounding is needed -- only the digits +# being dropped must be zeros. decimal_to_protobuf used to refuse every reduction (`delta < 0`), +# which rejected exact conversions: Decimal("1.50") at scale 1, and the negative scale that +# protobuf_to_decimal itself produces for a value like 1E+3. Values requiring real rounding are +# still refused, as setScale does without a rounding mode. +@pytest.mark.parametrize( + "decimal, scale, unscaled, out_scale", + [ + ("12.3400", 2, 1234, 2), # trailing zeros dropped, exact + ("1.50", 1, 15, 1), + ("-1.50", 1, -15, 1), + ("1000", -3, 1, -3), # negative scale, exact + ("-1000", -3, -1, -3), + ("0.00", 0, 0, 0), + ("12.34", 4, 123400, 4), # widening still works + ("12.34", 2, 1234, 2), # exact match still works + ], +) +def test_proto_decimal_narrows_scale_losslessly(decimal, scale, unscaled, out_scale): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert int.from_bytes(msg.value, byteorder="big", signed=True) == unscaled + assert msg.scale == out_scale + + +@pytest.mark.parametrize("decimal, scale", [("12.345", 2), ("1.01", 1), ("999", -1)]) +def test_proto_decimal_rejects_lossy_scale(decimal, scale): + with pytest.raises(ValueError, match="Scale provided does not match the decimal"): + decimal_to_protobuf(Decimal(decimal), scale) + + +# Both rescaling directions used to build a power of ten before deciding anything, which the +# requested scale sizes: `10**-delta` for the exactness check when narrowing, `10**delta` for +# the coefficient when widening. Measured against this function before the guards: +# +# scale -1e7 -> 5.2s to reach a ValueError +# scale -1e8 -> 178s to reach the same ValueError +# scale 1e7 -> 5.3s and a 4.1 MB field written +# scale 1e8 -> did not finish inside 240s +# +# BigDecimal.setScale(scale) is the reference for the whole function, and it answers these +# without the arithmetic. Measured against the JDK: +# +# setScale(1, -1e7) THROW ArithmeticException: Rounding necessary +# setScale(0, -1e9) OK, scale=-1000000000, instant - a zero has no digits to lose +# setScale(0.00, -1e9) OK, same +# setScale(1, 1e7) OK, precision 10000001 (1.4s - the JVM pays here too) +# setScale(1, 1e9) THROW ArithmeticException: BigInteger would overflow ... +# setScale(1E+1000000000, 0) THROW, same +@pytest.mark.parametrize( + "decimal, scale", + [ + # Narrowing a non-zero value past its trailing zeros: the JVM's "Rounding necessary". + ("1", -10000000), + ("1", -1000000000), + ("1.23", -1000000000), + # Widening past what a BigInteger coefficient can hold. + ("1", 1000000000), + ("1E+1000000000", 0), + ("12.34", 2000000000), + ], +) +def test_proto_decimal_rejects_the_scales_java_rejects(decimal, scale): + with pytest.raises(ValueError): + decimal_to_protobuf(Decimal(decimal), scale) + + +@pytest.mark.parametrize( + "decimal, scale, unscaled", + [ + # A zero narrows to any scale, which is what setScale does with a zero coefficient. + ("0", -1000000000, 0), + ("0.00", -1000000000, 0), + ("0E+10", -1000000000, 0), + # And the ordinary cases keep working. + ("1000", -3, 1), + ("1.50", 1, 15), + ], +) +def test_proto_decimal_accepts_the_scales_java_accepts(decimal, scale, unscaled): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert int.from_bytes(msg.value, byteorder="big", signed=True) == unscaled + assert msg.scale == scale + + +# A timing bound, because the cost *is* the defect: the answer was already right, it just took +# 5.2s to give at this scale and 178s one power of ten further out. The fixed path measures +# 0.000s, so a one-second budget separates them by three orders of magnitude and cannot flake. +def test_proto_decimal_rejects_a_wide_scale_without_the_arithmetic(): + import time + + start = time.monotonic() + with pytest.raises(ValueError): + decimal_to_protobuf(Decimal("1"), -10000000) + with pytest.raises(ValueError): + decimal_to_protobuf(Decimal("1"), 10000000000) + assert time.monotonic() - start < 1.0 + + +# `protobuf_to_decimal` deliberately does **not** apply the message's precision, and these are +# the cases that tell the two policies apart. +# +# Java reads it as `new BigDecimal(unscaled, scale, new MathContext(precision))`, so a precision +# narrower than the coefficient's digit count *rounds the value on read*. Every client - this one +# included - writes precision as the unscaled value's own digit count, which makes that +# MathContext a guaranteed no-op, so it only ever has an effect on a message from a foreign +# producer carrying a declared/column precision. There its effect is to silently round data the +# producer sent exactly, which is why six of the seven clients ignore it and this path now does +# too. See ~/Documents/decimals.md section 2. +# +# The values below are what the reference would have returned for the same inputs, kept in the +# comments so the divergence is legible: unscaled 125 at precision 2 is 1.3E+2 on the JVM. +@pytest.mark.parametrize( + "unscaled, scale, precision, expected", + [ + ("12325", 0, 4, "12325"), # JVM, rounding to 4 digits: 1.233E+4 + ("125", 0, 2, "125"), # JVM: 1.3E+2 + ("-125", 0, 2, "-125"), # JVM: -1.3E+2 + ("12315", 0, 4, "12315"), # JVM: 1.232E+4 + ("135", 0, 2, "135"), # JVM: 1.4E+2 + ("12345", 2, 3, "123.45"), # JVM: 123 + # Where precision is at least the digit count - i.e. everything this client family + # writes - the two policies agree, and always did. + ("1234", 2, 4, "12.34"), + ("0", 2, 1, "0.00"), + ("1", -3, 1, "1E+3"), # negative scale survives + # And precision 0, which the reference cannot produce but three of our write paths did + # until now: MathContext(0) is UNLIMITED, so this agreed too. + ("12345", 2, 0, "123.45"), + ], +) +def test_protobuf_to_decimal_ignores_precision(unscaled, scale, precision, expected): + msg = decimal_pb2.Decimal( + value=int(unscaled).to_bytes(16, byteorder="big", signed=True), + scale=scale, + precision=precision, + ) + assert str(protobuf_to_decimal(msg)) == expected + + +# The other half of section 2: the same message read through the CEL binding and through the +# serde must give the same value. It did not - this path applied precision and that one never +# has - so unscaled 125 at precision 2 was 1.3E+2 here and 125 there. +def test_both_read_paths_agree_on_precision(): + from confluent_kafka.schema_registry.confluent.type.decimal_utils import from_proto_decimal + + msg = decimal_pb2.Decimal(value=(125).to_bytes(2, byteorder="big", signed=True), scale=0, precision=2) + assert protobuf_to_decimal(msg) == from_proto_decimal(msg) + + +# decimal_to_protobuf left `precision` at 0, which the reference cannot produce - +# `BigDecimal.precision()` is never less than 1, zero's precision being 1 - so a JVM consumer +# rewrites such a message on its next touch. It now carries the unscaled value's digit count, +# derived from the integer actually written so a rescale cannot leave it stale. +@pytest.mark.parametrize( + "decimal, scale, precision", + [ + ("12.34", 2, 4), + ("1000", -3, 1), # unscaled 1 + ("0", 0, 1), # BigDecimal.ZERO.precision() == 1 + ("0.00", 2, 1), + ("-1.50", 1, 2), # the sign is not a digit + ("12.3400", 2, 4), # after the exact narrowing, not before + ], +) +def test_decimal_to_protobuf_writes_the_digit_count(decimal, scale, precision): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert msg.precision == precision + assert msg.scale == scale + + +# The digit count is counted arithmetically, not through `len(str(abs(unscaled)))`. CPython caps +# str <-> int conversion at 4300 digits, so the string form raised `ValueError: Exceeds the limit +# (4300 digits) for integer string conversion` for a coefficient this function otherwise accepts +# (`_MAX_COEFFICIENT_DIGITS` is 646456993 here) - and raised it *after* `result.value` had been +# assigned. Every expected value below is measured on the JDK, which is what the count mirrors: +# BigDecimal("1").setScale(4300) -> precision 4301 +# BigDecimal("1").setScale(5000) -> precision 5001 +# BigDecimal("1").setScale(1000000) -> precision 1000001 +# BigDecimal("0").setScale(5000) -> precision 1 +@pytest.mark.parametrize( + "decimal, scale, precision", + [ + ("1", 4299, 4300), # the widest the string form managed + ("1", 4300, 4301), # the first it refused + ("1", 5000, 5001), + ("1", 100000, 100001), + ("12.34", 5000, 5002), # digits and delta both contribute + # Zero has the digit tuple (0,) at every scale, so it needs the special case; the + # reference agrees that it is 1 and not 1 + delta. + ("0", 5000, 1), + ("0", 100000, 1), + ], +) +def test_decimal_to_protobuf_counts_wide_precision_without_str(decimal, scale, precision): + msg = decimal_to_protobuf(Decimal(decimal), scale) + assert msg.precision == precision + assert msg.scale == scale + # And the value itself still round-trips, so the count describes what was written. + assert protobuf_to_decimal(msg) == Decimal(decimal) + + +# The arithmetic count has to agree with the string form everywhere the string form is legal, +# which is what makes it a refactor rather than a second implementation. It is exact in both +# rescale directions only because this function refuses an inexact narrowing: widening appends +# `delta` zeros with no carry, and narrowing only drops digits already proven to be zeros. A +# rounding rescale could carry - 9.9 to scale 0 is 10, one digit becoming two - and would need +# the count taken after the fact. +def test_precision_count_matches_the_string_form(): + checked = 0 + for text in [ + "0", + "0.00", + "-0.00", + "1", + "12.34", + "1.50", + "1000", + "0.001", + "-999.5", + "9.9", + "99.99", + "1E+3", + "1E-3", + "100", + "123456789012345678901234567890", + "0.0000000001", + "-1.50", + "1.000000", + "9" * 100, + "1" + "0" * 200, + ]: + for scale in range(-8, 12): + try: + msg = decimal_to_protobuf(Decimal(text), scale) + except ValueError: + continue # an inexact narrowing, which this function refuses + unscaled = int.from_bytes(msg.value, byteorder="big", signed=True) + assert msg.precision == len(str(abs(unscaled))), (text, scale) + checked += 1 + assert checked > 200 diff --git a/tests/schema_registry/_sync/test_proto_serdes.py b/tests/schema_registry/_sync/test_proto_serdes.py index a1228e825..93f09c185 100644 --- a/tests/schema_registry/_sync/test_proto_serdes.py +++ b/tests/schema_registry/_sync/test_proto_serdes.py @@ -17,10 +17,11 @@ # import os import sys +import time import pytest -from confluent_kafka.schema_registry import Schema, header_schema_id_serializer +from confluent_kafka.schema_registry import Metadata, MetadataProperties, Schema, header_schema_id_serializer from confluent_kafka.schema_registry._sync.protobuf import ProtobufDeserializer, ProtobufSerializer from confluent_kafka.schema_registry._sync.schema_registry_client import SchemaRegistryClient from confluent_kafka.schema_registry._sync.serde import ( @@ -34,6 +35,20 @@ ) from confluent_kafka.schema_registry.common.serde import SubjectNameStrategyType from confluent_kafka.schema_registry.protobuf import _schema_to_str +from confluent_kafka.schema_registry.rules.encryption.encrypt_executor import ( + Clock, + EncryptionExecutor, + FieldEncryptionExecutor, +) +from confluent_kafka.schema_registry.schema_registry_client import ( + Rule, + RuleKind, + RuleMode, + RuleParams, + RuleSet, + ServerConfig, +) +from confluent_kafka.schema_registry.serde import RuleConditionError from confluent_kafka.serialization import MessageField, SerializationContext, SerializationError # Add proto directory to sys.path to resolve protobuf import dependencies @@ -47,9 +62,22 @@ example_pb2, map_widget_pb2, nested_pb2, + newerwidget_pb2, + newwidget_pb2, test_pb2, + widget_pb2, ) + +class FakeClock(Clock): + + def __init__(self): + self.fixed_now = int(round(time.time() * 1000)) + + def now(self) -> int: + return self.fixed_now + + _BASE_URL = "mock://" # _BASE_URL = "http://localhost:8081" _TOPIC = "topic1" @@ -263,6 +291,662 @@ def test_proto_cycle(): assert obj == obj2 +def test_proto_cel_condition(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.name == 'Kafka'", + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_condition_fail(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + "message.name != 'Kafka'", + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_proto_cel_field_transform(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.TRANSFORM, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "typeName == 'STRING' ; value + '-suffix'", + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + obj2 = example_pb2.Author( + name='Kafka-suffix', + id=123, + picture=b'foobar', + works=['The Castle-suffix', 'TheTrial-suffix'], + oneof_string='oneof-suffix', + ) + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2 == newobj + + +def test_proto_cel_field_condition(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'name' ; value == 'Kafka'", + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + newobj = deser(obj_bytes, ser_ctx) + assert obj == newobj + + +def test_proto_cel_field_condition_fail(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL_FIELD", + None, + None, + "name == 'name' ; value != 'Kafka'", + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_proto_cel_decimal_passes(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.gt(decimal("12.34"), decimal("10.00"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_decimal_fails(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.lt(decimal("12.34"), decimal("10.00"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_proto_cel_decimal_arithmetic(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.add(decimal("12.34"), decimal("1.66")), decimal("14.00"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_decimal_mod(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.mod(decimal("10"), decimal("3")), decimal("1"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_decimal_greatest_least(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.greatest(decimal("2.5"), decimal("9.99")), decimal("9.99")) ' + '&& decimals.eq(decimals.least(decimal("2.5"), decimal("9.99")), decimal("2.5"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_decimal_sqrt(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.eq(decimals.sqrt(decimal("144")), decimal("12"))', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_decimal_to_double(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'double(decimal("100.50")) == 100.5', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author(name='Kafka', id=123, picture=b'foobar') + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_timestamp_passes(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.updated_at) < now', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(nested_pb2.NestedMessage.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = nested_pb2.NestedMessage() + obj.user_id.kafka_user_id = 'u1' + obj.updated_at.seconds = 1577836800 # 2020-01-01 UTC + ser = ProtobufSerializer(nested_pb2.NestedMessage, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(nested_pb2.NestedMessage, deser_conf, client) + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_cel_timestamp_fails(): + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule = Rule( + "test-cel", + "", + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'timestamp(message.updated_at) > now', + None, + None, + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(nested_pb2.NestedMessage.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = nested_pb2.NestedMessage() + obj.user_id.kafka_user_id = 'u1' + obj.updated_at.seconds = 1577836800 # 2020-01-01 UTC, before now + ser = ProtobufSerializer(nested_pb2.NestedMessage, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + with pytest.raises(SerializationError) as e: + ser(obj, ser_ctx) + assert isinstance(e.value.__cause__, RuleConditionError) + + +def test_proto_encryption(): + executor = FieldEncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule_conf = {'secret': 'mysecret'} + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT", + ["PII"], + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + # reset encrypted fields + assert obj.name != 'Kafka' + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client, rule_conf=rule_conf) + executor.executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_payload_encryption(): + executor = EncryptionExecutor.register_with_clock(FakeClock()) + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + ser_conf = {'auto.register.schemas': False, 'use.latest.version': True, 'use.deprecated.format': False} + rule_conf = {'secret': 'mysecret'} + rule = Rule( + "test-encrypt", + "", + RuleKind.TRANSFORM, + RuleMode.WRITEREAD, + "ENCRYPT_PAYLOAD", + None, + RuleParams({"encrypt.kek.name": "kek1", "encrypt.kms.type": "local-kms", "encrypt.kms.key.id": "mykey"}), + None, + None, + "ERROR,NONE", + False, + ) + client.register_schema( + _SUBJECT, + Schema(_schema_to_str(example_pb2.Author.DESCRIPTOR.file), "PROTOBUF", [], None, RuleSet(None, None, [rule])), + ) + obj = example_pb2.Author( + name='Kafka', id=123, picture=b'foobar', works=['The Castle', 'TheTrial'], oneof_string='oneof' + ) + ser = ProtobufSerializer(example_pb2.Author, client, conf=ser_conf, rule_conf=rule_conf) + dek_client = executor.client + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deser_conf = {'use.deprecated.format': False} + deser = ProtobufDeserializer(example_pb2.Author, deser_conf, client, rule_conf=rule_conf) + executor.client = dek_client + obj2 = deser(obj_bytes, ser_ctx) + assert obj == obj2 + + +def test_proto_jsonata_fully_compatible(): + rule1_to_2 = "$merge([$sift($, function($v, $k) {$k != 'size'}), {'height': $.'size'}])" + rule2_to_1 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'size': $.'height'}])" + rule2_to_3 = "$merge([$sift($, function($v, $k) {$k != 'height'}), {'length': $.'height'}])" + rule3_to_2 = "$merge([$sift($, function($v, $k) {$k != 'length'}), {'height': $.'length'}])" + + conf = {'url': _BASE_URL} + client = SchemaRegistryClient.new_client(conf) + + client.set_config(_SUBJECT, ServerConfig(compatibility_group='application.version')) + + client.register_schema( + _SUBJECT, + Schema( + _schema_to_str(widget_pb2.Widget.DESCRIPTOR.file), + "PROTOBUF", + [], + Metadata(None, MetadataProperties({"application.version": "v1"}), None), + None, + ), + ) + + rule1 = Rule( + "rule1", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule1_to_2, None, None, False + ) + rule2 = Rule( + "rule2", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule2_to_1, None, None, False + ) + client.register_schema( + _SUBJECT, + Schema( + _schema_to_str(newwidget_pb2.NewWidget.DESCRIPTOR.file), + "PROTOBUF", + [], + Metadata(None, MetadataProperties({"application.version": "v2"}), None), + RuleSet([rule1, rule2], None), + ), + ) + + rule3 = Rule( + "rule3", "", RuleKind.TRANSFORM, RuleMode.UPGRADE, "JSONATA", None, None, rule2_to_3, None, None, False + ) + rule4 = Rule( + "rule4", "", RuleKind.TRANSFORM, RuleMode.DOWNGRADE, "JSONATA", None, None, rule3_to_2, None, None, False + ) + client.register_schema( + _SUBJECT, + Schema( + _schema_to_str(newerwidget_pb2.NewerWidget.DESCRIPTOR.file), + "PROTOBUF", + [], + Metadata(None, MetadataProperties({"application.version": "v3"}), None), + RuleSet([rule3, rule4], None), + ), + ) + + obj = widget_pb2.Widget(name='alice', size=123, version=1) + obj2 = newwidget_pb2.NewWidget(name='alice', height=123, version=1) + obj3 = newerwidget_pb2.NewerWidget(name='alice', length=123, version=1) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v1'}, + 'use.deprecated.format': False, + } + ser = ProtobufSerializer(widget_pb2.Widget, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj, ser_ctx) + + deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v2'}, + 'use.deprecated.format': False, + } + ser = ProtobufSerializer(newwidget_pb2.NewWidget, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj2, ser_ctx) + + deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + ser_conf = { + 'auto.register.schemas': False, + 'use.latest.version': False, + 'use.latest.with.metadata': {'application.version': 'v3'}, + 'use.deprecated.format': False, + } + ser = ProtobufSerializer(newerwidget_pb2.NewerWidget, client, conf=ser_conf) + ser_ctx = SerializationContext(_TOPIC, MessageField.VALUE) + obj_bytes = ser(obj3, ser_ctx) + + deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3) + + +def deserialize_with_all_versions(client, ser_ctx, obj_bytes, obj, obj2, obj3): + deser_conf = {'use.latest.with.metadata': {'application.version': 'v1'}, 'use.deprecated.format': False} + deser = ProtobufDeserializer(widget_pb2.Widget, deser_conf, client) + newobj = deser(obj_bytes, ser_ctx) + assert obj.size == newobj.size + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v2'}, 'use.deprecated.format': False} + deser = ProtobufDeserializer(newwidget_pb2.NewWidget, deser_conf, client) + newobj = deser(obj_bytes, ser_ctx) + assert obj2.height == newobj.height + + deser_conf = {'use.latest.with.metadata': {'application.version': 'v3'}, 'use.deprecated.format': False} + deser = ProtobufDeserializer(newerwidget_pb2.NewerWidget, deser_conf, client) + newobj = deser(obj_bytes, ser_ctx) + assert obj3.length == newobj.length + + def test_associated_name_strategy_with_association(): """Test that AssociatedNameStrategy returns subject from association""" conf = {'url': _BASE_URL} diff --git a/tests/schema_registry/conftest.py b/tests/schema_registry/conftest.py index a50fed41c..cc91f84f6 100644 --- a/tests/schema_registry/conftest.py +++ b/tests/schema_registry/conftest.py @@ -57,6 +57,9 @@ "test_azure_aead.py", "test_azure_client.py", "test_azure_driver.py", + "test_cel_field_value_types.py", + "test_cel_message_transform.py", + "test_cel_null_avro_field.py", "test_cel_validator.py", "test_dlq_action.py", "test_encrypt_executor.py", @@ -64,20 +67,25 @@ "test_inline_tags.py", "test_proto_transform.py", "test_validate_message.py", + "test_variant_utils.py", "_async/test_avro.py", "_async/test_avro_serdes.py", "_async/test_avro_serdes_rules.py", + "_async/test_cel_avro_message_transform.py", "_async/test_config_rules.py", "_async/test_dlq_serdes.py", "_async/test_json_serdes_rules.py", + "_async/test_proto_serdes.py", "_async/test_proto_serdes_rules.py", "_async/test_validation_serdes.py", "_sync/test_avro.py", "_sync/test_avro_serdes.py", "_sync/test_avro_serdes_rules.py", + "_sync/test_cel_avro_message_transform.py", "_sync/test_config_rules.py", "_sync/test_dlq_serdes.py", "_sync/test_json_serdes_rules.py", + "_sync/test_proto_serdes.py", "_sync/test_proto_serdes_rules.py", "_sync/test_validation_serdes.py", ] diff --git a/tests/schema_registry/data/proto/value_type_rules.proto b/tests/schema_registry/data/proto/value_type_rules.proto new file mode 100644 index 000000000..7b0261d6d --- /dev/null +++ b/tests/schema_registry/data/proto/value_type_rules.proto @@ -0,0 +1,61 @@ +syntax = "proto3"; + +package tests; + +import "confluent/meta.proto"; +import "confluent/type/decimal.proto"; +import "google/protobuf/timestamp.proto"; + +// Inline rules on the two value types, at message and field level, so both can be exercised +// over an UNSET field. Mirrors the reference's own fixture so the clients are compared against +// the same rules. +message InlineValueTypes { + option (.confluent.message_meta) = { + rules: [ + {name: "msgDec", expr: "decimals.gt(this.amount, decimal('10.00'))"}, + {name: "msgTs", expr: "this.ts > timestamp('2000-01-01T00:00:00Z')"} + ] + }; + + .confluent.type.Decimal amount = 1 [(.confluent.field_meta) = { + rules: [{name: "fldDec", expr: "decimals.gt(this, decimal('10.00'))"}] + }]; + google.protobuf.Timestamp ts = 2 [(.confluent.field_meta) = { + rules: [{name: "fldTs", expr: "this > timestamp('2000-01-01T00:00:00Z')"}] + }]; + string label = 3; +} + +// A container message carrying BOTH inline rules and tags, so one fixture serves every way a +// rule can reach a value type inside a container: inline rules, message-level selection, a +// tagged CEL_FIELD rule, and a message-level transform. +message ValueTypeContainers { + option (.confluent.message_meta) = { + rules: [ + {name: "msgArr", expr: "decimals.gt(this.amounts[0], decimal('1.00'))"}, + {name: "msgNested", expr: "decimals.gt(this.nested.inner, decimal('1.00'))"} + ] + }; + + repeated .confluent.type.Decimal amounts = 1 [(.confluent.field_meta) = { + tags: ["AMOUNTS"], + rules: [{name: "fldArr", expr: "decimals.gt(this[0], decimal('1.00'))"}] + }]; + map amount_map = 2 [(.confluent.field_meta) = { + tags: ["AMOUNTMAP"] + }]; + ValueTypeNested nested = 3; + // Tagged so a rule can target a singular scalar: a singular condition still raises where a + // repeated one does not. + string label = 4 [(.confluent.field_meta) = { tags: ["LABEL"] }]; + // A tagged repeated scalar. The element type decides whether the list branch reaches + // transformLeaf or the value-type leaf path. + repeated string codes = 5 [(.confluent.field_meta) = { tags: ["CODES"] }]; +} + +message ValueTypeNested { + .confluent.type.Decimal inner = 1 [(.confluent.field_meta) = { + tags: ["INNER"], + rules: [{name: "fldInner", expr: "decimals.gt(this, decimal('1.00'))"}] + }]; +} diff --git a/tests/schema_registry/data/proto/value_type_rules_pb2.py b/tests/schema_registry/data/proto/value_type_rules_pb2.py new file mode 100644 index 000000000..2e3d5a8fb --- /dev/null +++ b/tests/schema_registry/data/proto/value_type_rules_pb2.py @@ -0,0 +1,72 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: value_type_rules.proto +# Protobuf Python Version: 7.35.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +from google.protobuf import timestamp_pb2 as google_dot_protobuf_dot_timestamp__pb2 + +import confluent_kafka.schema_registry.confluent.meta_pb2 as confluent_dot_meta__pb2 +import confluent_kafka.schema_registry.confluent.type.decimal_pb2 as confluent_dot_type_dot_decimal__pb2 + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\n\x16value_type_rules.proto\x12\x05tests\x1a\x14\x63onfluent/meta.proto\x1a\x1c\x63onfluent/type/decimal.proto\x1a\x1fgoogle/protobuf/timestamp.proto\"\xcf\x02\n\x10InlineValueTypes\x12[\n\x06\x61mount\x18\x01 \x01(\x0b\x32\x17.confluent.type.DecimalB2\x82\x44/\"-\n\x06\x66ldDec\x1a#decimals.gt(this, decimal(\'10.00\'))\x12^\n\x02ts\x18\x02 \x01(\x0b\x32\x1a.google.protobuf.TimestampB6\x82\x44\x33\"1\n\x05\x66ldTs\x1a(this > timestamp(\'2000-01-01T00:00:00Z\')\x12\r\n\x05label\x18\x03 \x01(\t:o\x82\x44l\"4\n\x06msgDec\x1a*decimals.gt(this.amount, decimal(\'10.00\'))\"4\n\x05msgTs\x1a+this.ts > timestamp(\'2000-01-01T00:00:00Z\')\"\xf2\x03\n\x13ValueTypeContainers\x12g\n\x07\x61mounts\x18\x01 \x03(\x0b\x32\x17.confluent.type.DecimalB=\x82\x44:\x1a\x07\x41MOUNTS\"/\n\x06\x66ldArr\x1a%decimals.gt(this[0], decimal(\'1.00\'))\x12M\n\namount_map\x18\x02 \x03(\x0b\x32).tests.ValueTypeContainers.AmountMapEntryB\x0e\x82\x44\x0b\x1a\tAMOUNTMAP\x12&\n\x06nested\x18\x03 \x01(\x0b\x32\x16.tests.ValueTypeNested\x12\x19\n\x05label\x18\x04 \x01(\tB\n\x82\x44\x07\x1a\x05LABEL\x12\x19\n\x05\x63odes\x18\x05 \x03(\tB\n\x82\x44\x07\x1a\x05\x43ODES\x1aI\n\x0e\x41mountMapEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12&\n\x05value\x18\x02 \x01(\x0b\x32\x17.confluent.type.Decimal:\x02\x38\x01:z\x82\x44w\"7\n\x06msgArr\x1a-decimals.gt(this.amounts[0], decimal(\'1.00\'))\"<\n\tmsgNested\x1a/decimals.gt(this.nested.inner, decimal(\'1.00\'))\"u\n\x0fValueTypeNested\x12\x62\n\x05inner\x18\x01 \x01(\x0b\x32\x17.confluent.type.DecimalB:\x82\x44\x37\x1a\x05INNER\".\n\x08\x66ldInner\x1a\"decimals.gt(this, decimal(\'1.00\'))b\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'value_type_rules_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_INLINEVALUETYPES'].fields_by_name['amount']._loaded_options = None + _globals['_INLINEVALUETYPES'].fields_by_name[ + 'amount' + ]._serialized_options = b'\202D/\"-\n\006fldDec\032#decimals.gt(this, decimal(\'10.00\'))' + _globals['_INLINEVALUETYPES'].fields_by_name['ts']._loaded_options = None + _globals['_INLINEVALUETYPES'].fields_by_name[ + 'ts' + ]._serialized_options = b'\202D3\"1\n\005fldTs\032(this > timestamp(\'2000-01-01T00:00:00Z\')' + _globals['_INLINEVALUETYPES']._loaded_options = None + _globals['_INLINEVALUETYPES']._serialized_options = ( + b'\202Dl\"4\n\006msgDec\032*decimals.gt(this.amount, decimal(\'10.00\'))\"4\n\005msgTs\032+this.ts > timestamp(\'2000-01-01T00:00:00Z\')' + ) + _globals['_VALUETYPECONTAINERS_AMOUNTMAPENTRY']._loaded_options = None + _globals['_VALUETYPECONTAINERS_AMOUNTMAPENTRY']._serialized_options = b'8\001' + _globals['_VALUETYPECONTAINERS'].fields_by_name['amounts']._loaded_options = None + _globals['_VALUETYPECONTAINERS'].fields_by_name[ + 'amounts' + ]._serialized_options = b'\202D:\032\007AMOUNTS\"/\n\006fldArr\032%decimals.gt(this[0], decimal(\'1.00\'))' + _globals['_VALUETYPECONTAINERS'].fields_by_name['amount_map']._loaded_options = None + _globals['_VALUETYPECONTAINERS'].fields_by_name['amount_map']._serialized_options = b'\202D\013\032\tAMOUNTMAP' + _globals['_VALUETYPECONTAINERS'].fields_by_name['label']._loaded_options = None + _globals['_VALUETYPECONTAINERS'].fields_by_name['label']._serialized_options = b'\202D\007\032\005LABEL' + _globals['_VALUETYPECONTAINERS'].fields_by_name['codes']._loaded_options = None + _globals['_VALUETYPECONTAINERS'].fields_by_name['codes']._serialized_options = b'\202D\007\032\005CODES' + _globals['_VALUETYPECONTAINERS']._loaded_options = None + _globals['_VALUETYPECONTAINERS']._serialized_options = ( + b'\202Dw\"7\n\006msgArr\032-decimals.gt(this.amounts[0], decimal(\'1.00\'))\"<\n\tmsgNested\032/decimals.gt(this.nested.inner, decimal(\'1.00\'))' + ) + _globals['_VALUETYPENESTED'].fields_by_name['inner']._loaded_options = None + _globals['_VALUETYPENESTED'].fields_by_name[ + 'inner' + ]._serialized_options = b'\202D7\032\005INNER\".\n\010fldInner\032\"decimals.gt(this, decimal(\'1.00\'))' + _globals['_INLINEVALUETYPES']._serialized_start = 119 + _globals['_INLINEVALUETYPES']._serialized_end = 454 + _globals['_VALUETYPECONTAINERS']._serialized_start = 457 + _globals['_VALUETYPECONTAINERS']._serialized_end = 955 + _globals['_VALUETYPECONTAINERS_AMOUNTMAPENTRY']._serialized_start = 758 + _globals['_VALUETYPECONTAINERS_AMOUNTMAPENTRY']._serialized_end = 831 + _globals['_VALUETYPENESTED']._serialized_start = 957 + _globals['_VALUETYPENESTED']._serialized_end = 1074 +# @@protoc_insertion_point(module_scope) diff --git a/tests/schema_registry/data/proto/value_types.proto b/tests/schema_registry/data/proto/value_types.proto new file mode 100644 index 000000000..79d8bbb93 --- /dev/null +++ b/tests/schema_registry/data/proto/value_types.proto @@ -0,0 +1,19 @@ +syntax = "proto3"; + +package tests; + +import "confluent/meta.proto"; +import "confluent/type/decimal.proto"; +import "confluent/type/variant.proto"; +import "google/protobuf/timestamp.proto"; + +// The three value types protobuf carries as messages, alongside a plain scalar and an +// optional field used to exercise absence. Tagged so a CEL_FIELD rule can be scoped to one +// field at a time. +message ValueTypes { + .confluent.type.Decimal amount = 1 [(.confluent.field_meta) = { tags: ["AMOUNT"] }]; + google.protobuf.Timestamp ts = 2 [(.confluent.field_meta) = { tags: ["TS"] }]; + .confluent.type.Variant data = 3 [(.confluent.field_meta) = { tags: ["DATA"] }]; + string label = 4 [(.confluent.field_meta) = { tags: ["LABEL"] }]; + int32 count = 5; +} diff --git a/tests/schema_registry/data/proto/value_types_pb2.py b/tests/schema_registry/data/proto/value_types_pb2.py new file mode 100644 index 000000000..1872ee6cb --- /dev/null +++ b/tests/schema_registry/data/proto/value_types_pb2.py @@ -0,0 +1,43 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: value_types.proto +# Protobuf Python Version: 7.35.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +from google.protobuf import timestamp_pb2 as google_dot_protobuf_dot_timestamp__pb2 + +import confluent_kafka.schema_registry.confluent.meta_pb2 as confluent_dot_meta__pb2 +import confluent_kafka.schema_registry.confluent.type.decimal_pb2 as confluent_dot_type_dot_decimal__pb2 +import confluent_kafka.schema_registry.confluent.type.variant_pb2 as confluent_dot_type_dot_variant__pb2 + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\n\x11value_types.proto\x12\x05tests\x1a\x14\x63onfluent/meta.proto\x1a\x1c\x63onfluent/type/decimal.proto\x1a\x1c\x63onfluent/type/variant.proto\x1a\x1fgoogle/protobuf/timestamp.proto\"\xcf\x01\n\nValueTypes\x12\x34\n\x06\x61mount\x18\x01 \x01(\x0b\x32\x17.confluent.type.DecimalB\x0b\x82\x44\x08\x1a\x06\x41MOUNT\x12/\n\x02ts\x18\x02 \x01(\x0b\x32\x1a.google.protobuf.TimestampB\x07\x82\x44\x04\x1a\x02TS\x12\x30\n\x04\x64\x61ta\x18\x03 \x01(\x0b\x32\x17.confluent.type.VariantB\t\x82\x44\x06\x1a\x04\x44\x41TA\x12\x19\n\x05label\x18\x04 \x01(\tB\n\x82\x44\x07\x1a\x05LABEL\x12\r\n\x05\x63ount\x18\x05 \x01(\x05\x62\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'value_types_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_VALUETYPES'].fields_by_name['amount']._loaded_options = None + _globals['_VALUETYPES'].fields_by_name['amount']._serialized_options = b'\202D\010\032\006AMOUNT' + _globals['_VALUETYPES'].fields_by_name['ts']._loaded_options = None + _globals['_VALUETYPES'].fields_by_name['ts']._serialized_options = b'\202D\004\032\002TS' + _globals['_VALUETYPES'].fields_by_name['data']._loaded_options = None + _globals['_VALUETYPES'].fields_by_name['data']._serialized_options = b'\202D\006\032\004DATA' + _globals['_VALUETYPES'].fields_by_name['label']._loaded_options = None + _globals['_VALUETYPES'].fields_by_name['label']._serialized_options = b'\202D\007\032\005LABEL' + _globals['_VALUETYPES']._serialized_start = 144 + _globals['_VALUETYPES']._serialized_end = 351 +# @@protoc_insertion_point(module_scope) diff --git a/tests/schema_registry/test_cel_field_value_types.py b/tests/schema_registry/test_cel_field_value_types.py new file mode 100644 index 000000000..1f1c8da18 --- /dev/null +++ b/tests/schema_registry/test_cel_field_value_types.py @@ -0,0 +1,220 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +``CEL_FIELD`` rules over protobuf decimal and timestamp fields. + +Avro carries these two as logical types on a primitive, so the field is a leaf and a field rule +reaches it. Protobuf carries them as messages, so the walk used to descend *past* the field and +transform ``value``/``scale`` or ``seconds``/``nanos`` one at a time - meaning a rule tagged for +the field never fired at all, and the record came back unchanged with no error. A silent no-op +is the worst of the three possible outcomes: the rule author gets no signal. + +This is the port of the JVM client's #4538 (``isCelLeafMessage``). Variant is deliberately not a +leaf - it is a record in Avro too, so skipping it is the behaviour that matches, and a variant +is reached with a message-level ``CEL`` rule instead. +""" + +from decimal import Decimal + +import pytest + +from confluent_kafka.schema_registry.common.protobuf import get_type, transform +from confluent_kafka.schema_registry.rules.cel.cel_field_executor import CelFieldExecutor +from confluent_kafka.schema_registry.schema_registry_client import Rule, RuleKind, RuleMode, Schema +from confluent_kafka.schema_registry.serde import FieldType, RuleContext, RuleError + +from .data.proto import value_type_rules_pb2, value_types_pb2 + +_SCHEMA = """syntax = "proto3"; +package tests; +message ValueTypes {} +""" + +# 0x04D2 = 1234 unscaled, i.e. 12.34 at scale 2. +_UNSCALED = bytes([0x04, 0xD2]) + + +def _message(): + msg = value_types_pb2.ValueTypes() + msg.amount.value = _UNSCALED + msg.amount.precision = 8 + msg.amount.scale = 2 + msg.ts.seconds = 1700000000 + msg.ts.nanos = 123000000 + msg.label = "hi" + return msg + + +def _run(expr, kind, tag, msg=None): + msg = msg if msg is not None else _message() + rule = Rule("r", None, kind, RuleMode.WRITE, "CEL_FIELD", [tag] if tag else None, None, expr, None, None, False) + ctx = RuleContext( + None, + None, + None, + Schema(_SCHEMA, "PROTOBUF"), + "t-value", + RuleMode.WRITE, + rule, + 0, + [rule], + {"tests.ValueTypes.amount": {"AMOUNT"}, "tests.ValueTypes.ts": {"TS"}, "tests.ValueTypes.data": {"DATA"}}, + None, + ) + ft = CelFieldExecutor().new_transform(ctx) + return transform(ctx, msg.DESCRIPTOR, msg, ft) + + +def _decimal_of(msg): + unscaled = int.from_bytes(msg.amount.value, "big", signed=True) + return Decimal(unscaled).scaleb(-msg.amount.scale) + + +def test_field_types_match_the_avro_counterpart(): + """The declared type is what makes CEL_FIELD apply at all: a RECORD is skipped outright.""" + fields = value_types_pb2.ValueTypes.DESCRIPTOR.fields_by_name + + assert get_type(fields["amount"]) == FieldType.BYTES + assert get_type(fields["ts"]) == FieldType.LONG + # Variant stays a record, as in Avro - not a leaf. + assert get_type(fields["data"]) == FieldType.RECORD + + +def test_decimal_condition_fires(): + """C4. Before the port this raised nothing because the rule never ran.""" + _run('decimals.gt(decimal(value), decimal("10.00"))', RuleKind.CONDITION, "AMOUNT") + + +def test_decimal_condition_fails_when_it_should(): + """The must-fail twin. Without it the test above would also pass if no rule ran at all - + which is exactly how the defect hid.""" + with pytest.raises(Exception): + _run('decimals.gt(decimal(value), decimal("1000.00"))', RuleKind.CONDITION, "AMOUNT") + + +def test_timestamp_condition_fires(): + _run('value > timestamp("2000-01-01T00:00:00Z")', RuleKind.CONDITION, "TS") + + +def test_timestamp_condition_fails_when_it_should(): + with pytest.raises(Exception): + _run('value > timestamp("2050-01-01T00:00:00Z")', RuleKind.CONDITION, "TS") + + +def test_decimal_transform_is_written_back(): + """C5. The rule returns a Python Decimal; it has to be encoded back into the message.""" + out = _run('decimals.add(decimal(value), decimal("1.00"))', RuleKind.TRANSFORM, "AMOUNT") + + assert _decimal_of(out) == Decimal("13.34") + assert out.amount.scale == 2 + # Not merely the original left alone, which is what the defect looked like. + assert out.amount.value != _UNSCALED + + +def test_timestamp_transform_is_written_back(): + out = _run('value + duration("60s")', RuleKind.TRANSFORM, "TS") + + assert out.ts.seconds == 1700000060 + assert out.ts.nanos == 123000000 + + +def test_identity_transform_round_trips(): + """The pass-through: the cheapest check that the encode inverts the decode exactly.""" + out = _run("value", RuleKind.TRANSFORM, "AMOUNT") + + assert _decimal_of(out) == Decimal("12.34") + assert out.amount.scale == 2 + + +def test_variant_is_still_skipped(): + """Variant is a record in both formats, so a field rule must not reach it. The rule below + would raise if it ran, so passing means it was skipped.""" + msg = _message() + msg.data.metadata = b"\x01\x01\x00\x04name" + msg.data.value = b"\x02\x01\x00\x00\x06\x15alice" + + out = _run('variants.type(value) == "not-a-type"', RuleKind.CONDITION, "DATA", msg) + + assert out.data.metadata == b"\x01\x01\x00\x04name" + + +def test_a_wrong_result_type_is_reported(): + """A rule returning something that is neither a decimal nor the message is a rule-authoring + mistake; it must be named rather than written back as a default.""" + with pytest.raises(RuleError, match="expected a decimal"): + _run('"not a decimal"', RuleKind.TRANSFORM, "AMOUNT") + + +# A repeated value-type field needs its rule's result rebuilt *per element*. The walk applies the +# rule to each element, so what comes back is a list of Decimals; only the singular case was +# rebuilt, and writing the raw list failed with "Expected a message object, but got Decimal(...)". +# So a field rule over a repeated decimal could not be written back at all (the reference answers `[2.11, 3.22]`). +_CONTAINER_SCHEMA = """syntax = "proto3"; +package tests; +message ValueTypeContainers {} +""" + + +def _container_message(): + msg = value_type_rules_pb2.ValueTypeContainers() + for unscaled in (111, 222): + d = msg.amounts.add() + d.value = unscaled.to_bytes(2, "big") + d.precision = 8 + d.scale = 2 + msg.amount_map["a"].value = (333).to_bytes(2, "big") + msg.amount_map["a"].precision = 8 + msg.amount_map["a"].scale = 2 + msg.label = "hi" + return msg + + +def _run_container(expr, tag): + msg = _container_message() + rule = Rule("r", None, RuleKind.TRANSFORM, RuleMode.WRITE, "CEL_FIELD", [tag], None, expr, None, None, False) + ctx = RuleContext( + None, None, None, Schema(_CONTAINER_SCHEMA, "PROTOBUF"), "t-value", RuleMode.WRITE, rule, 0, [rule], None, None + ) + ft = CelFieldExecutor().new_transform(ctx) + return transform(ctx, msg.DESCRIPTOR, msg, ft) + + +def _amounts(msg): + return [Decimal(int.from_bytes(d.value, "big", signed=True)).scaleb(-d.scale) for d in msg.amounts] + + +def test_repeated_decimal_transform_is_written_back_per_element(): + out = _run_container('decimals.add(decimal(value), decimal("1.00"))', "AMOUNTS") + + assert _amounts(out) == [Decimal("2.11"), Decimal("3.22")] + + +def test_repeated_decimal_identity_transform_round_trips(): + """The must-pass twin: an identity rule hands back the message it was given, and the + per-element rebuild has to accept that as readily as a computed decimal.""" + out = _run_container("value", "AMOUNTS") + + assert _amounts(out) == [Decimal("1.11"), Decimal("2.22")] + + +def test_a_tagged_map_field_is_left_alone(): + """The reference does *not* transform a map value through a tag on the map field - the tag + does not reach the entry's value leaf - so matching it means leaving the map unchanged.""" + out = _run_container('decimals.add(decimal(value), decimal("1.00"))', "AMOUNTMAP") + + unscaled = int.from_bytes(out.amount_map["a"].value, "big", signed=True) + assert Decimal(unscaled).scaleb(-out.amount_map["a"].scale) == Decimal("3.33") diff --git a/tests/schema_registry/test_cel_message_transform.py b/tests/schema_registry/test_cel_message_transform.py new file mode 100644 index 000000000..74eea48ba --- /dev/null +++ b/tests/schema_registry/test_cel_message_transform.py @@ -0,0 +1,958 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +Message-level ``CEL`` transforms over protobuf: the rule returns a map and the message is +rebuilt from it. + +Before this the executor returned the raw celpy map, which the protobuf serializer cannot +write - a decimal and a variant are messages in protobuf and celpy has no rendering for +either - so every message-level transform failed. The equivalent Avro path already worked +because an Avro record is a dict in this client and a celpy map is a dict subclass. + +The transform has **replace** semantics: the map is the new message, so an unnamed field is +dropped and a ``null`` clears its field. Those are covered here too, because they are the +part a rule author is most likely to be surprised by. +""" + +import math +from decimal import Decimal + +import pytest + +from confluent_kafka.schema_registry.confluent.type.variant_utils import Variant, parse_json +from confluent_kafka.schema_registry.rules.cel.cel_executor import CelExecutor +from confluent_kafka.schema_registry.schema_registry_client import Rule, RuleKind, RuleMode, Schema +from confluent_kafka.schema_registry.serde import RuleContext + +from .data.proto import value_types_pb2 + +_SCHEMA = """syntax = "proto3"; +package tests; +import "confluent/type/decimal.proto"; +import "confluent/type/variant.proto"; +import "google/protobuf/timestamp.proto"; +message ValueTypes { + .confluent.type.Decimal amount = 1; + google.protobuf.Timestamp ts = 2; + .confluent.type.Variant data = 3; + string label = 4; + int32 count = 5; +} +""" + +# 0x04D2 = 1234 unscaled, i.e. 12.34 at scale 2. +_UNSCALED = bytes([0x04, 0xD2]) + + +def _message(): + msg = value_types_pb2.ValueTypes() + msg.amount.value = _UNSCALED + msg.amount.precision = 8 + msg.amount.scale = 2 + msg.ts.seconds = 1700000000 + msg.ts.nanos = 123000000 + variant = parse_json('{"name":"alice"}') + msg.data.metadata = variant.metadata + msg.data.value = variant.value + msg.label = "hi" + msg.count = 7 + return msg + + +def _transform(expr, msg=None): + rule = Rule("r", None, RuleKind.TRANSFORM, RuleMode.WRITE, "CEL", None, None, expr, None, None, False) + ctx = RuleContext( + None, None, None, Schema(_SCHEMA, "PROTOBUF"), "t-value", RuleMode.WRITE, rule, 0, [rule], None, None + ) + return CelExecutor().transform(ctx, _message() if msg is None else msg) + + +def _decimal_of(msg): + unscaled = int.from_bytes(msg.amount.value, "big", signed=True) + return Decimal(unscaled).scaleb(-msg.amount.scale) + + +_ALL = ( + '"amount": message.amount, "ts": message.ts, ' + '"data": message.data, "label": message.label, "count": message.count' +) + + +def test_pass_through_returns_a_message_unchanged(): + """C6. An identity transform is the cheapest regression test for a write-back path: + it fails for any breakage in the plumbing, without depending on the computation.""" + result = _transform("{" + _ALL + "}") + + assert isinstance(result, value_types_pb2.ValueTypes) + assert _decimal_of(result) == Decimal("12.34") + assert result.ts.seconds == 1700000000 + assert result.ts.nanos == 123000000 + assert Variant(result.data.value, result.data.metadata).to_json() == '{"name":"alice"}' + assert result.label == "hi" + assert result.count == 7 + + +def test_computed_decimal_is_written_back(): + """C7, decimal. The rule returns a Python Decimal, which has to become a + confluent.type.Decimal message - unscaled bytes plus the scale.""" + result = _transform( + '{"amount": decimals.add(decimal(message.amount), decimal("1.00")), ' + '"ts": message.ts, "data": message.data, "label": message.label}' + ) + + assert _decimal_of(result) == Decimal("13.34") + assert result.amount.scale == 2 + # Not merely the original echoed back. + assert result.amount.value != _UNSCALED + + +def test_computed_timestamp_is_written_back(): + """C7, timestamp. The rule returns a datetime, which has to become a + google.protobuf.Timestamp - and keep its sub-second part.""" + result = _transform( + '{"amount": message.amount, "ts": message.ts + duration("60s"), ' + '"data": message.data, "label": message.label}' + ) + + assert result.ts.seconds == 1700000060 + assert result.ts.nanos == 123000000 + + +def test_computed_variant_is_written_back(): + """C7, variant. The rule returns a Variant, which has to become a + confluent.type.Variant message. Asserted through the decoded JSON rather than the + metadata bytes: metadata holds the field *names*, so {"name":"alice"} and + {"name":"bob"} share it and comparing metadata would prove nothing.""" + result = _transform( + '{"amount": message.amount, "ts": message.ts, ' + '"data": variants.parseJson("{\\"name\\":\\"bob\\"}"), "label": message.label}' + ) + + assert Variant(result.data.value, result.data.metadata).to_json() == '{"name":"bob"}' + + +def test_scalar_can_be_replaced(): + result = _transform('{"label": "changed", "count": 9}') + + assert result.label == "changed" + assert result.count == 9 + + +def test_a_field_the_rule_does_not_name_is_dropped(): + """Replace semantics, and the consequence most likely to surprise: a rule naming only + the field it changes discards everything else. Intended, but silent on protobuf - + proto3 has no required fields, so nothing catches it.""" + result = _transform('{"label": "changed"}') + + assert result.label == "changed" + assert not result.HasField("amount") + assert not result.HasField("ts") + assert not result.HasField("data") + assert result.count == 0 + + +def test_null_clears_a_field(): + """The idiom for preserving absence across a transform that echoes a field: + `has(x) ? x : null`. Without a null arm there would be no way to express it.""" + result = _transform('{"amount": null, "ts": message.ts, "data": message.data, "label": message.label}') + + assert not result.HasField("amount") + assert result.HasField("ts") + assert result.label == "hi" + + +def test_echoing_an_absent_field_materialises_it(): + """The other face of replace: reading an absent field produces its default, so echoing + it writes that default back and `has()` flips from False to True. This documents the + behaviour rather than endorsing it - `has(x) ? x : null` is the way to avoid it.""" + absent = value_types_pb2.ValueTypes() + absent.label = "hi" + assert not absent.HasField("amount") + + echoed = _transform('{"amount": message.amount, "label": message.label}', absent) + assert echoed.HasField("amount") + + guarded = _transform('{"amount": has(message.amount) ? message.amount : null, "label": message.label}', absent) + assert not guarded.HasField("amount") + + +def test_condition_rules_are_unaffected(): + """A CONDITION returns a bool, which must not be run through the message rebuild.""" + rule = Rule( + "r", + None, + RuleKind.CONDITION, + RuleMode.WRITE, + "CEL", + None, + None, + 'decimals.gt(message.amount, decimal("10.00"))', + None, + None, + False, + ) + ctx = RuleContext( + None, None, None, Schema(_SCHEMA, "PROTOBUF"), "t-value", RuleMode.WRITE, rule, 0, [rule], None, None + ) + + assert CelExecutor().transform(ctx, _message()) is True + + +def _wrapper_descriptor(): + """A message with one field per protobuf wrapper type, plus a Duration. + + Built at runtime rather than added to value_types.proto so this needs no regenerated + ``_pb2`` fixture (the checked-in ones are deliberately free of protoc's runtime-version + gate, which a regeneration would reintroduce). + """ + from google.protobuf import descriptor_pb2, descriptor_pool, duration_pb2, message_factory, wrappers_pb2 + + fdp = descriptor_pb2.FileDescriptorProto() + fdp.name, fdp.package, fdp.syntax = "wrappers_probe.proto", "tests.wrap", "proto3" + fdp.dependency.extend(["google/protobuf/wrappers.proto", "google/protobuf/duration.proto"]) + msg = fdp.message_type.add() + msg.name = "Wrapped" + types = [ + "StringValue", + "BytesValue", + "Int32Value", + "Int64Value", + "UInt32Value", + "UInt64Value", + "FloatValue", + "DoubleValue", + "BoolValue", + "Duration", + ] + for number, type_name in enumerate(types, start=1): + field = msg.field.add() + field.name = type_name.lower() + field.number = number + field.type = descriptor_pb2.FieldDescriptorProto.TYPE_MESSAGE + field.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + field.type_name = ".google.protobuf." + type_name + field.json_name = type_name.lower() + + pool = descriptor_pool.DescriptorPool() + for dep in (wrappers_pb2.DESCRIPTOR, duration_pb2.DESCRIPTOR): + proto = descriptor_pb2.FileDescriptorProto() + dep.CopyToProto(proto) + pool.Add(proto) + pool.Add(fdp) + desc = pool.FindMessageTypeByName("tests.wrap.Wrapped") + return desc, message_factory.GetMessageClass(desc) + + +def _wrapped_message(): + _, cls = _wrapper_descriptor() + msg = cls() + msg.stringvalue.value = "hello" + msg.bytesvalue.value = b"\x01\x02" + msg.int32value.value = 7 + msg.int64value.value = 2**40 + msg.uint32value.value = 9 + msg.uint64value.value = 2**40 + 1 + msg.floatvalue.value = 1.5 + msg.doublevalue.value = 2.25 + msg.boolvalue.value = True + msg.duration.seconds, msg.duration.nanos = 3, 500000000 + return msg + + +def test_wrappers_and_duration_survive_an_identity_transform(): + """The CEL binding unwraps a wrapper to the scalar it holds and a Duration to a CEL + duration, so the write-back has to put them back. It only handled Decimal/Timestamp/ + Variant/mapping, so every one of these fields came back **empty** - silent data loss on + an identity transform. The JVM gets this right for free: its message-level write-back + goes through protobuf JSON, whose parser reads "hello" into a StringValue and "3s" into a + Duration. + """ + from confluent_kafka.schema_registry.rules.cel.constraints import _msg_to_cel + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + original = _wrapped_message() + assert convert(_msg_to_cel(original), original) == original + + +# A CEL timestamp is a datetime and a CEL duration a timedelta, both microsecond-resolution, so +# a nanos field cannot survive the conversion. It does not have to: an echoed value is copied +# from the message it was read from. Without that, a rule rewriting some *other* field turned +# nanos 1 into 0 - while the decimal in the same message came back byte-identical, because its +# binding already kept the source. These values are chosen to discriminate: the tests above use +# 123000000 and 500000000, whole microsecond counts that cannot detect the truncation. +@pytest.mark.parametrize("nanos", [123456789, 1, 999999999]) +def test_an_echoed_timestamp_keeps_its_nanos(nanos): + msg = _message() + msg.ts.nanos = nanos + + # Identity, and a rule that rewrites a sibling field and merely passes ts along. + identity = _transform("{" + _ALL + "}", msg) + sibling = _transform("{" + _ALL.replace('"label": message.label', '"label": message.label + "!"') + "}", msg) + + assert identity.ts.nanos == nanos + assert sibling.ts.nanos == nanos + assert sibling.label == "hi!" + + +@pytest.mark.parametrize("nanos", [123456789, 1, 999999999]) +def test_an_echoed_duration_keeps_its_nanos(nanos): + from confluent_kafka.schema_registry.rules.cel.constraints import _msg_to_cel + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + original = _wrapped_message() + original.duration.seconds, original.duration.nanos = 3, nanos + + assert convert(_msg_to_cel(original), original).duration.nanos == nanos + + +# ...and a *computed* timestamp still lands on the microsecond ceiling, which is inherent to +# datetime and already documented for `timestamp(x, 9)` and `string(ts)`. Pinned so the fix +# above is not mistaken for nanosecond arithmetic. +def test_a_computed_timestamp_keeps_the_microsecond_ceiling(): + msg = _message() + msg.ts.nanos = 123456789 + + result = _transform("{" + _ALL.replace('"ts": message.ts', '"ts": message.ts + duration("0s")') + "}", msg) + + assert result.ts.nanos == 123456000 + + +def test_negative_duration_keeps_matching_signs(): + """A Duration's seconds and nanos must share a sign; timedelta normalises microseconds to + be non-negative (-3.5s is days=-1, seconds=86396, microseconds=500000), so splitting it + by floor division produced seconds=-3 with nanos=+499000000.""" + from confluent_kafka.schema_registry.rules.cel.constraints import _msg_to_cel + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _wrapper_descriptor() + original = cls() + original.duration.seconds, original.duration.nanos = -3, -500000000 + result = convert(_msg_to_cel(original), original) + assert (result.duration.seconds, result.duration.nanos) == (-3, -500000000) + + +# A wrapper's `value` field had its own copy of the scalar conversions, and the copy was the +# unguarded one: an Int32Value took 1.9 as 1 and a BoolValue took the string "false" as *true*, +# while the identical plain fields refused both. The JVM draws no such distinction - +# JsonFormat's parseWrapperFieldValue hands the value to the same parseFieldValue a plain field +# goes through, so the accept/reject sets are identical. Measured against protobuf-java 4.35.1 +# with each wrapper as a nested field: +# Int32Value <- 2 / 2.0 -> 2; <- 1.9, 2147483648, true -> REJECT +# BoolValue <- true -> true; <- 0 -> REJECT "Invalid bool value: 0" +# BytesValue <- 5 -> REJECT; FloatValue <- 1.0e40 -> REJECT "Out of range float value" +# DoubleValue <- 3 -> 3.0; <- true -> REJECT "Not a double value: true" +def test_a_wrapper_field_is_narrowed_like_a_plain_scalar(): + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _wrapper_descriptor() + original = cls() + + # Exact conversions still reach the wrapper. + assert convert({"int32value": celtypes.DoubleType(2.0)}, original).int32value.value == 2 + assert convert({"doublevalue": celtypes.IntType(3)}, original).doublevalue.value == 3.0 + assert convert({"boolvalue": celtypes.BoolType(True)}, original).boolvalue.value is True + assert convert({"stringvalue": celtypes.StringType("ok")}, original).stringvalue.value == "ok" + assert convert({"bytesvalue": celtypes.BytesType(b"ab")}, original).bytesvalue.value == b"ab" + + # The same rejections a plain field of that type makes. + for field, value in [ + ("int32value", celtypes.DoubleType(1.9)), + ("int32value", celtypes.IntType(2**31)), + ("int32value", celtypes.BoolType(True)), + ("boolvalue", celtypes.IntType(0)), + ("boolvalue", celtypes.StringType("false")), + ("stringvalue", celtypes.IntType(1)), + ("bytesvalue", celtypes.IntType(5)), + ("doublevalue", celtypes.BoolType(True)), + ("floatvalue", celtypes.DoubleType(1e40)), + ("uint32value", celtypes.IntType(-1)), + ]: + with pytest.raises(ValueError): + convert({field: value}, original) + + +def test_message_level_decimal_sets_precision_like_the_field_level_writer(): + """Java's ProtobufResultWriter sets precision from the value (``dec.precision()``); this + writer left it at zero, so the same computed decimal produced different + confluent.type.Decimal bytes depending on the rule's scope. Safe to set here because this + writer never rescales, so len(digits) is the digit count of the unscaled value written. + """ + from confluent_kafka.schema_registry.common.protobuf import set_decimal_message + from confluent_kafka.schema_registry.confluent.type import decimal_pb2 + from confluent_kafka.schema_registry.confluent.type.decimal_utils import to_proto_decimal + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import _set_decimal + + # (value, java BigDecimal.precision(), java scale()) + for text, precision, scale in [ + ("12.34", 4, 2), + ("12.3400", 6, 4), + ("1E+3", 1, -3), + ("0.00", 1, 2), + ("100", 3, 0), + ]: + message_level = decimal_pb2.Decimal() + _set_decimal(message_level, Decimal(text)) + field_level = decimal_pb2.Decimal() + set_decimal_message(field_level, Decimal(text)) + # to_proto_decimal is the third writer; all three must agree byte for byte, so a + # consumer that honours precision cannot see one value two ways. + standalone = to_proto_decimal(Decimal(text)) + + assert (message_level.precision, message_level.scale) == (precision, scale), text + assert message_level.SerializeToString() == field_level.SerializeToString(), text + assert standalone.SerializeToString() == field_level.SerializeToString(), text + + +# `mul` is no longer guarded on its result's shape - each library's own exponent range is +# delegated and documented, and multiplication is measurably cheap at any width. So a value +# whose scale no int32 can carry now reaches the writer instead of being refused by the +# operator: two 1e2147483647 operands multiply exactly, and need a scale of -4294967294. +# Assigning that raises a bare `ValueError: Value out of range` from the protobuf runtime; +# the writer names the value and the field instead. +def test_a_scale_that_does_not_fit_int32_is_refused_at_the_wire(): + from confluent_kafka.schema_registry.confluent.type import decimal_pb2 + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import _set_decimal + + for text in ["1e4294967294", "1e-4294967294"]: + with pytest.raises(ValueError, match="does not fit the int32 scale field"): + _set_decimal(decimal_pb2.Decimal(), Decimal(text)) + + # The boundary itself still writes, from both sides. + for text, scale in [("1e-2147483647", 2147483647), ("1e2147483647", -2147483647)]: + target = decimal_pb2.Decimal() + _set_decimal(target, Decimal(text)) + assert target.scale == scale, text + + +# The coefficient, not the scale, is what actually bounds what this client can write. The wire +# form is the unscaled integer in base 256, and decimal <-> binary radix conversion is +# quadratic: CPython caps str <-> int at 4300 digits for exactly that reason (measured in the +# C++ sibling, whose own codec takes 0.04 s at 10**4 digits, 4.2 s at 10**5 and ~420 s at +# 10**6, with mpdecimal's mpd_qexport_u32 only about 10x better and the same quadratic shape). +# +# So the cap is pre-existing - `int("9" * 5000)` has always raised - and reached callers as +# CPython's "Exceeds the limit (4300 digits) for integer string conversion", naming neither +# the decimal nor the field. It now names both. +def test_a_coefficient_past_what_can_be_encoded_is_refused(): + from confluent_kafka.schema_registry.confluent.type import decimal_pb2 + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import _set_decimal + + for digits in [4301, 5000, 100000]: + with pytest.raises(ValueError, match="past the 4300 this client can encode"): + _set_decimal(decimal_pb2.Decimal(), Decimal("9" * digits)) + + # The boundary itself still writes, and so does anything narrower. + for digits in [1, 38, 4300]: + target = decimal_pb2.Decimal() + _set_decimal(target, Decimal("9" * digits)) + assert target.precision == digits + + +# The field-level twin of the above. `set_decimal_message` is the write-back for a +# confluent.type.Decimal *field*, and it had the unhelpful version of the same failure: the +# `int(...)` inside it is a str -> int conversion, so CPython raised "Exceeds the limit (4300 +# digits) for integer string conversion" naming neither the decimal nor the field. Both paths +# now report it the same way, from one shared constant - there were two constants called +# `_MAX_COEFFICIENT_DIGITS` in this client with different values. +def test_the_field_level_writer_reports_a_wide_coefficient_too(): + from confluent_kafka.schema_registry.common.protobuf import set_decimal_message + from confluent_kafka.schema_registry.confluent.type import decimal_pb2 + + for digits in [4301, 5000, 100000]: + with pytest.raises(ValueError, match="past the 4300 this client can encode"): + set_decimal_message(decimal_pb2.Decimal(), Decimal("9" * digits)) + + # And the two writers agree exactly, boundary included, so which path produced a decimal + # cannot change whether it is accepted. + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import _set_decimal + + for digits in [1, 38, 4300]: + field_level = decimal_pb2.Decimal() + set_decimal_message(field_level, Decimal("9" * digits)) + message_level = decimal_pb2.Decimal() + _set_decimal(message_level, Decimal("9" * digits)) + assert field_level.SerializeToString() == message_level.SerializeToString() + + +def _scalar_descriptor(): + """A message with one field of each scalar kind, built at runtime.""" + from google.protobuf import descriptor_pb2, descriptor_pool, message_factory + + fdp = descriptor_pb2.FileDescriptorProto() + fdp.name, fdp.package, fdp.syntax = "int_coercion.proto", "tests.ic", "proto3" + choices = fdp.enum_type.add() + choices.name = "Choice" + choices.value.add(name="ZERO", number=0) + choices.value.add(name="ONE", number=1) + msg = fdp.message_type.add() + msg.name = "Ints" + spec = [ + ("i32", 1, descriptor_pb2.FieldDescriptorProto.TYPE_INT32, 1), + ("u32", 2, descriptor_pb2.FieldDescriptorProto.TYPE_UINT32, 1), + ("dbl", 3, descriptor_pb2.FieldDescriptorProto.TYPE_DOUBLE, 1), + ("codes", 4, descriptor_pb2.FieldDescriptorProto.TYPE_STRING, 3), + ("text", 5, descriptor_pb2.FieldDescriptorProto.TYPE_STRING, 1), + ("flag", 6, descriptor_pb2.FieldDescriptorProto.TYPE_BOOL, 1), + ("blob", 7, descriptor_pb2.FieldDescriptorProto.TYPE_BYTES, 1), + ("choice", 8, descriptor_pb2.FieldDescriptorProto.TYPE_ENUM, 1), + ("flt", 9, descriptor_pb2.FieldDescriptorProto.TYPE_FLOAT, 1), + ("u64", 10, descriptor_pb2.FieldDescriptorProto.TYPE_UINT64, 1), + ] + for name, number, ftype, label in spec: + field = msg.field.add() + field.name, field.number, field.type, field.label = name, number, ftype, label + field.json_name = name + if ftype == descriptor_pb2.FieldDescriptorProto.TYPE_ENUM: + field.type_name = ".tests.ic.Choice" + + pool = descriptor_pool.DescriptorPool() + pool.Add(fdp) + desc = pool.FindMessageTypeByName("tests.ic.Ints") + return desc, message_factory.GetMessageClass(desc) + + +# int() silently truncated, so a CEL double of 1.9 landed in an int32 field as 1 and a +# fractional Decimal lost its fraction. The JVM's message-level write-back goes through a +# protobuf JSON parse, which refuses a non-integral value and range-checks the result. +# Measured against protobuf-java 4.35.1: +# Int32Value <- 1.9 -> REJECT "Not an int32 value: 1.9" +# Int32Value <- 2.0 -> 2 +# Int32Value <- 2147483648 -> REJECT "Not an int32 value" +def test_integer_fields_reject_non_integral_and_out_of_range(): + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + + def write(field, value): + return convert({field: value}, original) + + # Integral values are accepted, whatever their Python type. + assert write("i32", 2.0).i32 == 2 + assert write("i32", 2).i32 == 2 + assert write("i32", Decimal("3")).i32 == 3 + # A double field still takes a fractional value. + assert write("dbl", 1.9).dbl == 1.9 + + for value in (1.9, Decimal("1.5")): + with pytest.raises(ValueError, match="non-integral"): + write("i32", value) + with pytest.raises(ValueError, match="out of range"): + write("i32", 2**31) + with pytest.raises(ValueError, match="out of range"): + write("u32", -1) + # Both boolean spellings. bool is an int subclass in Python, and celtypes.BoolType + # subclasses int rather than bool - so guarding only `bool` caught the case CEL never + # produces while a real CEL `true` was written as 1. protobuf JSON refuses true for an + # integer field ("Not an int32 value: true"). + from celpy import celtypes + + for value in (True, celtypes.BoolType(True), celtypes.BoolType(False)): + with pytest.raises(ValueError, match="bool"): + write("i32", value) + + +# The other three scalar arms narrowed unconditionally, so a wrong-typed result was accepted +# and silently changed meaning: bytes(5) fabricated five NUL bytes, bool("false") wrote true, +# float(True) wrote 1.0, and str() turned any value at all into a string field's text. +# +# Every rejection below is one protobuf's own JSON parser makes, which is what the JVM's +# write-back parses the result map with. Measured against protobuf-java 4.35.1: +# bool <- 0, "TRUE", "" -> REJECT "Invalid bool value" +# bytes <- 5, [97, 98] -> REJECT +# float <- true -> REJECT "Not a double value: true" +# int <- 1.9, true -> REJECT "Not an int32 value" +# +# That parser is also lenient the other way - it stringifies a number into a string field, +# reads "true"/"false" as a bool and a numeric string as a number - and these tests pin the +# decision *not* to follow it. Those coercions only exist because its input crossed a JSON +# transport, which this writer does not; each one turns a rule-authoring mistake into data. +def test_string_fields_take_only_a_string(): + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + assert convert({"text": celtypes.StringType("ok")}, original).text == "ok" + + # The JVM would stringify each of these (1 -> "1", true -> "true"). Deliberately not. + for value in ( + celtypes.IntType(1), + celtypes.DoubleType(1.5), + celtypes.BoolType(True), + celtypes.BytesType(b"ab"), + celtypes.ListType([celtypes.IntType(1)]), + ): + with pytest.raises(ValueError, match="to string field 'text'"): + convert({"text": value}, original) + + +def test_bool_fields_do_not_use_python_truthiness(): + """The string "false" is the case that matters: truthiness wrote *true* for it.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + assert convert({"flag": celtypes.BoolType(True)}, original).flag is True + assert convert({"flag": celtypes.BoolType(False)}, original).flag is False + + # A number is refused by the JVM too, so bool(0) accepted what it rejects. + for value in (celtypes.IntType(0), celtypes.IntType(1), celtypes.DoubleType(1.0)): + with pytest.raises(ValueError, match="to bool field 'flag'"): + convert({"flag": value}, original) + # "true"/"false" are the JVM's own lenient spellings, not followed here; the rest it + # rejects outright. + for value in ( + celtypes.StringType("true"), + celtypes.StringType("false"), + celtypes.StringType("TRUE"), + celtypes.StringType("yes"), + celtypes.StringType(""), + ): + with pytest.raises(ValueError, match="to bool field 'flag'"): + convert({"flag": value}, original) + + +def test_bytes_fields_take_only_a_byte_string(): + """bytes(5) fabricates five NUL bytes out of a number the JVM refuses.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + assert convert({"blob": celtypes.BytesType(b"ab")}, original).blob == b"ab" + + with pytest.raises(ValueError, match="to bytes field 'blob'"): + convert({"blob": celtypes.IntType(5)}, original) + for value in (celtypes.BoolType(True), celtypes.ListType([celtypes.IntType(97), celtypes.IntType(98)])): + with pytest.raises(ValueError, match="to bytes field 'blob'"): + convert({"blob": value}, original) + # The JVM base64-decodes a string here, because base64 is how bytes cross its JSON + # transport. This writer builds against the descriptor, so a CEL string is text that was + # never encoded and is not reinterpreted as bytes. + with pytest.raises(ValueError, match="to bytes field 'blob'"): + convert({"blob": celtypes.StringType("YWI=")}, original) + + +def test_float_fields_take_only_a_number(): + """A bool gets the same guard an integer field gives it, and so does a numeric string.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + # Every numeric kind a rule can compute still reaches a float field. + assert convert({"dbl": celtypes.DoubleType(1.5)}, original).dbl == 1.5 + assert convert({"dbl": celtypes.IntType(3)}, original).dbl == 3.0 + assert convert({"dbl": celtypes.UintType(3)}, original).dbl == 3.0 + assert convert({"dbl": Decimal("1.5")}, original).dbl == 1.5 + + for value in (True, celtypes.BoolType(True), celtypes.BoolType(False)): + with pytest.raises(ValueError, match="bool"): + convert({"dbl": value}, original) + # "1.5" and "NaN" are the JVM's lenient numeric strings, not followed here. + for value in (celtypes.StringType("1.5"), celtypes.StringType("NaN"), celtypes.StringType("abc")): + with pytest.raises(ValueError, match="to float field 'dbl'"): + convert({"dbl": value}, original) + + +def test_a_float_field_range_checks_the_narrowing(): + """CEL has one floating type, so a `float` field is a narrowing that can overflow. + float(1e40) gave inf; the JVM says "Out of range float value: 1.0e40". The 1e-6 slack and + the pass-through for NaN/infinity are both JsonFormat.parseFloat's own behaviour.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + # approx because the value comes back as its float32 self, not the double that went in. + assert convert({"flt": celtypes.DoubleType(3.4028235e38)}, original).flt == pytest.approx(3.4028235e38) + assert convert({"flt": celtypes.DoubleType(1.5)}, original).flt == 1.5 + assert math.isnan(convert({"flt": celtypes.DoubleType(float("nan"))}, original).flt) + assert math.isinf(convert({"flt": celtypes.DoubleType(float("inf"))}, original).flt) + + for value in (celtypes.DoubleType(1e40), celtypes.DoubleType(-1e40)): + with pytest.raises(ValueError, match="out of range float value"): + convert({"flt": value}, original) + # A double field takes the same value: only the 32-bit narrowing is range-checked. + assert convert({"dbl": celtypes.DoubleType(1e40)}, original).dbl == 1e40 + + +# Overflow has to be judged on the *source*, not the result. `float()` saturates a finite but +# too-large value to an infinity, and a range check on the result reads that infinity as one the +# rule asked for and lets it through - so Decimal("1e1000") was written as inf, and a double +# field had no range check at all. A wide Python int is the same case reported differently: +# float(10**400) raises OverflowError, which escaped as a raw Python exception rather than a +# rule error. An *explicitly* non-finite value does pass, because protobuf JSON has canonical +# spellings for those. Measured against protobuf-java 4.35.1: +# +# double <- 1e308 1.0E308 +# double <- 1e309, 1e1000, -1e1000 REJECT "Out of range double value" +# double <- "Infinity", "-Infinity", "NaN" accepted as-is +# float <- 1e39, 1e1000 REJECT "Out of range float value" +# float <- "Infinity", "NaN" accepted as-is +def test_a_finite_value_that_overflows_is_refused(): + from decimal import Decimal as D + + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + + # The double field had no range check at all, so all four of these were written as inf. + for value in (D("1e1000"), D("-1e1000"), D("1e309"), D("-1e309")): + with pytest.raises(ValueError, match="out of range value for float field 'dbl'"): + convert({"dbl": value}, original) + # The float field's check was evaded by the saturation itself. + for value in (D("1e1000"), D("1e39")): + with pytest.raises(ValueError, match="out of range"): + convert({"flt": value}, original) + # A Python int wider than a double raised OverflowError, not a rule error. + with pytest.raises(ValueError, match="out of range value for float field 'dbl'"): + convert({"dbl": 10**400}, original) + + # The widest value that still fits, so the guard cannot be off by an order of magnitude. + assert convert({"dbl": D("1e308")}, original).dbl == 1e308 + assert convert({"dbl": celtypes.IntType(3)}, original).dbl == 3.0 + + +def test_an_explicitly_non_finite_value_still_passes(): + """The JVM's parser takes protobuf JSON's canonical "Infinity"/"-Infinity"/"NaN", so a + rule that computes one deliberately is not an overflow.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + assert math.isinf(convert({"dbl": celtypes.DoubleType(float("inf"))}, original).dbl) + assert math.isinf(convert({"dbl": celtypes.DoubleType(float("-inf"))}, original).dbl) + assert math.isnan(convert({"dbl": celtypes.DoubleType(float("nan"))}, original).dbl) + assert math.isinf(convert({"flt": celtypes.DoubleType(float("inf"))}, original).flt) + assert math.isnan(convert({"flt": celtypes.DoubleType(float("nan"))}, original).flt) + # And a Decimal cannot be non-finite without saying so either. + from decimal import Decimal as D + + assert math.isnan(convert({"dbl": D("NaN")}, original).dbl) + + +# The same class in the integer arm, found while checking the above: `int()` was called before +# the range check, so int(Decimal("1e100000000")) spent minutes building a hundred million +# digits to reach a rejection its magnitude already settled. The JVM reports that off the token +# without building the number. The timing bound is generous: the fixed path is instant. +def test_an_out_of_range_integer_is_refused_without_building_it(): + import time + from decimal import Decimal as D + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + start = time.monotonic() + for value in (D("1e100000000"), D("-1e100000000"), D("1e30")): + with pytest.raises(ValueError, match="out of range"): + convert({"i32": value}, original) + assert time.monotonic() - start < 1.0 + # A non-finite Decimal is not an integer either. + for value in (D("NaN"), D("Infinity")): + with pytest.raises(ValueError): + convert({"i32": value}, original) + # And the values that do fit still convert. + assert convert({"i32": D("2")}, original).i32 == 2 + assert convert({"u64": D("18446744073709551615")}, original).u64 == 2**64 - 1 + + +def test_an_enum_still_takes_a_symbol_name(): + """Not a coercion: a name is protobuf JSON's canonical enum form and CEL has no enum + type, so a string is the only way a rule can name a symbol. The JVM accepts it too.""" + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + assert convert({"choice": celtypes.StringType("ONE")}, original).choice == 1 + assert convert({"choice": celtypes.IntType(1)}, original).choice == 1 + with pytest.raises(ValueError, match="bool"): + convert({"choice": celtypes.BoolType(True)}, original) + + +def _oneof_descriptor(): + """A message with a two-member oneof and a snake_case field, built at runtime.""" + from google.protobuf import descriptor_pb2, descriptor_pool, message_factory + + fdp = descriptor_pb2.FileDescriptorProto() + fdp.name, fdp.package, fdp.syntax = "oneof_probe.proto", "tests.oo", "proto3" + msg = fdp.message_type.add() + msg.name = "Choice" + msg.oneof_decl.add(name="choice") + field_type = descriptor_pb2.FieldDescriptorProto + msg.field.add( + name="a", number=1, type=field_type.TYPE_INT32, label=field_type.LABEL_OPTIONAL, oneof_index=0, json_name="a" + ) + msg.field.add( + name="b", number=2, type=field_type.TYPE_INT32, label=field_type.LABEL_OPTIONAL, oneof_index=0, json_name="b" + ) + msg.field.add( + name="total_amount", + number=3, + type=field_type.TYPE_INT32, + label=field_type.LABEL_OPTIONAL, + json_name="totalAmount", + ) + + pool = descriptor_pool.DescriptorPool() + pool.Add(fdp) + desc = pool.FindMessageTypeByName("tests.oo.Choice") + return desc, message_factory.GetMessageClass(desc) + + +# Two result entries can name the same slot, and applying both left the outcome to the order the +# rule happened to write them in: `{a: 1, b: 2}` kept b and `{b: 2, a: 1}` kept a, both reported +# as a successful transform. JsonFormat refuses both shapes, and the two have *opposite* null +# handling - mergeField's hasField test sits before its null early-return, mergeOneofField's +# after. Measured against protobuf-java 4.35.1: +# +# {"a":1,"b":2} / {"b":2,"a":1} REJECT "...belonging to the same oneof has already +# been set" +# {"a":1,"b":null} / {"a":null,"b":2} accept - a null is treated as absent +# {"total_amount":1,"totalAmount":2} REJECT "Field p.M.total_amount has already been set." +# {"total_amount":1,"totalAmount":null} REJECT - the same, because the value was already set +# {"total_amount":null,"totalAmount":null} accept - neither null set anything +@pytest.mark.parametrize( + "values", + [ + {"a": 1, "b": 2}, + {"b": 2, "a": 1}, + {"a": 1, "b": 2, "total_amount": 3}, + ], +) +def test_two_members_of_one_oneof_are_rejected(values): + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _oneof_descriptor() + celified = {k: celtypes.IntType(v) for k, v in values.items()} + with pytest.raises(ValueError, match="more than one member of oneof"): + convert(celified, cls()) + + +@pytest.mark.parametrize( + "values", + [ + {"total_amount": 1, "totalAmount": 2}, + {"total_amount": 1, "totalAmount": None}, + ], +) +def test_naming_one_field_twice_is_rejected(values): + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _oneof_descriptor() + celified = {k: (None if v is None else celtypes.IntType(v)) for k, v in values.items()} + with pytest.raises(ValueError, match="names field 'tests.oo.Choice.total_amount' twice"): + convert(celified, cls()) + + +# The accepted half, which is where the two rules differ: a null does not count towards a oneof +# collision but does count as having set a field. +def test_a_null_does_not_collide_with_its_oneof_sibling(): + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _oneof_descriptor() + assert convert({"a": celtypes.IntType(1), "b": None}, cls()).a == 1 + assert convert({"a": None, "b": celtypes.IntType(2)}, cls()).b == 2 + assert convert({"a": None, "b": None}, cls()).WhichOneof("choice") is None + # Two nulls for one field set nothing, so neither is a duplicate. + assert convert({"total_amount": None, "totalAmount": None}, cls()).total_amount == 0 + # A null first, then a value, is the same: the null set nothing to collide with. + assert convert({"total_amount": None, "totalAmount": celtypes.IntType(1)}, cls()).total_amount == 1 + # And one member of the oneof plus an unrelated field is fine. + out = convert({"a": celtypes.IntType(1), "total_amount": celtypes.IntType(3)}, cls()) + assert (out.a, out.total_amount) == (1, 3) + + +# A protobuf map value cannot be null either, and dropping the entry reported success while +# deleting it. The JVM's write-back parse says "Map value cannot be null." - measured against +# protobuf-java 4.35.1 on {"mp": {"a":1,"b":null}}. +def test_a_null_map_value_is_rejected(): + from celpy import celtypes + + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _map_descriptor() + original = cls() + assert dict(convert({"counts": {"a": celtypes.IntType(1)}}, original).counts) == {"a": 1} + with pytest.raises(ValueError, match="null value to map field 'counts'"): + convert({"counts": {"a": celtypes.IntType(1), "b": None}}, original) + + +def _map_descriptor(): + """A message with a map field, built at runtime.""" + from google.protobuf import descriptor_pb2, descriptor_pool, message_factory + + fdp = descriptor_pb2.FileDescriptorProto() + fdp.name, fdp.package, fdp.syntax = "map_probe.proto", "tests.mp", "proto3" + msg = fdp.message_type.add() + msg.name = "Counts" + entry = msg.nested_type.add() + entry.name = "CountsEntry" + entry.options.map_entry = True + field_type = descriptor_pb2.FieldDescriptorProto + entry.field.add(name="key", number=1, type=field_type.TYPE_STRING, label=field_type.LABEL_OPTIONAL, json_name="key") + entry.field.add( + name="value", number=2, type=field_type.TYPE_INT32, label=field_type.LABEL_OPTIONAL, json_name="value" + ) + msg.field.add( + name="counts", + number=1, + type=field_type.TYPE_MESSAGE, + type_name=".tests.mp.Counts.CountsEntry", + label=field_type.LABEL_REPEATED, + json_name="counts", + ) + + pool = descriptor_pool.DescriptorPool() + pool.Add(fdp) + desc = pool.FindMessageTypeByName("tests.mp.Counts") + return desc, message_factory.GetMessageClass(desc) + + +# A protobuf repeated field cannot hold null. Dropping the element changed the list's length +# and hid the mistake; the JVM says "Repeated field elements cannot be null in field: ...". +def test_null_element_in_a_repeated_field_is_rejected(): + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import convert + + _, cls = _scalar_descriptor() + original = cls() + assert list(convert({"codes": ["a", "b"]}, original).codes) == ["a", "b"] + with pytest.raises(ValueError, match="cannot write null to repeated field 'codes'"): + convert({"codes": ["a", None, "b"]}, original) diff --git a/tests/schema_registry/test_cel_null_avro_field.py b/tests/schema_registry/test_cel_null_avro_field.py new file mode 100644 index 000000000..9badd1c9f --- /dev/null +++ b/tests/schema_registry/test_cel_null_avro_field.py @@ -0,0 +1,177 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +A ``CEL_FIELD`` rule over the *null* branch of an Avro ``["null", T]`` union must be evaluated, +not skipped. + +Avro's null is a first-class value, and the reference binds it as CEL null so a rule can guard +with ``value == null``. Skipping the field instead removes that capability and is *silent*: a +rule that never ran and a rule that ran and passed produce the same result, so nothing in a +positive-only test can tell them apart. + +Two guards used to prevent it: one in ``CelFieldExecutor`` (``if field_value is None``), and the +blanket ``if message is None`` at the top of the Avro walk. The reference has neither - it guards +a null only in the *record* case, where there are no fields to walk - and leaves the decision to +each format's walk. The protobuf walk still skips an unset field, which is correct there: a field +with presence that is unset has no value, and writing one back would materialise it. +""" + +import json +from decimal import Decimal + +import pytest + +from confluent_kafka.schema_registry.common.avro import transform +from confluent_kafka.schema_registry.rules.cel.cel_field_executor import CelFieldExecutor +from confluent_kafka.schema_registry.schema_registry_client import Rule, RuleKind, RuleMode, Schema +from confluent_kafka.schema_registry.serde import RuleConditionError, RuleContext, RuleError + +_SCHEMA = { + "type": "record", + "name": "Nullable", + "fields": [ + { + "name": "amount", + "type": ["null", {"type": "bytes", "logicalType": "decimal", "precision": 8, "scale": 2}], + "confluent:tags": ["AMOUNT"], + }, + {"name": "note", "type": ["null", "string"], "confluent:tags": ["NOTE"]}, + {"name": "plain", "type": "string"}, + ], +} + + +# The tag map the walk would otherwise read off the parsed schema. Without it the rule matches +# no field and the walk returns the record untouched - which looks exactly like a pass. +_INLINE_TAGS = {"Nullable.amount": {"AMOUNT"}, "Nullable.note": {"NOTE"}} + + +def _run(expr, amount): + rule = Rule("r", None, RuleKind.CONDITION, RuleMode.WRITE, "CEL_FIELD", ["AMOUNT"], None, expr, None, None, False) + ctx = RuleContext( + None, + None, + None, + Schema(json.dumps(_SCHEMA), "AVRO"), + "t-value", + RuleMode.WRITE, + rule, + 0, + [rule], + _INLINE_TAGS, + None, + ) + ft = CelFieldExecutor().new_transform(ctx) + return transform(ctx, _SCHEMA, {"amount": amount, "plain": "hi"}, ft) + + +def test_null_field_reaches_the_rule(): + """`value == null` can only be true if the null was bound and the rule ran.""" + assert _run("value == null", None) is not None + + +def test_null_field_was_not_merely_skipped(): + """The discriminator: `value != null` is false on a null, so it must FAIL. + + Without it, the test above is satisfied by a rule that never ran - a skipped field + reports no violation either. + """ + with pytest.raises((RuleConditionError, RuleError)): + _run("value != null", None) + + +def test_unguarded_rule_on_a_null_raises(): + """An expression that cannot handle a null fails loudly, and says why. + + At this level the evaluator's error propagates as-is; the serializer wraps it in a + RuleError naming the rule. Either way it is loud, which is the point - the reference + raises here rather than passing silently. + """ + with pytest.raises(Exception, match="cannot convert null"): + _run('decimals.gt(decimal(value), decimal("10.00"))', None) + + +def test_a_present_value_still_evaluates_normally(): + """The must-pass twin: removing the skip must not break the ordinary case.""" + assert _run('decimals.gt(decimal(value), decimal("10.00"))', Decimal("12.34")) is not None + with pytest.raises((RuleConditionError, RuleError)): + _run('decimals.gt(decimal(value), decimal("100.00"))', Decimal("12.34")) + + +def _run_transform(expr, tag, record): + rule = Rule("r", None, RuleKind.TRANSFORM, RuleMode.WRITE, "CEL_FIELD", [tag], None, expr, None, None, False) + ctx = RuleContext( + None, + None, + None, + Schema(json.dumps(_SCHEMA), "AVRO"), + "t-value", + RuleMode.WRITE, + rule, + 0, + [rule], + _INLINE_TAGS, + None, + ) + ft = CelFieldExecutor().new_transform(ctx) + return transform(ctx, _SCHEMA, record, ft) + + +def test_value_written_to_a_null_branch_leaves_the_null_notation(): + """A rule that fills a null branch must land on the value branch. + + Re-emitting ``("null", x)`` is silent data loss - fastavro encodes it as null and drops x. + The reference resolves the branch from the datum, so the notation has to go. + """ + out = _run_transform('"recovered"', "NOTE", {"note": ("null", None), "plain": "hi"}) + assert out["note"] == "recovered" + + +def test_a_present_branch_keeps_its_notation(): + """The must-pass twin: tuple notation still selects the branch it names.""" + out = _run_transform('value + "!"', "NOTE", {"note": ("string", "a"), "plain": "hi"}) + assert out["note"] == ("string", "a!") + + +def test_a_branch_that_no_longer_fits_is_re_resolved(): + """A rule that changes the value's type moves it off a branch that cannot hold it. + + The reference keeps no branch at all, so a string turned into an int lands on the int + branch; re-emitting ("string", 3) makes fastavro refuse the record instead. + """ + schema = { + "type": "record", + "name": "Nullable", + "fields": [{"name": "note", "type": ["string", "int"], "confluent:tags": ["NOTE"]}], + } + rule = Rule("r", None, RuleKind.TRANSFORM, RuleMode.WRITE, "CEL_FIELD", ["NOTE"], None, "3", None, None, False) + ctx = RuleContext( + None, + None, + None, + Schema(json.dumps(schema), "AVRO"), + "t-value", + RuleMode.WRITE, + rule, + 0, + [rule], + {"Nullable.note": {"NOTE"}}, + None, + ) + ft = CelFieldExecutor().new_transform(ctx) + out = transform(ctx, schema, {"note": ("string", "a")}, ft) + assert out["note"] == 3 diff --git a/tests/schema_registry/test_cel_validator.py b/tests/schema_registry/test_cel_validator.py index be211e6be..1f5de7404 100644 --- a/tests/schema_registry/test_cel_validator.py +++ b/tests/schema_registry/test_cel_validator.py @@ -21,6 +21,7 @@ import datetime +import celpy import pytest from celpy import celtypes from google.protobuf import descriptor_pb2, message_factory, wrappers_pb2 @@ -28,6 +29,8 @@ from google.protobuf.timestamp_pb2 import Timestamp from confluent_kafka.schema_registry.common.protobuf import validate_message as validate_protobuf +from confluent_kafka.schema_registry.confluent.type import decimal_pb2, variant_pb2 +from confluent_kafka.schema_registry.confluent.type import variant_utils as vu from confluent_kafka.schema_registry.rules.cel.cel_executor import _value_to_cel from confluent_kafka.schema_registry.rules.cel.cel_validator import CelValidator from confluent_kafka.schema_registry.serde import RuleError, ValidationRule @@ -146,6 +149,877 @@ def test_protobuf_map_field_binds_a_map(validator): assert validator.execute(rule("'a' in this"), fd, message.labels) is True +# A ``confluent.type.Decimal`` proto field is bound into CEL as a celpy MessageType wrapper +# (the same shape produced whether the Decimal is the whole message or a nested field), so +# ``decimal(...)`` must unwrap it and dispatch the ``decimals.*`` operators against it. This +# mirrors the JVM client's CelValidatorDecimalTest, which reads a ``confluent.type.Decimal`` +# field via ``decimal(this)``. +def test_decimal_unwraps_a_confluent_type_decimal_message(validator): + # 12.34 = unscaled 1234 (0x04D2) at scale 2. + d = decimal_pb2.Decimal(value=(1234).to_bytes(2, "big"), scale=2) + assert validator.execute(rule("decimals.gt(decimal(this), decimal('10.00'))"), d.DESCRIPTOR, d) is True + assert validator.execute(rule("decimals.lt(decimal(this), decimal('10.00'))"), d.DESCRIPTOR, d) is False + + +# A decimal reached by *selection* rather than bound directly must compare numerically too. +# The boundary conversion only sees what is bound, so ``this.a`` stays a confluent.type.Decimal +# message; comparing those structurally - field by field over unscaled bytes and scale - calls +# 1.50 and 1.5 unequal. Containers and ``in`` follow the same rule, or they contradict ``==``. +@pytest.mark.parametrize( + "expr,expected", + [ + ("this.a == this.b", True), + ("this.a != this.b", False), + ("[this.a] == [this.b]", True), + ("{'k': this.a} == {'k': this.b}", True), + ("this.a in [this.b]", True), + ("decimals.eq(this.a, this.b)", True), + # Negative controls. + ("this.a == decimal('9')", False), + ("[this.a] == [decimal('9')]", False), + ("this.a in [decimal('9')]", False), + ], +) +def test_nested_proto_decimal_equality(validator, expr, expected): + def dec(unscaled, scale): + return decimal_pb2.Decimal(value=unscaled.to_bytes(2, "big"), scale=scale) + + # 1.50 (unscaled 150, scale 2) and 1.5 (unscaled 15, scale 1) - one number, two encodings. + holder = {"a": dec(150, 2), "b": dec(15, 1)} + assert validator.execute(rule(expr), None, holder) is expected + + +# Overriding the equality operators must not disturb anything that has no decimal in it. +@pytest.mark.parametrize( + "expr,expected", + [ + ("1 == 1", True), + ("1 == 2", False), + ("1 != 2", True), + ("'a' == 'a'", True), + ("[1, 2] == [1, 2]", True), + ("[1, 2] == [2, 1]", False), + ("{'a': 1} == {'a': 1}", True), + ("2 in [1, 2]", True), + ("3 in [1, 2]", False), + ("b'x' == b'x'", True), + ("null == null", True), + ], +) +def test_equality_unchanged_without_decimals(validator, expr, expected): + assert validator.execute(rule(expr), None, {"unused": 1}) is expected + + +# Cross-client parity: a bare ``confluent.type.Decimal`` field is usable with ``decimals.*``, +# ``==``, ``string()`` and ``double()`` with **no ``decimal(...)`` call** on it. The +# discriminating case is the scale-differing equality: a client comparing decimals by their +# protobuf encoding (unscaled bytes plus scale, field by field) answers False for +# ``decimal("12.340")``, because 12.34 and 12.340 are the same number in two encodings. +_BARE_PROTO_DECIMAL_CASES = [ + # Bare: no constructor call on the field. + ('decimals.eq(this, decimal("12.34"))', True), + ('decimals.gt(this, decimal("10.00"))', True), + # The wrapped form must keep working (decimal(...) re-entry). + ('decimals.eq(decimal(this), decimal("12.34"))', True), + # `==` is numeric on it: 12.34 equals 12.340 despite the differing scale. + ('this == decimal("12.340")', True), + ('this != decimal("12.340")', False), + ('decimals.lt(this, decimal("100"))', True), + # Negative control: a false comparison must still be False. + ('decimals.gt(this, decimal("100"))', False), + ('string(this) == "12.34"', True), + ('double(this) == 12.34', True), +] + + +@pytest.mark.parametrize("expr,expected", _BARE_PROTO_DECIMAL_CASES) +def test_proto_decimal_needs_no_constructor(validator, expr, expected): + # 12.34 = unscaled 1234 at scale 2. + d = decimal_pb2.Decimal(value=(1234).to_bytes(2, "big"), scale=2) + assert validator.execute(rule(expr), d.DESCRIPTOR, d) is expected + + +# The Python decimal layer must match java.math.BigDecimal's EXACT/unbounded semantics +# for add/sub/mul/mod, setScale/quantize (round/trunc/floor/ceil), and scaleb — rather +# than the thread-local default context (prec=28) which silently rounds or hard-errors on +# values with >28 significant digits. Only div/sqrt cap at 38 digits. These are the +# Java-reference regression cases (#30 exact arithmetic, #31 negative-scale round/trunc, +# #32 no-cap floor/ceil, #33 exact mod, #34 exact decimal-from-bytes). +@pytest.mark.parametrize( + "expr, expected", + [ + # #30 exact add/mul — no silent rounding of the >28-digit result. + ('string(decimals.add(decimal("1E38"), decimal("1")))', "100000000000000000000000000000000000001"), + ( + 'string(decimals.mul(decimal("12345678901234567890"), ' 'decimal("98765432109876543210")))', + "1219326311370217952237463801111263526900", + ), + # Scale preservation still holds for ordinary-magnitude operands. + ('string(decimals.mul(decimal("2.0"), decimal("3.0")))', "6.00"), + ('string(decimals.add(decimal("1.5"), decimal("1.25")))', "2.75"), + # #31 negative-scale round/trunc — quantize target Decimal(1).scaleb(-scale), + # so scale=-2 rounds/truncates to the hundreds place (not to an integer). + ('string(decimals.round(decimal("1234.5"), -2))', "1200"), + ('string(decimals.trunc(decimal("1234"), -2))', "1200"), + # #32 no 28-digit cap on floor (30-digit value passes through, no error). + ('string(decimals.floor(decimal("123456789012345678901234567890")))', "123456789012345678901234567890"), + # #33 exact mod — quotient exceeds 38 digits, but remainder is exact. + ('string(decimals.mod(decimal("1E40"), decimal("3")))', "1"), + # #34 decimal(dyn) from a >28-digit string round-trips exactly. + ('string(decimal("12345678901234567890123456789012345"))', "12345678901234567890123456789012345"), + ], +) +def test_decimal_ops_match_java_bigdecimal_exact_semantics(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) == expected + + +# An exact div/sqrt result carries the reference's *preferred* scale, not the quotient's own +# natural scale: `dividend.scale - divisor.scale` for divide, `scale / 2` truncated toward +# zero for square root. Trailing zeros are kept down to it and padded up to it, never +# stripped below. +# +# Division needs no code of ours: the decimal arithmetic spec's ideal exponent for divide is +# `exponent(dividend) - exponent(divisor)`, which is the same quantity. Square root is where +# libmpdec and the reference part company - the spec floors `exponent / 2` where BigDecimal +# truncates `scale / 2` - so they agree on an even scale and differ by one on an odd one. +@pytest.mark.parametrize( + "expr, expected", + [ + # divide: already native, pinned so a future context change cannot drift it. + ('string(decimals.div(decimal("10.0"), decimal("2.0")))', "5"), + ('string(decimals.div(decimal("10.0"), decimal("2")))', "5.0"), + ('string(decimals.div(decimal("6.0"), decimal("3")))', "2.0"), + ('string(decimals.div(decimal("10.00"), decimal("2")))', "5.00"), + ('string(decimals.div(decimal("1.000"), decimal("0.1")))', "10.00"), + ('string(decimals.div(decimal("-6.0"), decimal("3")))', "-2.0"), + # ...but never below the exact quotient's own scale: 10/4 is 2.5 at a preferred 0. + ('string(decimals.div(decimal("10"), decimal("4")))', "2.5"), + ('string(decimals.div(decimal("1.0"), decimal("8")))', "0.125"), + # An inexact quotient keeps all 38 digits; padding it would claim digits it lacks. + ('string(decimals.div(decimal("1.00000"), decimal("3")))', "0." + "3" * 38), + # sqrt, even scale: libmpdec and the reference already agree. + ('string(decimals.sqrt(decimal("4.00")))', "2.0"), + ('string(decimals.sqrt(decimal("0.0001")))', "0.01"), + ('string(decimals.sqrt(decimal("100.0000")))', "10.00"), + # sqrt, odd scale: this is the divergence. Native libmpdec answers "3.0" and "4.00". + ('string(decimals.sqrt(decimal("9.0")))', "3"), + ('string(decimals.sqrt(decimal("400.0")))', "20"), + ('string(decimals.sqrt(decimal("16.000")))', "4.0"), + # An inexact root is left at full precision. + ('string(decimals.sqrt(decimal("2")))', "1.4142135623730950488016887242096980786"), + ], +) +def test_exact_div_and_sqrt_carry_the_preferred_scale(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) == expected + + +# The scale itself, not its rendering. `toPlainString` hides the difference for a zero and +# for a negative scale - "0", "0" and "500" read the same at several scales - but the scale +# is a field of the `confluent.type.Decimal` encoding, so it has to be asserted directly. +# A zero is the case a strip loop cannot handle: it takes the preferred scale outright, in +# both directions, because the reference returns `zeroValueOf(preferredScale)`. +# The preferred scale does not override the 38-digit context precision. The reference pads +# toward the preferred scale only while the result still fits in ``mc.precision`` significant +# digits, and stops short otherwise, so the target is +# ``min(preferred, minimal_scale + (38 - minimal_precision))`` floored at the minimal scale. +# +# Division needs no code of ours - libmpdec caps its own ideal exponent, as the decimal +# arithmetic spec requires - but it is pinned here anyway. Square root goes through +# ``_apply_preferred_scale``, which quantizes in ``_EXACT_CONTEXT`` and so had nothing capping +# it: ``sqrt(1.<100 zeros>)`` padded to 51 significant digits against the reference's 38. +@pytest.mark.parametrize( + "expr, expected", + [ + # 37 zeros is exactly 38 significant digits: the last reachable preferred scale. + ('string(decimals.div(decimal("1.0000000000000000000000000000000000000"), decimal("1")))', "1." + "0" * 37), + # 40 and 100 would need 41 and 101 digits; both stop at 37. + ('string(decimals.div(decimal("1.0000000000000000000000000000000000000000"), decimal("1")))', "1." + "0" * 37), + ( + 'string(decimals.div(decimal("1.' + "0" * 100 + '"), decimal("1")))', + "1." + "0" * 37, + ), + # The cap is on *precision*, not on scale, so a value that spends digits before the + # padding starts reaches a higher scale: 0.5 gets to 38 where 1 gets only to 37... + ( + 'string(decimals.div(decimal("1.' + "0" * 100 + '"), decimal("2")))', + "0.5" + "0" * 37, + ), + # ...and 0.125 also gets to 38, from a minimal scale of 3 rather than 1. + ( + 'string(decimals.div(decimal("1.' + "0" * 100 + '"), decimal("8")))', + "0.125" + "0" * 35, + ), + # sqrt: preferred 20 fits, 37 is exactly the ceiling, 50 does not fit. + ('string(decimals.sqrt(decimal("1.0000000000000000000000000000000000000000")))', "1." + "0" * 20), + ( + 'string(decimals.sqrt(decimal("1.' + "0" * 74 + '")))', + "1." + "0" * 37, + ), + ( + 'string(decimals.sqrt(decimal("1.' + "0" * 100 + '")))', + "1." + "0" * 37, + ), + ], +) +def test_the_preferred_scale_cannot_exceed_the_context_precision(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) == expected + + +# A zero is exempt from that cap: it is one digit at any scale, so it keeps the full preferred +# scale. Measured on the reference: ``0.<100 zeros> / 1`` is scale 100 at precision 1. +@pytest.mark.parametrize( + "expr, expected_scale", + [ + ( + 'decimals.div(decimal("0.' + "0" * 100 + '"), decimal("1"))', + 100, + ), + ( + 'decimals.div(decimal("0.' + "0" * 100 + '"), decimal("3.0"))', + 99, + ), + ], +) +def test_a_zero_is_exempt_from_the_precision_cap(expr, expected_scale): + from decimal import Decimal + + from confluent_kafka.schema_registry.rules.cel import decimal_funcs + + a = Decimal("0." + "0" * 100) + b = Decimal("1") if expected_scale == 100 else Decimal("3.0") + q = decimal_funcs._decimals_div(a, b) + assert -q.as_tuple().exponent == expected_scale + + +@pytest.mark.parametrize( + "literal, expected_scale", + [ + ("0", 0), + ("0.0", 0), + ("0.00", 1), + ("0.000", 1), + ("9.0", 0), + ("16.000", 1), + ("4E+2", -1), + ("1E+4", -2), + ("250E+3", -1), + ], +) +def test_sqrt_preferred_scale_is_the_scale_not_the_rendering(literal, expected_scale): + from decimal import Decimal + + from confluent_kafka.schema_registry.rules.cel import decimal_funcs + + root = decimal_funcs._decimals_sqrt(Decimal(literal)) + assert -root.as_tuple().exponent == expected_scale + + +# #34 decimal(bytes, scale): a 38-digit unscaled value at scale 5 must round-trip +# exactly through _from_bytes_scale (no rounding to the 28-digit default context). +def test_decimal_from_bytes_scale_is_exact(validator): + unscaled = 12345678901234567890123456789012345678 # 38 digits + raw = unscaled.to_bytes(16, "big", signed=True) + result = validator.execute(rule("string(decimal(this, 5))"), None, raw) + assert result == "123456789012345678901234567890123.45678" + + +# ``decimal()`` / ``decimal()`` must match java.math.BigDecimal's +# ``new BigDecimal(String)`` / ``BigDecimal.valueOf(double)``, which throw +# NumberFormatException on non-finite values, underscore digit-grouping, and +# surrounding whitespace. Python's ``Decimal(str)`` silently accepts all of +# these — building a poisoned NaN/Infinity Decimal or a wrongly-parsed 1000 — +# so the constructor must reject them (surfaced as a RuleError). The +# ``decimal(bytes, scale)`` path parses no string and is unaffected. +@pytest.mark.parametrize( + "expr", + [ + 'decimal("NaN") > decimal("0")', + 'decimal("Infinity") > decimal("0")', + 'decimal("-Infinity") > decimal("0")', + 'decimal("-inf") > decimal("0")', + 'decimal("sNaN") > decimal("0")', + 'decimal("1_000") > decimal("0")', + # Surrounding whitespace: Java rejects; Python's Decimal strips it. + "decimal(' 5 ') > decimal('0')", + # An exponent that will not fit BigDecimal's signed-int scale. Measured against the + # JVM, the accepted band is symmetric -- |exponent| <= INT32_MAX -- with both + # +/-2147483648 a NumberFormatException there. Python's Decimal accepts them, and + # rendering one as fixed-point would try to materialise billions of digits. + 'decimal("1e-2147483648") > decimal("0")', + 'decimal("1e2147483648") > decimal("0")', + 'decimal("1E+2147483648") > decimal("0")', + ], +) +def test_decimal_rejects_inputs_java_bigdecimal_rejects(validator, expr): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule(expr), None, 1) + + +# The other side of the band: the widest exponents the JVM *accepts* must keep working, so +# the guard above cannot be off by one. +@pytest.mark.parametrize("expr", ['decimal("1e-2147483647")', 'decimal("1e2147483647")']) +def test_decimal_accepts_the_widest_exponents_java_accepts(validator, expr): + assert validator.execute(rule(f"{expr} != decimal(\"0\")"), None, 1) is True + + +# A NaN/Infinity double routed through ``decimal()`` must also be +# rejected — Java's ``BigDecimal.valueOf(double)`` throws on non-finite doubles. +@pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf")]) +def test_decimal_rejects_non_finite_double(validator, value): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule("decimal(this) > decimal('0')"), None, value) + + +# Legitimate finite decimals must still parse: ordinary decimals, scientific +# notation, negatives, negative zero, a leading '+', and a finite double. +@pytest.mark.parametrize( + "expr, expected", + [ + ('string(decimal("123.45"))', "123.45"), + ('string(decimal("1e40"))', "10000000000000000000000000000000000000000"), + ('string(decimal("-0.5"))', "-0.5"), + # BigDecimal has no negative zero: BigDecimal("-0").toPlainString() is "0" (signum 0), + # and a scale survives the sign being dropped ("-0.00" -> "0.00"). + ('string(decimal("-0"))', "0"), + ('string(decimal("-0.00"))', "0.00"), + ('string(decimals.round(decimal("-0.4"), 0))', "0"), + ('string(decimal("+5"))', "5"), + ], +) +def test_decimal_accepts_legitimate_finite_values(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) == expected + + +def test_decimal_accepts_finite_double(validator): + assert validator.execute(rule("string(decimal(this))"), None, 1.5) == "1.5" + + +# CEL ``==``/``!=`` on two Decimal values must be NUMERIC (scale-insensitive), matching +# ``decimals.eq`` (java.math.BigDecimal.compareTo) rather than an equals() that also +# compares scale. Python's ``decimal.Decimal.__eq__`` is already numeric +# (``Decimal("2.0") == Decimal("2.00")`` is True), and celpy dispatches ``==`` to it, +# so this is a regression guard — no code change is required. +@pytest.mark.parametrize( + "expr, expected", + [ + ('decimal("2.0") == decimal("2.00")', True), + ('decimal("2.0") == decimal("2.0")', True), + ('decimal("2.0") == decimal("2.1")', False), + ('decimal("2.0") != decimal("2.00")', False), + ('decimal("2.0") != decimal("2.1")', True), + ], +) +def test_decimal_equality_is_numeric_scale_insensitive(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) is expected + + +# -------------------------------------------------------------------------------------- +# Variant CEL functions +# -------------------------------------------------------------------------------------- + +_VARIANT_JSON = '{"name":"alice","age":30,"scores":[10,20,30],"nested":{"x":1},"explicit":null}' + + +# `this` is bound to a JSON string; variants.parseJson(this) turns it into a Variant, then +# the variants.* accessors navigate and extract. Covers the null model (absent vs +# variant-null), path/field/index navigation, typed extraction, and toJson. +@pytest.mark.parametrize( + "expr", + [ + "variants.type(variants.parseJson(this)) == 'object'", + "variants.as(variants.field(variants.parseJson(this), 'name'), 'string') == 'alice'", + "variants.as(variants.field(variants.parseJson(this), 'age'), 'int') == 30", + # Absent (missing field) vs present-but-variant-null (explicit JSON null). + "variants.field(variants.parseJson(this), 'missing') == null", + "variants.isNull(variants.field(variants.parseJson(this), 'explicit'))", + "!variants.isNull(variants.field(variants.parseJson(this), 'missing'))", + "variants.as(variants.path(variants.parseJson(this), '$.nested.x'), 'int') == 1", + "variants.as(variants.index(" "variants.field(variants.parseJson(this), 'scores'), 2), 'int') == 30", + # tryAs returns CEL null on a type mismatch (age is not a string). + "variants.tryAs(variants.field(variants.parseJson(this), 'age'), 'string') == null", + "variants.toJson(variants.field(variants.parseJson(this), 'nested')) == '{\"x\":1}'", + ], +) +def test_variant_functions_over_parsed_json(validator, expr): + assert validator.execute(rule(expr), None, _VARIANT_JSON) is True + + +# Variant `==` is equality of the encoding. Before Variant grew __eq__ this was object identity, +# while a confluent.type.Variant protobuf field has always compared bytes. Sound but incomplete: +# equal bytes mean equal values, but one value has many encodings. +@pytest.mark.parametrize( + "expr, expected", + [ + ('variants.parseJson("1") == variants.parseJson("1")', True), + ('variants.parseJson("1") != variants.parseJson("1")', False), + ('variants.parseJson("{}") == variants.parseJson("{}")', True), + ("variants.parseJson('{\"a\":1}') == variants.parseJson('{\"a\":1}')", True), + ('variants.parseJson("1") == variants.parseJson("2")', False), + # Incomplete, as documented: an int and a double are two encodings. + ('variants.parseJson("1") == variants.parseJson("1.0")', False), + # Containers recurse with the same equality. + ('[variants.parseJson("1")] == [variants.parseJson("1")]', True), + ('[variants.parseJson("1")] == [variants.parseJson("2")]', False), + # Navigation: the same position in an identical parent. + ( + "variants.field(variants.parseJson('{\"a\":1}'), 'a') == " + "variants.field(variants.parseJson('{\"a\":1}'), 'a')", + True, + ), + # A field holding 1 is not the standalone variant 1: it carries its parent's metadata + # dictionary, which is part of the comparison. + ( + "variants.field(variants.parseJson('{\"a\":1}'), 'a') == variants.parseJson(\"1\")", + False, + ), + # And the whole document, reached two ways, is the same variant. + ("variants.parseJson(this) == variants.parseJson(this)", True), + ], +) +def test_variant_equality_is_over_the_encoding(validator, expr, expected): + assert validator.execute(rule(expr), None, _VARIANT_JSON) is expected + + +# A CEL `uint` is a distinct type, and every declared overload here takes an `int` +# (`SimpleType.INT` in the reference), so a uint argument has no matching overload. The typed +# runtimes say so before evaluating - cel-java, cel-go and cel-es all report "found no matching +# overload" at compile time, measured - and the Rust client refuses at runtime. celpy is untyped +# and `celtypes.UintType` subclasses `int`, so every `isinstance(x, int)` guard accepted a uint +# until each one excluded it explicitly. +@pytest.mark.parametrize( + "expr", + [ + 'decimals.round(decimal("1"), 1u) == decimal("1.0")', + 'decimals.trunc(decimal("1.29"), 1u) == decimal("1.2")', + 'decimals.round(decimal("1"), -1u) == decimal("0")', + 'timestamp(0u) == timestamp("1970-01-01T00:00:00Z")', + # Both arguments of the precision form, the second of which is easy to miss. + 'timestamp(0u, 0) == timestamp("1970-01-01T00:00:00Z")', + 'timestamp(0, 0u) == timestamp("1970-01-01T00:00:00Z")', + "variants.index(variants.parseJson('[10,20,30]'), 1u) != null", + ], +) +def test_a_uint_has_no_overload_where_an_int_is_declared(validator, expr): + with pytest.raises(Exception): + validator.execute(rule(expr), None, _VARIANT_JSON) + + +# The int forms these mirror, so the guards cannot be satisfied by refusing everything. +@pytest.mark.parametrize( + "expr", + [ + 'decimals.round(decimal("1"), 1) == decimal("1.0")', + 'decimals.trunc(decimal("1.29"), 1) == decimal("1.2")', + 'timestamp(0) == timestamp("1970-01-01T00:00:00Z")', + 'timestamp(0, 0) == timestamp("1970-01-01T00:00:00Z")', + "variants.index(variants.parseJson('[10,20,30]'), 1) != null", + ], +) +def test_the_int_forms_still_answer(validator, expr): + assert validator.execute(rule(expr), None, _VARIANT_JSON) is True + + +_MAX_TS_MICROS = 253402300799 * 1_000_000 + 999_999 +_MIN_TS_MICROS = -62135596800 * 1_000_000 + + +def _timestamp_variant(micros): + from confluent_kafka.schema_registry.confluent.type.variant_utils import VariantBuilder + + b = VariantBuilder() + b.append_timestamp_tz(micros) + return b.build() + + +# A variant timestamp spans the whole int64 range while a CEL timestamp is 0001-9999, so an +# out-of-range value is reachable from data. It used to be built anyway, leaving an instant that +# could not be rendered but could still be compared - `< now` answered a confident false for a +# value that is not a time. Refused now, and routed through the as/tryAs split so a rule can +# guard, matching the reference's variantGetTimestamp. +@pytest.mark.parametrize("micros", [0, _MAX_TS_MICROS, _MIN_TS_MICROS]) +def test_variant_as_timestamp_accepts_the_range(validator, micros): + v = _timestamp_variant(micros) + assert validator.execute(rule('variants.as(this, "timestamp") == variants.as(this, "timestamp")'), None, v) is True + # tryAs answers a timestamp, not null - otherwise the guard test below proves nothing. + assert validator.execute(rule('variants.tryAs(this, "timestamp") == null'), None, v) is False + + +@pytest.mark.parametrize( + "micros", + [ + 9223372036854775807, + -9223372036854775807, + _MAX_TS_MICROS + 1_000_000, + ], +) +def test_variant_as_timestamp_refuses_out_of_range(validator, micros): + v = _timestamp_variant(micros) + with pytest.raises(Exception): + validator.execute(rule('variants.as(this, "timestamp") != null'), None, v) + # tryAs answers CEL null instead, so a rule can guard on it. + assert validator.execute(rule('variants.tryAs(this, "timestamp") == null'), None, v) is True + + +# An Avro `variant` logical-type field decodes to a Variant (via the logical type registered +# in common/avro.py), which then flows into CEL through variant(this). +def test_avro_variant_field_into_cel(validator): + import io + + import fastavro + + import confluent_kafka.schema_registry.common.avro # noqa: F401 (registers the logical type) + + schema = fastavro.parse_schema( + { + "type": "record", + "name": "confluent.type.Variant", + "logicalType": "variant", + "fields": [{"name": "metadata", "type": "bytes"}, {"name": "value", "type": "bytes"}], + } + ) + built = vu.parse_json('{"name":"alice","age":30}') + value, metadata = built.value, built.metadata + buf = io.BytesIO() + fastavro.schemaless_writer(buf, schema, vu.Variant(value, metadata)) + buf.seek(0) + decoded = fastavro.schemaless_reader(buf, schema) + assert isinstance(decoded, vu.Variant) + assert ( + validator.execute( + rule("variants.as(variants.field(variant(this), 'name'), 'string') == 'alice'"), None, decoded + ) + is True + ) + + +# A confluent.type.Variant proto field is bound into CEL as a celpy MessageType wrapper; +# variant(...) must unwrap it, mirroring the decimal test above and the JVM client. +def test_proto_variant_field_into_cel(validator): + built = vu.parse_json('{"name":"alice","age":30}') + value, metadata = built.value, built.metadata + v = variant_pb2.Variant(value=value, metadata=metadata) + expr = "variants.as(variants.field(variant(this), 'name'), 'string') == 'alice'" + assert validator.execute(rule(expr), v.DESCRIPTOR, v) is True + + +# Cross-client parity: a variant value is usable with the variants.* accessors with **no +# variant(...) call**, in both formats, and the wrapped form keeps working alongside it. The +# accessors are plain Python functions that coerce their subject, so they take whatever the +# decoder produced -- a vu.Variant from the Avro logical type, or a proto message. +_BARE_VARIANT_CASES = [ + # Bare: no constructor call. + ("variants.type(this) == 'object'", True), + ("variants.as(variants.field(this, 'name'), 'string') == 'alice'", True), + ("variants.as(variants.path(this, '$.age'), 'int') == 30", True), + # The wrapped form must keep working (variant(...) re-entry). + ("variants.as(variants.field(variant(this), 'name'), 'string') == 'alice'", True), + # A missing key is CEL null, not an error. + ("variants.field(this, 'nope') == null", True), + # Negative control. + ("variants.as(variants.field(this, 'name'), 'string') == 'bob'", False), +] + + +# ``variants.isNull`` must coerce its receiver like every other accessor. It is declared over +# dyn, so a bare variant field reaches it; a receiver check that only accepts the client's own +# Variant type answers False for the shapes a variant-typed field actually decodes to, reporting +# "not null" for a variant that holds an explicit JSON null. The bare-object cases above cannot +# catch this: isNull on an object is False either way, so only a variant that *is* null +# discriminates. +@pytest.mark.parametrize( + "expr,expected", + [ + ("variants.isNull(this)", True), + # The wrapped form has always worked and must keep working. + ("variants.isNull(variant(this))", True), + ], +) +def test_proto_variant_is_null_coerces_bare_receiver(validator, expr, expected): + built = vu.parse_json("null") + v = variant_pb2.Variant(value=built.value, metadata=built.metadata) + assert validator.execute(rule(expr), v.DESCRIPTOR, v) is expected + + +def test_proto_variant_is_null_false_for_non_null(validator): + built = vu.parse_json("5") + v = variant_pb2.Variant(value=built.value, metadata=built.metadata) + assert validator.execute(rule("variants.isNull(this)"), v.DESCRIPTOR, v) is False + + +@pytest.mark.parametrize("expr,expected", _BARE_VARIANT_CASES) +def test_avro_variant_needs_no_constructor(validator, expr, expected): + import io + + import fastavro + + import confluent_kafka.schema_registry.common.avro # noqa: F401 (registers the logical type) + + schema = fastavro.parse_schema( + { + "type": "record", + "name": "confluent.type.Variant", + "logicalType": "variant", + "fields": [{"name": "metadata", "type": "bytes"}, {"name": "value", "type": "bytes"}], + } + ) + built = vu.parse_json('{"name":"alice","age":30}') + buf = io.BytesIO() + fastavro.schemaless_writer(buf, schema, vu.Variant(built.value, built.metadata)) + buf.seek(0) + decoded = fastavro.schemaless_reader(buf, schema) + assert validator.execute(rule(expr), None, decoded) is expected + + +@pytest.mark.parametrize("expr,expected", _BARE_VARIANT_CASES) +def test_proto_variant_needs_no_constructor(validator, expr, expected): + built = vu.parse_json('{"name":"alice","age":30}') + v = variant_pb2.Variant(value=built.value, metadata=built.metadata) + assert validator.execute(rule(expr), v.DESCRIPTOR, v) is expected + + +# A string is rejected by variant(...) with a redirect to parseJson. +# An *absent* variant -- a protobuf field left unset, or an Avro variant record whose byte +# fields are empty -- carries no metadata, so there is nothing to read. It reads as CEL null and +# every accessor propagates that, rather than the Variant constructor raising on a metadata +# version byte that isn't there. +_ABSENT_VARIANT_CASES = [ + "variants.type(this) == null", + # isNull is False, not an error: an absent variant is not a JSON null. + "!variants.isNull(this)", + "variants.field(this, 'name') == null", + "variants.path(this, '$.name') == null", + "variants.toJson(this) == null", + # The explicit constructor reports it as CEL null too, like variant(null). + "variant(this) == null", +] + + +@pytest.mark.parametrize("expr", _ABSENT_VARIANT_CASES) +def test_absent_proto_variant_reads_as_null(validator, expr): + v = variant_pb2.Variant(value=b"", metadata=b"") + assert validator.execute(rule(expr), v.DESCRIPTOR, v) is True + + +@pytest.mark.parametrize("expr", _ABSENT_VARIANT_CASES) +def test_absent_avro_variant_reads_as_null(validator, expr): + # The mapping an Avro variant record decodes to, with empty byte fields. + assert validator.execute(rule(expr), None, {"metadata": b"", "value": b""}) is True + + +def test_explicit_null_variant_is_not_absent(validator): + # Absent must stay distinguishable from a variant that genuinely holds JSON null: the + # former is CEL null, the latter a present variant whose type is NULL. + assert validator.execute(rule("variants.isNull(variants.parseJson('null'))"), None, "null") is True + assert validator.execute(rule("variants.type(variants.parseJson('null')) != null"), None, "null") is True + + +def test_variant_from_empty_metadata_bytes_is_rejected(validator): + # Passing empty metadata explicitly is a rule-authoring mistake rather than an absent + # field, so it is reported instead of yielding null. + with pytest.raises(Exception) as exc: + validator.execute(rule("variants.type(variant(b'', b'')) == 'object'"), None, "x") + # The validator wraps rule failures, so the explanation is on the cause chain. + chain = [] + err = exc.value + while err is not None: + chain.append(str(err)) + err = err.__cause__ + assert any("metadata is empty" in m for m in chain), chain + + +def test_variant_rejects_string_input(validator): + with pytest.raises(RuleError, match="Could not execute"): + validator.execute(rule("variants.type(variant(this)) == 'object'"), None, "not-a-variant") + + +# variant(null) yields CEL null instead of erroring (matching the Java reference), and it +# composes: a null flows through the accessors as absent. +@pytest.mark.parametrize( + "expr", + [ + "variant(null) == null", + "variants.field(variant(null), 'k') == null", + # An absent field is null, and variant(null) of it is still null. + "variant(variants.field(variants.parseJson(this), 'missing')) == null", + ], +) +def test_variant_of_null_is_cel_null(validator, expr): + assert validator.execute(rule(expr), None, _VARIANT_JSON) is True + + +# Non-finite doubles round-trip through CEL as bareword NaN/Infinity/-Infinity (Confluent +# Java contract). Bareword literals parse (Python json.loads accepts them by default). +@pytest.mark.parametrize("tok", ["NaN", "Infinity", "-Infinity"]) +def test_variant_non_finite_bareword_roundtrip_through_cel(validator, tok): + expr = "variants.toJson(variants.parseJson(this)) == '%s'" % tok + assert validator.execute(rule(expr), None, tok) is True + + +# variants.tryParseJson of empty/whitespace-only input is a soft failure -> CEL null, +# while the strict variants.parseJson raises (surfaced as a RuleError). +@pytest.mark.parametrize("src", ["", " ", "\t\n"]) +def test_variant_try_parse_json_empty_is_cel_null(validator, src): + assert validator.execute(rule("variants.tryParseJson(this) == null"), None, src) is True + + +@pytest.mark.parametrize("src", ["", " "]) +def test_variant_parse_json_empty_raises(validator, src): + with pytest.raises(RuleError, match="Could not execute"): + validator.execute(rule("variants.type(variants.parseJson(this)) == 'object'"), None, src) + + +# -------------------------------------------------------------------------------------- +# timestamp(value, precision) +# -------------------------------------------------------------------------------------- + + +# ``timestamp(value, precision)`` must split the epoch value into whole microseconds with +# exact integer FLOOR division (mirroring Java TimestampUtils' Math.floorDiv/floorMod), +# not float division that rounds half-to-even and drops precision. datetime resolution is +# one microsecond, so sub-microsecond nanos are floored away (an inherent, Java-matching +# limit), but the microsecond itself must never round up, and negative epochs must floor +# toward negative infinity. +@pytest.mark.parametrize( + "expr", + [ + # nanos floor to the microsecond (1500 ns -> 1 us, not rounded up to 2). + 'timestamp(1500, 9) == timestamp("1970-01-01T00:00:00.000001Z")', + # 999999500 ns floors to .999999, not rounded up to the next whole second. + 'timestamp(999999500, 9) == timestamp("1970-01-01T00:00:00.999999Z")', + # Negative epoch floors toward -inf: -500 ns -> the microsecond before the epoch. + 'timestamp(-500, 9) == timestamp("1969-12-31T23:59:59.999999Z")', + # A large micros value keeps its microsecond (float division would have lost it). + 'timestamp(253402300799000001, 6) == ' 'timestamp("9999-12-31T23:59:59.000001Z")', + # millis/micros/seconds precisions are exact. + 'timestamp(1500, 3) == timestamp("1970-01-01T00:00:01.500000Z")', + 'timestamp(1, 6) == timestamp("1970-01-01T00:00:00.000001Z")', + 'timestamp(1, 0) == timestamp("1970-01-01T00:00:01Z")', + ], +) +def test_timestamp_precision_floors_with_exact_integer_arithmetic(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +def test_timestamp_bool_reports_bool_not_int(validator): + # celtypes.BoolType subclasses int (MRO: BoolType -> int -> object) and *not* + # bool, so a plain ``isinstance(v, bool)`` guard never fires for a CEL bool and + # the value used to be misreported as a unitless raw int. + with pytest.raises(RuleError) as excinfo: + validator.execute(rule("timestamp(true) == timestamp(1)"), None, 1) + assert "cannot convert bool" in str(excinfo.value.__cause__) + + +@pytest.mark.parametrize("precision", [1, 2, 4, 5, 7, 8, 10, -3]) +def test_timestamp_rejects_precision_outside_the_set(validator, precision): + # With the unit a number rather than a name, rejecting anything outside + # {0, 3, 6, 9} is the only thing between a typo and a silently wrong instant. + with pytest.raises(RuleError) as excinfo: + validator.execute(rule(f"timestamp(1700000000, {precision}) == timestamp(0)"), None, 1) + assert "unknown precision" in str(excinfo.value.__cause__) + + +def test_timestamp_datetime_components_form_still_works(validator): + # celpy's components form takes three or more args, so it never collides with + # the two-arg precision form. + assert validator.execute(rule('timestamp(2009, 2, 13) == timestamp("2009-02-13T00:00:00Z")'), None, 1) is True + + +# -------------------------------------------------------------------------------------- +# stdlib timestamp(...) — the single-int epoch-seconds overload every other client has +# -------------------------------------------------------------------------------------- + + +# celpy binds ``timestamp`` straight to celtypes.TimestampType, which accepts a +# datetime, a string, or an int followed by *at least two more* args (datetime +# components) — but rejects a lone int. cel-java (int64_to_timestamp), Go, C++ and C# +# all read a single int as epoch SECONDS, so the client registers its own "timestamp" +# that adds that overload and delegates every other form to the base implementation. +@pytest.mark.parametrize( + "expr", + [ + # The regression: a bare int is epoch seconds. + 'timestamp(1700000000) == timestamp("2023-11-14T22:13:20Z")', + 'timestamp(0) == timestamp("1970-01-01T00:00:00Z")', + # Negative / pre-epoch ints. + 'timestamp(-1) == timestamp("1969-12-31T23:59:59Z")', + 'timestamp(-2208988800) == timestamp("1900-01-01T00:00:00Z")', + # Matches timestamp(value, 0) exactly. + 'timestamp(1700000000) == timestamp(1700000000, 0)', + # The result is a real UTC-aware timestamp, usable with the timestamp methods. + "timestamp(1700000000).getFullYear() == 2023", + # Forwarded to the base implementation: the datetime-components form needs + # arity >= 3 to reach TimestampType, so the override must not swallow it. + 'timestamp(2009, 2, 13) == timestamp("2009-02-13T00:00:00Z")', + 'timestamp(2009, 2, 13, 23, 31, 30) == timestamp("2009-02-13T23:31:30Z")', + # Forwarded: RFC 3339 strings, including the lenient form celpy accepts. + 'timestamp("2023-11-14T22:13:20Z") == timestamp(1700000000)', + 'timestamp("2020-01-01 00:00:00") == timestamp("2020-01-01T00:00:00Z")', + # Forwarded: a timestamp is passed through unchanged. + 'timestamp(timestamp("2023-11-14T22:13:20Z")) == timestamp(1700000000)', + ], +) +def test_timestamp_int_is_epoch_seconds_and_other_forms_still_work(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +def test_timestamp_bool_raises_rather_than_meaning_epoch_second_one(validator): + # BoolType subclasses int, so an unguarded int check would read true as 1. + with pytest.raises(RuleError) as excinfo: + validator.execute(rule('timestamp(true) == timestamp("1970-01-01T00:00:01Z")'), None, 1) + assert "cannot convert bool" in str(excinfo.value.__cause__) + + +def test_timestamp_out_of_range_int_is_a_cel_error(validator): + with pytest.raises(RuleError) as excinfo: + validator.execute(rule("timestamp(9223372036854775807) == timestamp(0)"), None, 1) + cause = excinfo.value.__cause__ + assert isinstance(cause, celpy.CELEvalError) + assert "out of range" in str(cause) + + +# The (value, precision) overload floors the epoch to whole microseconds before the range +# check, where the reference checks Math.floorDiv(value, unitsPerSecond). Nested floor +# division by positive divisors composes, so the accept/reject boundary is the same one - +# verified here at both ends. The finer precisions cannot reach it at all: an int64 count of +# nanoseconds spans only 1677..2262. +@pytest.mark.parametrize( + "expr, expected", + [ + ("timestamp(253402300799, 0)", "9999-12-31T23:59:59Z"), + ("timestamp(-62135596800, 0)", "0001-01-01T00:00:00Z"), + ("timestamp(253402300799999, 3)", "9999-12-31T23:59:59.999Z"), + ("timestamp(9223372036854775807, 9)", "2262-04-11T23:47:16.854775Z"), + ("timestamp(-9223372036854775808, 9)", "1677-09-21T00:12:43.145224Z"), + ], +) +def test_timestamp_precision_range_boundary(validator, expr, expected): + assert validator.execute(rule(f"string({expr})"), None, 1) == expected + + +@pytest.mark.parametrize("expr", ["timestamp(253402300800, 0)", "timestamp(-62135596801, 0)"]) +def test_timestamp_precision_past_the_boundary_is_refused(validator, expr): + with pytest.raises(RuleError) as excinfo: + validator.execute(rule(f"string({expr})"), None, 1) + assert "out of range" in str(excinfo.value.__cause__) + + +# A protobuf Timestamp whose nanos fall outside the contract's [0, 999999999] is normalized +# into the neighbouring instant rather than refused, which is what cel-java does with the same +# message (measured: nanos=-1 gives the preceding nanosecond, nanos=1000000000 the next +# second). Validating the contract here would refuse values the reference accepts. +@pytest.mark.parametrize( + "nanos, expected", + [ + (-1, "1969-12-31T23:59:59.999999Z"), + (1_000_000_000, "1970-01-01T00:00:01Z"), + (1_500_000_000, "1970-01-01T00:00:01.500Z"), + ], +) +def test_proto_timestamp_nanos_outside_the_contract_is_normalized(validator, nanos, expected): + ts = Timestamp(seconds=0, nanos=nanos) + assert validator.execute(rule("string(timestamp(this))"), None, ts) == expected + + # -------------------------------------------------------------------------------------- # `now` end to end through the protobuf walker, mirroring the JVM client's test # -------------------------------------------------------------------------------------- @@ -357,3 +1231,498 @@ def test_written_wrapper_is_the_value_it_holds(validator): assert validator.execute(rule("has(this.name)"), descriptor, written) is True empty = message_factory.GetMessageClass(descriptor)() assert validator.execute(rule("has(this.name)"), descriptor, empty) is False + + +# -------------------------------------------------------------------------------------- +# string(timestamp) renders the sub-second component +# +# celpy's TimestampType.__str__ formats with strftime("%Y-%m-%dT%H:%M:%S%z") -- no %f -- so it +# dropped the fraction entirely: string(timestamp("...T22:13:20.123Z")) came back as +# "2023-11-14T22:13:20Z". The stored value was always correct (comparisons and getMilliseconds() +# agreed with the other clients), so this was a silent rendering divergence rather than an error. +# decimal_funcs._string now routes a Timestamp through timestamp_funcs.format_timestamp. +# +# Every expectation below is the verbatim output of the Java reference for the same expression. +@pytest.mark.parametrize( + ("expr", "expected"), + [ + ("string(timestamp(1700000000, 0))", "2023-11-14T22:13:20Z"), + ("string(timestamp(1700000000123, 3))", "2023-11-14T22:13:20.123Z"), + ("string(timestamp(1700000000123456, 6))", "2023-11-14T22:13:20.123456Z"), + # A whole millisecond keeps its trailing zeros (3-digit group), not ".1Z". + ("string(timestamp(1700000000100, 3))", "2023-11-14T22:13:20.100Z"), + ("string(timestamp('2023-11-14T22:13:20.5Z'))", "2023-11-14T22:13:20.500Z"), + # A zero fraction emits no decimal point at all. + ("string(timestamp(1700000000000, 3))", "2023-11-14T22:13:20Z"), + ("string(timestamp(0))", "1970-01-01T00:00:00Z"), + # Pre-epoch, where the fraction is a non-negative nano-of-second. + ("string(timestamp(-1500, 3))", "1969-12-31T23:59:58.500Z"), + # Rendered in UTC with a Z suffix whatever offset the literal carried. + ("string(timestamp('2020-01-01T00:00:00+05:00'))", "2019-12-31T19:00:00Z"), + ], +) +def test_string_timestamp_renders_subsecond(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) == expected + + +def test_string_timestamp_nanos_limited_to_microseconds(validator): + # datetime's resolution is one microsecond, so a nanosecond-precision value renders 6 digits + # where Java renders 9 (".123456789Z"). That is the same pre-existing limit that floors the + # value itself in timestamp_funcs._from_epoch, not something the formatting introduces. + assert ( + validator.execute(rule("string(timestamp(1700000000123456789, 9))"), None, 1) == "2023-11-14T22:13:20.123456Z" + ) + + +# Rescaling is bounded by this client's own width ceiling (`_SANE_WIDTH`, 10**7 digits), not +# by BigDecimal's. That is a deliberate divergence: BigInteger tops out at Integer.MAX_VALUE +# bits = 646456993 digits, and reproducing that bound across six decimal libraries is neither +# achievable nor the point. What matters is that a wide rescale becomes a *rule error* rather +# than resource exhaustion, which is the one failure a framework cannot attribute after the +# fact. Java is the only client in the family that fails cleanly on width; this stands in. +# +# The bound has to be enforced here rather than delegated, because libmpdec honours any int32 +# scale and materialises the whole coefficient: a scale of 2**31-1 costs 918 MB inside +# quantize and 9.2 GB once anything calls as_tuple() on the result. +# +# The quantizer itself is built in _EXACT_CONTEXT, not the ambient one, for an unrelated +# reason kept here because it is the same call: the ambient Emin of -999999 made +# `Decimal(1).scaleb(1000000)` raise decimal.Overflow, so a negative scale past a million was +# refused for values the JVM rounds happily - and Overflow is not an InvalidOperation, so it +# escaped as a raw Python exception rather than a rule error. +# +# So the accepted/rejected split below is *this client's* limit, and the JDK column is +# recorded only where the two now differ. +@pytest.mark.parametrize( + "expr", + [ + # Copilot's case, and the one that raised an uncaught decimal.Overflow. + 'string(decimals.round(decimal("1e1000000"), -1000000)) != ""', + 'string(decimals.round(decimal("1.23"), -1000000)) != ""', + 'string(decimals.trunc(decimal("1.23"), -1000000)) != ""', + # A wide scale in the other direction, at 10**6 - an order under the ceiling. + 'decimals.round(decimal("1e1000000"), 1000000) != decimal("0")', + # *Coarsening* a scale is free at any distance - the coefficient shrinks rather than + # grows - so none of these is bounded. Measured, all instant and all one digit wide: + # 1.23 at scale -1000000 / -100000000 / -2000000000, and 1e-1000000 and 1e-100000000 + # at scale 0. Java agrees: BigDecimal("1.23").setScale(-100000000) is precision 1. + # An `abs(shift) + digits` formula refused every one of them. + 'decimals.round(decimal("1.23"), -100000000) == decimal("0")', + 'decimals.round(decimal("1.23"), -1000000000) == decimal("0")', + 'decimals.trunc(decimal("1.23"), -1000000000) == decimal("0")', + 'decimals.round(decimal("1e-20000000")) == decimal("0")', + 'decimals.floor(decimal("1e-20000000")) == decimal("0")', + 'decimals.ceil(decimal("1e-20000000")) == decimal("1")', + 'decimals.trunc(decimal("1e-20000000")) == decimal("0")', + 'decimals.round(decimal("1e-100000000")) == decimal("0")', + # A coarsened result can leave the int32 scale domain - 1.23 at scale -2147483648 is + # 0E+2147483648 - and that is refused at the wire, not here, the same way mul's result + # is (see test_cel_message_transform). The operator itself is free. + 'decimals.round(decimal("1.23"), -2147483648) == decimal("0")', + 'decimals.round(decimal("1e1000000"), -2147483648) == decimal("0")', + # A no-op in Java too, via its `intScale >= v.scale()` early return, so no rescale + # happens and no bound applies. + 'string(decimals.trunc(decimal("1.23"), 2147483647)) == "1.23"', + # Zero rescales for free at any scale, so the width formula must exempt it - measured, + # both directions cost nothing and the result stays compact. BigDecimal agrees: + # `new BigDecimal(BigInteger.ZERO, 2147483647)` is precision 1. Without the exemption + # these are false rejections of values the reference handles. + 'decimals.round(decimal(b"", 2147483647), 0) == decimal("0")', + 'decimals.round(decimal("0"), 2147483647) == decimal("0")', + 'decimals.floor(decimal(b"", 2147483647)) == decimal("0")', + 'decimals.ceil(decimal(b"", 2147483647)) == decimal("0")', + ], +) +def test_round_accepts_the_wide_scales_within_this_clients_ceiling(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +@pytest.mark.parametrize( + "expr", + [ + # Only *expanding* a scale costs anything, and these are just past the 10**7 ceiling. + # The JDK accepts them - setScale(1e8) on 1.23 is a 100000001-digit BigDecimal, 952 MB + # measured here - and this client refuses them, by design. + 'decimals.round(decimal("1.23"), 100000000) != decimal("0")', + 'decimals.round(decimal("1.23"), 646456993) != decimal("0")', + 'decimals.round(decimal("1.23"), 1000000000) != decimal("0")', + 'decimals.round(decimal("1.23"), 2147483647) != decimal("0")', + # trunc is absent on purpose: Java early-returns when `intScale >= v.scale()`, which + # this client mirrors, so trunc only ever *coarsens* - and coarsening is free. It + # cannot reach an expanding rescale by any argument, which is why the one-argument + # forms below list it as unguarded rather than guarded. + ], +) +def test_round_rejects_the_scales_past_this_clients_ceiling(validator, expr): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule(expr), None, 1) + + +# The one-argument forms quantize to scale 0 and so are rescales too, but three of the five +# call sites did not go through the guarded helper - `round(x)`, `floor(x)` and `ceil(x)` +# reached `d.quantize(Decimal(1), ...)` directly. Each is a multi-GB allocation reachable from +# a rule that names no scale at all, which is the failure mode of a guard hung off one helper +# rather than off the operation. `trunc(x)` was safe only by accident, via its early return. +@pytest.mark.parametrize( + "expr", + [ + # Scale 0 is a coarsening for a fractional value, so the one-argument forms are + # bounded only when the value's *integer* part is what has to be built: 1e20000000 at + # scale 0 is a 20000001-digit coefficient. (Which also means these three call sites + # were unguarded for the wrong reason before - the bug was real, the demonstration + # of it was not.) + 'decimals.round(decimal("1e20000000")) != decimal("0")', + 'decimals.floor(decimal("1e20000000")) != decimal("0")', + 'decimals.ceil(decimal("1e20000000")) != decimal("0")', + ], +) +def test_the_one_argument_rounding_family_is_guarded_too(validator, expr): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule(expr), None, 1) + + +# An order of magnitude under the ceiling, the same expressions answer. +@pytest.mark.parametrize( + "expr", + [ + 'decimals.round(decimal("1e1000000")) != decimal("0")', + 'decimals.floor(decimal("1e1000000")) != decimal("0")', + 'decimals.ceil(decimal("1e1000000")) != decimal("0")', + # trunc never rescales here at all - its `scale >= current scale` early return fires + # for any value with a non-negative exponent - so it is unbounded by construction. + 'decimals.trunc(decimal("1e1000000")) != decimal("0")', + 'decimals.trunc(decimal("1e20000000")) != decimal("0")', + 'decimals.trunc(decimal("1e-2000000000"), -1000000000) == decimal("0")', + 'decimals.round(decimal("2.5")) == decimal("3")', + 'decimals.floor(decimal("-1.5")) == decimal("-2")', + 'decimals.ceil(decimal("1.5")) == decimal("2")', + 'decimals.trunc(decimal("-1.9")) == decimal("-1")', + ], +) +def test_the_one_argument_rounding_family_still_answers(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +# Rendering is the third width site, and it does not come from a rescale: `div` holds its +# coefficient to 38 digits while its exponent runs free, so the value below is cheap to +# compute and four billion characters to print. Measured: rendering a 10**8-digit value costs +# 204 MB. No zero shortcut here, unlike the rescale guard - a zero at an extreme scale renders +# as that many zeros. +@pytest.mark.parametrize( + "expr", + [ + 'string(decimals.div(decimal("1e-2147483647"), decimal("1e2147483647"))) != ""', + 'string(decimal("1e2147483647")) != ""', + 'string(decimal("1e-2147483647")) != ""', + 'string(decimal(b"", 2147483647)) != ""', + ], +) +def test_rendering_a_wide_plain_form_is_refused(validator, expr): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule(expr), None, 1) + + +def test_rendering_still_works_below_the_ceiling(validator): + assert validator.execute(rule('string(decimal("12.34")) == "12.34"'), None, 1) is True + assert validator.execute(rule('string(decimal("1e1000000")) != ""'), None, 1) is True + # The value the guard is computed from is the plain form, not the coefficient: this one + # has a single digit and a million-place exponent. + assert validator.execute(rule('string(decimal("1e-1000000")) != ""'), None, 1) is True + + +# variants.index is declared (DYN, INT) and variants.as / variants.tryAs (DYN, STRING), so a +# wrong-typed second argument fails to bind on the JVM whatever the receiver holds. Both +# checks used to come *after* the receiver was inspected, or not at all: +# +# * variants.index(anObject, 1.5) answered CEL null - the receiver was not an array, so the +# index's own type was never reached, and the argument error depended on runtime shape; +# * variants.tryAs(v, 1) stringified the 1 to "1", took the unknown-type branch and +# returned CEL null. Null is tryAs's answer for a type *mismatch*, so a call that names +# no type at all was indistinguishable from a variant of the wrong shape. +@pytest.mark.parametrize( + "expr", + [ + # A non-array receiver: the index type is still what is wrong with the call. + "variants.index(variants.parseJson('{\"a\":1}'), 1.5) == null", + "variants.index(variants.parseJson('{\"a\":1}'), true) == null", + # And an array receiver, where the check already fired. + "variants.index(variants.parseJson('[1,2]'), 1.5) == null", + "variants.index(variants.parseJson('[1,2]'), true) == null", + # A type name that is not a string. + "variants.tryAs(variants.parseJson('\"x\"'), 1) == null", + "variants.tryAs(variants.parseJson('\"x\"'), true) == null", + "variants.as(variants.parseJson('\"x\"'), 1) == 'x'", + # variants.path was the one overload in this family with no check at all: `str(path)` + # looked up the path "1" for variants.path(v, 1). Java declares it `(DYN, STRING)`, + # the same as variants.field, and Go, JS and C++ enforce that in the declared overload + # too - so a non-string path has no matching overload on any of them, whatever the + # receiver holds. Note that these three failed before the check as well, but for the + # wrong reason - "1" is a malformed JSONPath - so it is + # test_variant_path_rejects_a_non_string_path below that pins the type error itself. + "variants.path(variants.parseJson('{\"a\":1}'), 1) == null", + "variants.path(variants.parseJson('{\"a\":1}'), 1.5) == null", + "variants.path(variants.parseJson('{\"a\":1}'), true) == null", + # A receiver that is not a variant at all, which used to short-circuit to CEL null + # before the path's type was ever looked at. This one does distinguish the fix. + "variants.path(variants.tryParseJson('nope'), 1) == null", + ], +) +def test_variant_argument_types_are_checked_before_the_receiver(validator, expr): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule(expr), None, 1) + + +# The must-fail twins: the well-typed calls still work, in both receiver shapes. +@pytest.mark.parametrize( + "expr", + [ + "variants.index(variants.parseJson('[10,20,30]'), 2) != null", + "variants.index(variants.parseJson('[10,20,30]'), 9) == null", + # A non-array receiver is still CEL null, not an error, once the index type is right. + "variants.index(variants.parseJson('{\"a\":1}'), 0) == null", + "variants.tryAs(variants.parseJson('\"x\"'), 'string') == 'x'", + "variants.tryAs(variants.parseJson('\"x\"'), 'int') == null", + "variants.as(variants.parseJson('\"x\"'), 'string') == 'x'", + # A well-typed path still navigates, misses still answer CEL null, and a non-variant + # receiver is still CEL null rather than an error once the path's type is right. + "variants.as(variants.path(variants.parseJson('{\"a\":1}'), '$.a'), 'int') == 1", + "variants.path(variants.parseJson('{\"a\":1}'), '$.b') == null", + "variants.path(variants.tryParseJson('nope'), '$.a') == null", + ], +) +def test_variant_well_typed_arguments_still_work(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +# The type error itself, asserted on the message rather than through a rule, because the rule +# wrapper does not carry the detail - and because a stringified path happens to be a malformed +# JSONPath too, so "it raised" does not distinguish the two causes. The receiver is varied to +# show the check does not depend on it: Java rejects the call at compile time, so the argument +# error cannot be contingent on what the receiver holds. +@pytest.mark.parametrize( + "receiver", + [ + None, # a null receiver + celtypes.IntType(1), # not a variant at all + vu.parse_json('{"a":1}'), # a perfectly good object + vu.parse_json('[1,2]'), + ], +) +@pytest.mark.parametrize("path", [celtypes.IntType(1), celtypes.DoubleType(1.5), celtypes.BoolType(True), None]) +def test_variant_path_rejects_a_non_string_path(receiver, path): + from confluent_kafka.schema_registry.rules.cel import variant_funcs + + with pytest.raises(celpy.CELEvalError, match="variants.path: expected a string path"): + variant_funcs._path(receiver, path) + + +# The must-fail twin: a string path is accepted for every one of those receivers, and a +# non-variant receiver is still CEL null rather than an error. +def test_variant_path_accepts_a_string_path(): + from confluent_kafka.schema_registry.rules.cel import variant_funcs + + assert variant_funcs._path(vu.parse_json('{"a":1}'), "$.a") is not None + assert variant_funcs._path(vu.parse_json('{"a":1}'), "$.missing") is None + assert variant_funcs._path(None, "$.a") is None + # celtypes.StringType as well as str, since that is what a CEL literal produces. + assert variant_funcs._path(vu.parse_json('{"a":1}'), celtypes.StringType("$.a")) is not None + + +# Arithmetic is bounded by *width*, and the dividing line is not arithmetic vs. rescale - it +# is whether the operation has to align two exponents. `add` and `sub` do: the narrower +# operand is expanded into the wider one's positional frame before a single digit is computed. +# `remainder` is in the same family but is bounded by its integral quotient, which libmpdec +# short-circuits when the dividend is the smaller operand. `mul` does not align - it adds the +# exponents and multiplies the coefficients. `div` does not - it holds the coefficient to the +# context precision. Comparison does not - libmpdec short-circuits on the adjusted exponent. +# +# Measured, peak RSS, operands 1e2147483647 and 3: +# +# mul, div, <, ==, compare, min, neg, abs 13 MB +# add 1738 MB +# sub 1738 MB +# remainder 1733 MB +# add(1e2147483647, 1e-2147483647) 3125 MB +# remainder(1e-2147483647, 1e2147483647) 13 MB <- quotient is 0 +# remainder(1e2147483647, 1e2147483000) 13 MB <- quotient is 647 digits +# +# So three of six arithmetic operations reach a multi-GB allocation from a single expression +# over two operands each cheap to construct. An earlier design guarded `mul` and `div` on a +# prediction of BigDecimal's own domain errors - the operations that turn out to cost nothing +# - and this is the correction. `mul` is now unguarded; a result whose scale no int32 can +# carry is refused at the wire boundary instead, where it is actually a problem (see +# test_cel_message_transform). +@pytest.mark.parametrize( + "expr", + [ + # Alignment: the narrower operand expands into the wider one's frame. + 'decimals.add(decimal("1e2147483647"), decimal("1")) != decimal("0")', + 'decimals.add(decimal("1e-2147483647"), decimal("1")) != decimal("0")', + 'decimals.add(decimal("1e2147483647"), decimal("1e-2147483647")) != decimal("0")', + 'decimals.sub(decimal("1e2147483647"), decimal("1e-2147483647")) != decimal("0")', + # remainder, via the integral quotient it has to produce. + 'decimals.mod(decimal("1e2147483647"), decimal("3")) != decimal("0")', + 'decimals.mod(decimal("1e2147483647"), decimal("1e-2147483647")) != decimal("0")', + 'decimals.mod(decimal("1.5"), decimal("1e-2147483647")) != decimal("0")', + ], +) +def test_alignment_width_is_refused(validator, expr): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule(expr), None, 1) + + +# `decimals.mod` must be exact at any width. The Rust client computed it as +# `trunc(a/b) * b` through a division capped at its library's default 100-digit precision, and so +# returned a silently wrong residual past 100 digits (1e101 mod 3 came back as 10). libmpdec's +# remainder is exact, so this client was never affected - pinned so it stays that way, and +# because the previous coverage stopped at 1E40 (41 digits) and would not have caught it. +# +# Measured on the JDK: 10^k mod 3 is 1 for every k, and 10^200 mod 7 is 2 (10^6 = 1 mod 7, +# 200 mod 6 = 2). +@pytest.mark.parametrize( + "expr", + [ + 'string(decimals.mod(decimal("1e99"), decimal("3"))) == "1"', + 'string(decimals.mod(decimal("1e100"), decimal("3"))) == "1"', + 'string(decimals.mod(decimal("1e101"), decimal("3"))) == "1"', + 'string(decimals.mod(decimal("1e200"), decimal("3"))) == "1"', + 'string(decimals.mod(decimal("1e10000"), decimal("3"))) == "1"', + 'string(decimals.mod(decimal("1e200"), decimal("7"))) == "2"', + 'string(decimals.mod(decimal("-1e101"), decimal("3"))) == "-1"', + ], +) +def test_mod_is_exact_past_a_hundred_digits(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +# Expanding a *zero* is free, so the aligned frame is set by the operands that actually have +# digits - which means several of these turn on *which* operand expands rather than on how far +# apart the scales are. A zero also keeps whatever scale it was built with, so its adjusted +# exponent says nothing about the cost, which is what an earlier estimate got wrong. +# +# Every row measured on both libmpdec and the JDK, and they agree throughout: +# +# 0E+2e9 + 0E-2e9 free, 1 digit precision 1 +# 0E+2e9 + 1 free, 1 digit precision 1, scale 0 (the zero expands) +# 0E+2e9 mod 1E-2e9 free, 1 digit precision 1 +# 0E-2e9 mod 1E+2e9 free, 1 digit precision 1 +# 1 + 0E-2e9 1601 MB, 2e9+1 digits ArithmeticException (the *one* expands) +@pytest.mark.parametrize( + "expr", + [ + 'decimals.eq(decimals.add(decimal("0E+2000000000"), decimal("0E-2000000000")), decimal("0"))', + 'decimals.eq(decimals.sub(decimal("0E+2000000000"), decimal("0E-2000000000")), decimal("0"))', + 'decimals.eq(decimals.add(decimal("0E+2000000000"), decimal("1")), decimal("1"))', + 'decimals.eq(decimals.mod(decimal("0E+2000000000"), decimal("1E-2000000000")), decimal("0"))', + 'decimals.eq(decimals.mod(decimal("0E-2000000000"), decimal("1E+2000000000")), decimal("0"))', + 'decimals.eq(decimals.mod(decimal("0"), decimal("3")), decimal("0"))', + ], +) +def test_expanding_a_zero_operand_is_free(validator, expr): + assert validator.execute(rule(expr), None, 1) is True + + +# The row that must still be refused, because here the *one* is what expands into the zero's +# scale. A blanket zero exemption would have let this through. +def test_a_nonzero_operand_expanding_into_a_zeros_scale_is_refused(validator): + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule('decimals.add(decimal("1"), decimal("0E-2000000000")) != decimal("0")'), None, 1) + + +@pytest.mark.parametrize( + "expr, expected", + [ + # Ordinary arithmetic, unchanged. + ('string(decimals.mul(decimal("1.5"), decimal("2.5")))', "3.75"), + ('string(decimals.add(decimal("12.34"), decimal("1.5")))', "13.84"), + ('string(decimals.sub(decimal("12.34"), decimal("1.5")))', "10.84"), + ('string(decimals.mod(decimal("12.34"), decimal("1.5")))', "0.34"), + ('string(decimals.mod(decimal("1E40"), decimal("3")))', "1"), + # mul and div are not guarded at all, at any width - measured, they cost nothing. + # The first two were refused by the earlier design; both are exact and cheap. + ('decimals.mul(decimal("1e2147483647"), decimal("1e2147483647")) != decimal("0")', True), + ('decimals.mul(decimal("1e-2147483647"), decimal("1e-2147483647")) != decimal("0")', True), + ('decimals.mul(decimal("1e2147483647"), decimal("10")) != decimal("0")', True), + ('decimals.mul(decimal("1e2147483647"), decimal("1e-2147483647")) == decimal("1")', True), + ('decimals.div(decimal("1e-2147483647"), decimal("1e2147483647")) != decimal("0")', True), + # Comparison of two extreme operands, which short-circuits on the adjusted exponent. + ('decimals.lt(decimal("1e-2147483647"), decimal("1e2147483647"))', True), + ('decimal("1e2147483647") != decimal("1e-2147483647")', True), + # Alignment that stays narrow because the exponents are close, however extreme they + # both are. A guard on the operands' magnitudes rather than their difference would + # falsely refuse all of these. + ('decimals.add(decimal("1e2147483647"), decimal("1e2147483647")) != decimal("0")', True), + ('decimals.sub(decimal("1e2147483647"), decimal("1e2147483646")) != decimal("0")', True), + ('decimals.add(decimal("1e1000"), decimal("1e-1000")) != decimal("0")', True), + # remainder whose integral quotient is small, however far apart the operands are. + ('decimals.mod(decimal("1e-2147483647"), decimal("1e2147483647")) != decimal("0")', True), + ('decimals.mod(decimal("1e2147483647"), decimal("1e2147483000")) == decimal("0")', True), + ], +) +def test_the_cheap_operations_stay_unguarded(validator, expr, expected): + result = validator.execute(rule(expr), None, 1) + assert result is expected if expected is True else result == expected + + +# The coefficient arriving from the wire is the width risk on the (bytes, scale) constructor; +# the scale is not, because `scaleb` only sets the exponent. Checked from the byte count +# before `int.from_bytes` builds the integer - one byte carries about 2.41 decimal digits - +# and against the *encodable* ceiling rather than the computation one, since a coefficient +# this client cannot write back is not worth reading in. +def test_a_wide_coefficient_from_bytes_is_refused(validator): + # 4300 digits is about 1785 bytes, so this is comfortably past it without needing a + # multi-megabyte literal. + with pytest.raises(RuleError, match="Could not execute validation rule 'r'"): + validator.execute(rule('decimal(b"' + "\\x01" * 4000 + '", 0) != decimal("0")'), None, 1) + + +def test_an_ordinary_coefficient_from_bytes_still_works(validator): + assert validator.execute(rule('decimal(b"\\x04\\xd2", 2) == decimal("12.34")'), None, 1) is True + # An extreme scale on a small coefficient is fine: it only sets the exponent. + assert validator.execute(rule('decimal(b"\\x01", 2147483647) != decimal("0")'), None, 1) is True + + +# Registering a `string` / `double` entry *replaces* celpy's rather than extending it, so every +# type is the client's responsibility - and celpy's are `str` and `float`, which render or accept +# things CEL does not. Each row measured against cel-java 0.13.1. +@pytest.mark.parametrize( + "expr, expected", + [ + ('string(true)', 'true'), # was Python's "True" + ('string(false)', 'false'), + ('string(1)', '1'), + ('string(uint(3))', '3'), + ('string(1.5)', '1.5'), + ('string(b"ab")', 'ab'), + ('string("x")', 'x'), + ('string(duration("1s"))', '1s'), + ('string(timestamp(1))', '1970-01-01T00:00:01Z'), + ('string(decimal("12.30"))', '12.30'), + ('string(double("1.5"))', '1.5'), + ('string(double(decimal("100.50")))', '100.5'), + ], +) +def test_string_renders_the_declared_overloads_as_cel_does(validator, expr, expected): + assert validator.execute(rule(expr), None, 1) == expected + + +# The types CEL declares no overload for. Previously `string(null)` was "None", a list and a map +# were Python container reprs, and `double` took a bool as 0/1 and bytes as a number. +@pytest.mark.parametrize( + "expr, type_name", + [ + ('string(null)', 'null'), + ('string([1, 2])', 'list'), + ('string({"a": 1})', 'map'), + ('double(true)', 'bool'), + ('double(b"12")', 'bytes'), + ], +) +def test_string_and_double_refuse_an_undeclared_overload(validator, expr, type_name): + # A `double(...)` result is not a bool or a string, so it is wrapped to reach the validator. + wrapped = f"string({expr})" if expr.startswith("double") else expr + with pytest.raises(RuleError) as excinfo: + validator.execute(rule(wrapped), None, 1) + cause = str(excinfo.value.__cause__) + assert "found no matching overload" in cause + assert f"({type_name})" in cause diff --git a/tests/schema_registry/test_proto_builtin_deps.py b/tests/schema_registry/test_proto_builtin_deps.py new file mode 100644 index 000000000..171954915 --- /dev/null +++ b/tests/schema_registry/test_proto_builtin_deps.py @@ -0,0 +1,128 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import base64 + +import pytest +from google.protobuf import descriptor_pb2 +from google.protobuf.descriptor_pool import DescriptorPool + +from confluent_kafka.schema_registry.common.protobuf import _init_pool, _str_to_proto + + +def _schema_str(dep: str, message: str, name: str) -> str: + """A one-field schema importing ``dep``, base64-encoded the way the registry stores it.""" + fdp = descriptor_pb2.FileDescriptorProto() + fdp.name = name + fdp.package = "test" + fdp.syntax = "proto3" + fdp.dependency.append(dep) + msg = fdp.message_type.add() + msg.name = "M" + field = msg.field.add() + field.name, field.number, field.type, field.label = "f", 1, 11, 1 + field.type_name = ".confluent.type." + message + return base64.standard_b64encode(fdp.SerializeToString()).decode("ascii") + + +def _load(dep: str, message: str, name: str): + """The production path: registry text -> _str_to_proto -> pool.Add. + + ``DescriptorPool.Add()`` returns the ``FileDescriptor`` under the upb/C++ backend, but + the pure-Python backend's ``Add()`` has no ``return`` statement and always answers + ``None`` - not a bug, just an implementation detail neither backend's docstring + promises either way. The production serializer/deserializer never relies on it, always + following up with ``pool.FindFileByName()`` instead, so this does the same. + """ + pool = DescriptorPool() + _init_pool(pool) + pool.Add(_str_to_proto(name, _schema_str(dep, message, name))) + return pool.FindFileByName(name) + + +# The canonical import path is confluent/type/... - what the Java client registers and what +# ProtobufSchema declares - while the generated descriptors here were named confluent/types/... +# after the directory the Go client needs (`type` is a keyword there), which this client copied. +# A Java-registered schema used to fail with "Depends on file 'confluent/type/decimal.proto', +# but it has not been loaded". The descriptors are canonical now, and decimal's old path is a +# public-import stub registered alongside them: it declares nothing, so it re-exports +# confluent.type.Decimal without the second declaration a pool refuses ("duplicate symbol"). +@pytest.mark.parametrize( + "dep, message", + [ + ("confluent/type/decimal.proto", "Decimal"), + ("confluent/type/variant.proto", "Variant"), + ("confluent/types/decimal.proto", "Decimal"), + ], +) +def test_builtin_confluent_type_imports_resolve(dep, message): + fd = _load(dep, message, "test_%s.proto" % dep.replace("/", "_")) + assert fd.message_types_by_name["M"].fields[0].message_type.full_name == "confluent.type." + message + + +# Only decimal's old path is stubbed; anything else still has to fail, naming the import it +# could not find. Variant is in this list on purpose: it had not shipped under the old path, so +# nothing can be importing it, and pinning that makes adding a stub a deliberate act. +@pytest.mark.parametrize( + "dep", + [ + "confluent/type/nope.proto", + "confluent/types/nope.proto", + "confluent/types/variant.proto", + ], +) +def test_an_unknown_builtin_still_fails(dep): + with pytest.raises(Exception, match=dep): + _load(dep, "Decimal", "test_%s.proto" % dep.replace("/", "_")) + + +# The stub itself: registered under the old name, declaring nothing, and re-exporting the symbol +# through a public import. Each of the three is what keeps it from conflicting with the canonical +# file while still resolving - a declaration here would raise "duplicate symbol". +def test_the_legacy_stub_declares_nothing_and_reexports(): + pool = DescriptorPool() + _init_pool(pool) + + stub = pool.FindFileByName("confluent/types/decimal.proto") + assert stub.message_types_by_name == {} + assert [d.name for d in stub.public_dependencies] == ["confluent/type/decimal.proto"] + assert pool.FindMessageTypeByName("confluent.type.Decimal").file.name == "confluent/type/decimal.proto" + + +# The module path that shipped before the move. v2.15.0rc2's confluent/types/decimal_pb2.py had +# exactly two public names - DESCRIPTOR and Decimal - and both have to keep resolving there. +def test_the_old_module_path_still_exports_everything_it_shipped(): + from confluent_kafka.schema_registry.confluent.types import decimal_pb2 as legacy + + assert sorted(n for n in vars(legacy) if not n.startswith("_") and not n.startswith("confluent_dot")) == [ + "DESCRIPTOR", + "Decimal", + ] + assert legacy.Decimal(value=b"\x04\xd2", scale=2).DESCRIPTOR.full_name == "confluent.type.Decimal" + + +# What DESCRIPTOR describes did change, unavoidably: it is the stub's own file now, so it +# declares no message where the shipped one declared Decimal. Declaring it in both places is +# exactly what a pool refuses ("duplicate symbol 'confluent.type.Decimal'"), so this is the +# price of keeping the old path resolvable at all - pinned here as intended, not as drift. +# `Decimal.DESCRIPTOR` and the canonical file are where the message is found. +def test_the_old_descriptor_describes_the_stub_not_the_message(): + from confluent_kafka.schema_registry.confluent.types import decimal_pb2 as legacy + + assert legacy.DESCRIPTOR.name == "confluent/types/decimal.proto" + assert dict(legacy.DESCRIPTOR.message_types_by_name) == {} + assert [d.name for d in legacy.DESCRIPTOR.public_dependencies] == ["confluent/type/decimal.proto"] diff --git a/tests/schema_registry/test_variant_utils.py b/tests/schema_registry/test_variant_utils.py new file mode 100644 index 000000000..984643172 --- /dev/null +++ b/tests/schema_registry/test_variant_utils.py @@ -0,0 +1,678 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +# +# Copyright 2026 Confluent Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the Variant binary codec (reader, builder, and JSON conversion).""" + +import base64 +import decimal +import struct +import uuid as uuid_mod + +import pytest + +from confluent_kafka.schema_registry.confluent.type import variant_utils as vu +from confluent_kafka.schema_registry.confluent.type.variant_utils import ( + Variant, + VariantError, + VariantType, +) + +# Minimal empty metadata: version 1, offset_size 1, dictionary_size 0. Enough for any +# scalar/array value that references no object keys. +EMPTY_META = b"\x01\x00\x00" + + +def prim(code, payload=b""): + """A bare primitive Variant with the given type code and payload.""" + return Variant(bytes([code << 2]) + payload, EMPTY_META) + + +def decimal_value(code, scale, unscaled, width): + return prim(code, bytes([scale]) + unscaled.to_bytes(width, "little", signed=True)) + + +# -------------------------------------------------------------------------------------- +# parse_json + navigation +# -------------------------------------------------------------------------------------- + + +def test_parse_json_navigation_and_scalars(): + v = vu.parse_json('{"name":"alice","age":30,"scores":[10,20,30],"nested":{"x":1},"explicit":null}') + assert v.get_type() == VariantType.OBJECT + assert v.num_object_fields() == 5 + assert v.get_field_by_key("name").get_string() == "alice" + assert v.get_field_by_key("age").get_long() == 30 + assert v.get_field_by_key("scores").get_element_at_index(2).get_long() == 30 + assert v.get_field_by_key("scores").num_array_elements() == 3 + # A missing field is absent (None); an explicit JSON null is a present NULL-typed variant. + assert v.get_field_by_key("missing") is None + assert v.get_field_by_key("explicit").get_type() == VariantType.NULL + + +def test_field_by_key_binary_search_path(): + # More than the linear/binary-search threshold (32) fields exercises binary search. + obj = {("k%02d" % i): i for i in range(40)} + import json + + v = vu.parse_json(json.dumps(obj)) + assert v.get_field_by_key("k39").get_long() == 39 + assert v.get_field_by_key("k00").get_long() == 0 + assert v.get_field_by_key("k40") is None + + +def test_out_of_bounds_index_is_none(): + v = vu.parse_json("[1, 2, 3]") + assert v.get_element_at_index(0).get_long() == 1 + assert v.get_element_at_index(3) is None + assert v.get_element_at_index(-1) is None + + +# -------------------------------------------------------------------------------------- +# get_type for every primitive +# -------------------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "variant, expected", + [ + (prim(vu.NULL), VariantType.NULL), + (prim(vu.TRUE), VariantType.BOOLEAN), + (prim(vu.FALSE), VariantType.BOOLEAN), + (prim(vu.INT1, struct.pack(" INT2. + assert prim(vu.INT1, struct.pack(" INT4. + assert prim(vu.INT1, struct.pack("28 significant digits under the thread-local +# default context (prec=28). A 35-digit DECIMAL16 unscaled value at scale 5 round-trips. +def test_decimal_large_unscaled_is_not_rounded(): + unscaled = 12345678901234567890123456789012345 # 35 digits, > default prec 28 + result = decimal_value(vu.DECIMAL16, 5, unscaled, 16).get_decimal() + assert format(result, "f") == "123456789012345678901234567890.12345" + assert result == decimal.Decimal("123456789012345678901234567890.12345") + + +# -------------------------------------------------------------------------------------- +# to_json — the cross-language contract (matches the Java reference) +# -------------------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "variant, expected", + [ + # Instant (TZ): seconds always present, 'Z', 0/3/6/9 fractional grouping. + (prim(vu.TIMESTAMP, struct.pack(" raw, not "café" + ("日本語", '"日本語"'), # 日本語 -> raw + ('a"b', '"a\\"b"'), # quote still escaped + ("a\tb\nc", '"a\\tb\\nc"'), # control chars still escaped + ]: + b = vu.VariantBuilder() + b.append_string(text) + rendered = b.build().to_json() + assert rendered == expected + assert "\\u" not in rendered.replace("\\\\u", "") + + # Non-ASCII object keys must also pass through raw. + b = vu.VariantBuilder() + b.start_object() + b.append_key("café") + b.append_string("résumé") + b.end_object() + assert b.build().to_json() == '{"café":"résumé"}' + + +# -------------------------------------------------------------------------------------- +# non-finite doubles/floats (Confluent Java contract: bareword NaN/Infinity/-Infinity, +# diverging from Spark which quotes them) +# -------------------------------------------------------------------------------------- + + +def test_to_json_non_finite_double_is_bareword(): + # The builder must accept and store non-finite doubles, and to_json must emit the + # capitalized bareword tokens (no quotes, not lowercase nan/inf from str()). + for value, expected in [ + (float("nan"), "NaN"), + (float("inf"), "Infinity"), + (float("-inf"), "-Infinity"), + ]: + b = vu.VariantBuilder() + b.append_double(value) + assert b.build().to_json() == expected + + +def test_to_json_non_finite_float_is_bareword(): + for value, expected in [ + (float("nan"), "NaN"), + (float("inf"), "Infinity"), + (float("-inf"), "-Infinity"), + ]: + b = vu.VariantBuilder() + b.append_float(value) + assert b.build().to_json() == expected + + +def test_parse_json_non_finite_barewords_roundtrip(): + # Bareword non-finite literals parse (Python json.loads accepts them by default) and + # round-trip back to the same bareword tokens. + for tok in ("NaN", "Infinity", "-Infinity"): + assert vu.parse_json(tok).to_json() == tok + + +def test_parse_json_overflow_magnitude_becomes_infinity(): + # An out-of-range magnitude parses to a stored infinity and renders as the bareword. + assert vu.parse_json("1e400").to_json() == "Infinity" + + +# -------------------------------------------------------------------------------------- +# malformed input +# -------------------------------------------------------------------------------------- + + +def test_unsupported_metadata_version_raises(): + with pytest.raises(VariantError, match="version"): + Variant(b"\x00", b"\x02\x00\x00") # version 2 in metadata header + + +def test_parse_json_malformed_raises(): + with pytest.raises(ValueError): + vu.parse_json("{not json") + + +def test_parse_json_empty_or_whitespace_raises_value_error(): + # Empty/whitespace-only input must be a normal typed ValueError (json.JSONDecodeError + # is a ValueError subclass) so variants.tryParseJson catches it -> CEL null, rather + # than an unexpected crash. + for src in ("", " ", "\t\n"): + with pytest.raises(ValueError): + vu.parse_json(src) + + +def test_wrong_getter_raises(): + with pytest.raises(VariantError): + prim(vu.TRUE).get_string() + with pytest.raises(VariantError): + prim(vu.NULL).get_long() + + +# -------------------------------------------------------------------------------------- +# VariantBuilder (flat streaming writer) +# -------------------------------------------------------------------------------------- + + +def test_builder_matches_parse_json_byte_for_byte(): + # A big integer wider than 64 bits parses as a scale-0 DECIMAL16 - the one decimal + # form parse_json emits - so the programmatic decimal append can match it exactly. + big = 10**20 + src = ( + '{"id":42,"name":"hello","active":true,"score":3.5,' + '"amount":%d,"missing":null,"nums":[1,2,3],"nested":{"a":1}}' % big + ) + + b = vu.VariantBuilder() + b.start_object() + b.append_key("id") + b.append_byte(42) # parse_json encodes 42 as INT1 + b.append_key("name") + b.append_string("hello") + b.append_key("active") + b.append_boolean(True) + b.append_key("score") + b.append_double(3.5) + b.append_key("amount") + b.append_decimal((big).to_bytes(9, byteorder="big", signed=True), 0) + b.append_key("missing") + b.append_null() + b.append_key("nums") + b.start_array() + b.append_byte(1) + b.append_byte(2) + b.append_byte(3) + b.end_array() + b.append_key("nested") + b.start_object() + b.append_key("a") + b.append_byte(1) + b.end_object() + b.end_object() + built = b.build() + + parsed = vu.parse_json(src) + + # Canonical-equivalence via JSON. + assert built.to_json() == parsed.to_json() + # Byte-identical value + metadata. + assert built.value == parsed.value + assert built.metadata == parsed.metadata + + +def test_builder_native_decimal_overload_matches_bytes_overload(): + b1 = vu.VariantBuilder() + b1.append_decimal(decimal.Decimal("1.50")) + b2 = vu.VariantBuilder() + b2.append_decimal((150).to_bytes(2, byteorder="big", signed=True), 2) + assert b1.build().value == b2.build().value + assert b1.build().to_json() == "1.50" + + +def test_builder_root_scalar(): + b = vu.VariantBuilder() + b.append_long(1234567890123) + v = b.build() + assert v.get_type() == VariantType.LONG + assert v.get_long() == 1234567890123 + assert v.to_json() == "1234567890123" + + +def test_float_renders_float32_shortest(): + # Bug #7: the FLOAT case widened the f64 through the double formatter, emitting the + # f64-shortest string (e.g. "0.10000000149011612") instead of the float32-shortest + # string ("0.1") that Java Float.toString / Apache Arrow produce. + def render(f): + b = vu.VariantBuilder() + b.append_float(f) + return b.build().to_json() + + assert render(0.1) == "0.1" + assert render(0.3) == "0.3" + assert render(2.0) == "2.0" # integer ".0" preserved + + +def test_builder_append_key_outside_object_raises(): + b = vu.VariantBuilder() + with pytest.raises(VariantError): + b.append_key("x") + + +def test_builder_value_without_key_in_object_raises(): + b = vu.VariantBuilder() + b.start_object() + with pytest.raises(VariantError): + b.append_long(1) + + +def test_builder_build_with_open_container_raises(): + b = vu.VariantBuilder() + b.start_array() + b.append_long(1) + with pytest.raises(VariantError): + b.build() + + +def test_builder_unbalanced_end_raises(): + b = vu.VariantBuilder() + b.start_object() + with pytest.raises(VariantError): + b.end_array() + + +# -------------------------------------------------------------------------------------- +# variants.as('timestamp') extraction (bug #27): NANOS-precision variants must floor +# to microseconds using floor division (matching Java Math.floorDiv/floorMod), so that +# pre-epoch (negative) values round toward negative infinity rather than toward zero. +# celpy's TimestampType is datetime-backed (microsecond resolution), so the residual +# sub-microsecond nanoseconds Java keeps in its protobuf Timestamp cannot be represented. +# -------------------------------------------------------------------------------------- + +import datetime as _dt # noqa: E402 + +from confluent_kafka.schema_registry.rules.cel.variant_funcs import ( # noqa: E402 + _variant_as, + _variant_get_timestamp, +) + +_EPOCH_UTC = _dt.datetime(1970, 1, 1, tzinfo=_dt.timezone.utc) + + +def _java_nanos_to_micros(ns): + """Java TimestampUtils.fromEpochNanos split, floored to the microsecond that a + datetime can hold: sec = floorDiv(ns, 1e9), nanos = floorMod(ns, 1e9), then the + nanos field floored to micros. Equals floor(ns / 1000).""" + sec = ns // 1_000_000_000 # Math.floorDiv + nanos = ns - sec * 1_000_000_000 # Math.floorMod, 0 <= nanos < 1e9 + return sec * 1_000_000 + nanos // 1000 + + +def _nanos_variant(ns, ntz=False): + code = vu.TIMESTAMP_NANOS_NTZ if ntz else vu.TIMESTAMP_NANOS + return prim(code, struct.pack(" floors to epoch + 999, # sub-micro positive -> floors to 0 us + 1000, + 1577836800123456789, # 2020-01-01T00:00:00.123456789Z + -1, # 1 ns before epoch: floor -> -1 us (NOT 0) + -999, # sub-micro pre-epoch -> -1 us (NOT 0) + -1000, + -1500, # -1.5 us -> floor -> -2 us (NOT -1) + -1577836800123456789, # deep pre-1970 nanos timestamp + ], +) +def test_variant_as_timestamp_nanos_floors_to_micros_like_java(ns): + expected_micros = _java_nanos_to_micros(ns) + expected = _EPOCH_UTC + _dt.timedelta(microseconds=expected_micros) + # Both the TZ and NTZ nanos types extract identically (celpy carries no zone flag). + for ntz in (False, True): + result = _variant_get_timestamp(_nanos_variant(ns, ntz=ntz)) + assert result == expected + # And the full variants.as(...) dispatch path agrees. + assert _variant_as(_nanos_variant(ns, ntz=ntz), "timestamp", False) == expected + + +def test_variant_as_timestamp_nanos_uses_floor_not_truncation_for_negatives(): + # The whole point of bug #27: negative epoch nanos must floor, not truncate toward 0. + # -1 ns: floor gives -1 us; truncation toward zero would (wrongly) give 0 us. + trunc_wrong = _EPOCH_UTC + _dt.timedelta(microseconds=int(-1 / 1000)) # == epoch (0 us) + floored = _variant_get_timestamp(_nanos_variant(-1)) + assert floored == _EPOCH_UTC - _dt.timedelta(microseconds=1) + assert floored != trunc_wrong + + +def test_variant_as_timestamp_nanos_residual_precision_is_micros_only(): + # Documented type limit: datetime cannot hold sub-microsecond nanoseconds, so the + # trailing 789 ns of a NANOS value are dropped (Java keeps them in its Timestamp). + result = _variant_get_timestamp(_nanos_variant(1577836800123456789)) + assert result.microsecond == 123456 + assert result == _dt.datetime(2020, 1, 1, 0, 0, 0, 123456, tzinfo=_dt.timezone.utc) + + +def test_variant_as_timestamp_micros_types_are_used_as_is(): + # MICROS-precision variants store microseconds directly (no nanos division). + micros = 1577836800123456 + expected = _EPOCH_UTC + _dt.timedelta(microseconds=micros) + assert _variant_get_timestamp(prim(vu.TIMESTAMP, struct.pack(" 0 + assert child.value == doc.value # the whole buffer is shared + assert to_json_string(child) == "1" # the accessors honour pos + assert to_json_string(Variant(child.standalone_value_bytes(), child.metadata)) == "1" + + # A root variant is unaffected: its position is already zero. + assert doc.standalone_value_bytes() == doc.value + + +def test_the_write_back_paths_use_the_standalone_bytes(): + from confluent_kafka.schema_registry.common.avro import _variant_to_avro + from confluent_kafka.schema_registry.common.protobuf import variant_to_protobuf + from confluent_kafka.schema_registry.confluent.type.variant_utils import ( + Variant, + parse_json, + to_json_string, + ) + + doc = parse_json('{"a":1,"secret":"TOPSECRET"}') + child = doc.get_field_by_key("a") + + rec = _variant_to_avro(child, None) + assert to_json_string(Variant(rec["value"], rec["metadata"])) == "1" + + pb = variant_to_protobuf(child) + assert to_json_string(Variant(pb.value, pb.metadata)) == "1" + + # The CEL protobuf result writer is a third write-back site, and was missed when the first + # two were fixed. + from confluent_kafka.schema_registry.confluent.type import variant_pb2 + from confluent_kafka.schema_registry.rules.cel.protobuf_result_writer import _set_message + + class _Fd: + def __init__(self, name): + class _M: + pass + + self.message_type = _M() + self.message_type.full_name = name + + target = variant_pb2.Variant() + _set_message(target, _Fd("confluent.type.Variant"), child) + assert to_json_string(Variant(target.value, target.metadata)) == "1" + + # And a whole document still writes as itself. + root = _variant_to_avro(doc, None) + assert to_json_string(Variant(root["value"], root["metadata"])) == '{"a":1,"secret":"TOPSECRET"}' + + +def test_integer_size_ladder_reaches_four_bytes(): + # The format's offset_size is 1-4 bytes and the reference (VariantBuilder + # .getMinIntegerSize) has all four tiers; this helper stopped at three, so + # anything above 0xFFFFFF selected a 3-byte offset and then raised + # OverflowError from to_bytes(3) instead of encoding. + assert vu._integer_size(vu.U8_MAX) == 1 + assert vu._integer_size(vu.U8_MAX + 1) == 2 + assert vu._integer_size(vu.U16_MAX) == 2 + assert vu._integer_size(vu.U16_MAX + 1) == vu.U24_SIZE + assert vu._integer_size(vu.U24_MAX) == vu.U24_SIZE + assert vu._integer_size(vu.U24_MAX + 1) == vu.U32_SIZE + assert vu._integer_size(0xFFFFFFFF) == vu.U32_SIZE + + +def test_builder_container_past_u24_uses_four_byte_offsets(): + # A container whose data region exceeds 0xFFFFFF bytes. Reachable at the default + # 16 MiB size limit (which is 0xFFFFFF + 1) and plainly so above it. + big = "x" * (vu.U24_MAX + 1) + b = vu.VariantBuilder(size_limit=64 * 1024 * 1024) + b.start_array() + b.append_string(big) + b.append_long(7) + b.end_array() + v = b.build() + assert len(v.get_element_at_index(0).get_string()) == len(big) + assert v.get_element_at_index(1).get_long() == 7 + + +def test_builder_refuses_an_oversized_coefficient_as_a_variant_error(): + # str() past 4300 digits raises CPython's own ValueError naming an interpreter limit, + # so a caller catching VariantError saw an unrelated exception escape. + b = vu.VariantBuilder() + with pytest.raises(vu.VariantError, match="maximum precision"): + b.append_decimal(10**200000, 0) + + # The Decimal overload converts the coefficient to int on its way in, which raises the + # same interpreter ValueError one call earlier. + d = vu.VariantBuilder() + with pytest.raises(vu.VariantError, match="maximum precision"): + d.append_decimal(decimal.Decimal("9" * 5000)) + + # The boundary: a 38-digit coefficient still encodes, a 39-digit one does not. + ok = vu.VariantBuilder() + ok.start_array() + ok.append_decimal(10**38 - 1, 0) + ok.end_array() + assert ok.build().get_element_at_index(0).get_decimal() == decimal.Decimal(10**38 - 1) + with pytest.raises(vu.VariantError, match="maximum precision"): + vu.VariantBuilder().append_decimal(10**39, 0)