This is an automated email from the ASF dual-hosted git repository.
Abacn pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/master by this push:
new 83148cba130 Portable tuple logical type (#40081)
83148cba130 is described below
commit 83148cba130e3001c8815bf48462de6996792e40
Author: Yi Hu <[email protected]>
AuthorDate: Tue Sep 22 12:57:32 2026 -0400
Portable tuple logical type (#40081)
* Portable tuple logical type
* Fixed-length tuple logical type (`beam:logical_type:fixed_tuple:v1`):
Backed by a Row with anonymous fields (`f0`, `f1`, ...) and converts to/from
Python `tuple`, properly supporting heterogeneous elements (e.g. `Tuple[str,
int]`) while keep serialization exactly the same as non-portable TupleCoder.
* Variable-length tuple logical type (`beam:logical_type:var_tuple:v1`):
Backed by an Array and converts to/from Python `tuple` for `Tuple[T, ...]`.
This also keeps serialization the same as IterableCoder backing ArrayType.
* Static row encoding in `RowCoderImpl`: Added support for schema option
`beam:option:row:static_encoding`, omitting the field count and null-mask
bitset when fields are non-nullable to match `TupleCoder` wire efficiency.
* Symmetrical `_from_runner_api` on `LogicalType`: Enabled logical types to
reconstruct parameterized typing from their proto representations, keeping
`typing_from_runner_api` generic and preserving exact type annotations across
serialization boundaries.
* Add a test case for Beam SQL using python tuples
---
sdks/python/apache_beam/coders/coder_impl.pxd | 1 +
sdks/python/apache_beam/coders/coder_impl.py | 86 +++++----
sdks/python/apache_beam/coders/row_coder.py | 8 +
sdks/python/apache_beam/coders/row_coder_test.py | 49 +++++
sdks/python/apache_beam/transforms/sql_test.py | 34 ++--
.../typehints/native_type_compatibility.py | 20 +-
.../typehints/native_type_compatibility_test.py | 14 ++
sdks/python/apache_beam/typehints/row_type_test.py | 12 ++
sdks/python/apache_beam/typehints/schemas.py | 209 ++++++++++++++++-----
sdks/python/apache_beam/typehints/schemas_test.py | 11 +-
10 files changed, 341 insertions(+), 103 deletions(-)
diff --git a/sdks/python/apache_beam/coders/coder_impl.pxd
b/sdks/python/apache_beam/coders/coder_impl.pxd
index e64177e6fd3..6330a9f7b4d 100644
--- a/sdks/python/apache_beam/coders/coder_impl.pxd
+++ b/sdks/python/apache_beam/coders/coder_impl.pxd
@@ -297,6 +297,7 @@ cdef class RowCoderImpl(StreamCoderImpl):
cdef bint encoding_positions_are_trivial
cdef list components
cdef bint has_nullable_fields
+ cdef bint static_encoding
@cython.locals(i=int, nvals=libc.stdint.int64_t, running=int,
component_coder=CoderImpl,
null_mask=bytes, null_mask_c=char_ptr)
diff --git a/sdks/python/apache_beam/coders/coder_impl.py
b/sdks/python/apache_beam/coders/coder_impl.py
index 0bded25e05d..1f5f54dace8 100644
--- a/sdks/python/apache_beam/coders/coder_impl.py
+++ b/sdks/python/apache_beam/coders/coder_impl.py
@@ -61,6 +61,7 @@ from apache_beam.coders import observable
from apache_beam.coders.avro_record import AvroRecord
from apache_beam.internal import cloudpickle_pickler
from apache_beam.internal.cloudpickle import cloudpickle
+from apache_beam.typehints.schemas import _SCHEMA_OPTION_STATIC_ENCODING
from apache_beam.typehints.schemas import named_tuple_from_schema
from apache_beam.utils import proto_utils
from apache_beam.utils import windowed_value
@@ -1955,31 +1956,35 @@ class RowCoderImpl(StreamCoderImpl):
for i in self.encoding_positions)
self.has_nullable_fields = any(
field.type.nullable for field in self.schema.fields)
+ self.static_encoding = any(
+ opt.name == _SCHEMA_OPTION_STATIC_ENCODING and
+ opt.value.atomic_value.boolean for opt in self.schema.options)
def encode_to_stream(self, value, out, nested):
- out.write_var_int64(self.num_fields)
attrs = [getattr(value, name) for name in self.field_names]
- if self.has_nullable_fields:
- any_nulls = False
- for attr in attrs:
- if attr is None:
- any_nulls = True
- break
- if any_nulls:
- out.write_var_int64((self.num_fields + 7) // 8)
- # Pack the bits, little-endian, in consecutive bytes.
- running = 0
- for i, attr in enumerate(attrs):
- if i and i % 8 == 0:
- out.write_byte(running)
- running = 0
- running |= (attr is None) << (i % 8)
- out.write_byte(running)
+ if not self.static_encoding:
+ out.write_var_int64(self.num_fields)
+ if self.has_nullable_fields:
+ any_nulls = False
+ for attr in attrs:
+ if attr is None:
+ any_nulls = True
+ break
+ if any_nulls:
+ out.write_var_int64((self.num_fields + 7) // 8)
+ # Pack the bits, little-endian, in consecutive bytes.
+ running = 0
+ for i, attr in enumerate(attrs):
+ if i and i % 8 == 0:
+ out.write_byte(running)
+ running = 0
+ running |= (attr is None) << (i % 8)
+ out.write_byte(running)
+ else:
+ out.write_byte(0)
else:
out.write_byte(0)
- else:
- out.write_byte(0)
for i in range(self.num_fields):
if not self.encoding_positions_are_trivial:
@@ -2021,13 +2026,14 @@ class RowCoderImpl(StreamCoderImpl):
has_null_bits = np.zeros((n, ), dtype=np.uint8)
for k in range(n):
- out.write_var_int64(self.num_fields)
- if has_null_bits[k]:
- out.write_byte(null_bits_len)
- for i in range(null_bits_len):
- out.write_byte(null_bits[k, i])
- else:
- out.write_byte(0)
+ if not self.static_encoding:
+ out.write_var_int64(self.num_fields)
+ if has_null_bits[k]:
+ out.write_byte(null_bits_len)
+ for i in range(null_bits_len):
+ out.write_byte(null_bits[k, i])
+ else:
+ out.write_byte(0)
for i in range(self.num_fields):
if not self.encoding_positions_are_trivial:
i = self.encoding_positions_argsort[i]
@@ -2040,11 +2046,15 @@ class RowCoderImpl(StreamCoderImpl):
cython.cast(RowColumnEncoder, attrs[i]).encode_to_stream(k, out)
def decode_from_stream(self, in_stream, nested):
- nvals = in_stream.read_var_int64()
- null_mask_len = in_stream.read_var_int64()
- if null_mask_len:
- # pylint: disable=unused-variable
- null_mask_c = null_mask_py = in_stream.read(null_mask_len)
+ if self.static_encoding:
+ nvals = self.num_fields
+ null_mask_len = 0
+ else:
+ nvals = in_stream.read_var_int64()
+ null_mask_len = in_stream.read_var_int64()
+ if null_mask_len:
+ # pylint: disable=unused-variable
+ null_mask_c = null_mask_py = in_stream.read(null_mask_len)
# Note that if this coder's schema has *fewer* attributes than the encoded
# value, we just need to ignore the additional values, which will occur
@@ -2078,11 +2088,15 @@ class RowCoderImpl(StreamCoderImpl):
for k in range(n):
if in_stream.size() == 0:
break
- nvals = in_stream.read_var_int64()
- null_mask_len = in_stream.read_var_int64()
- if null_mask_len:
- # pylint: disable=unused-variable
- null_mask_c = null_mask = in_stream.read(null_mask_len)
+ if self.static_encoding:
+ nvals = self.num_fields
+ null_mask_len = 0
+ else:
+ nvals = in_stream.read_var_int64()
+ null_mask_len = in_stream.read_var_int64()
+ if null_mask_len:
+ # pylint: disable=unused-variable
+ null_mask_c = null_mask = in_stream.read(null_mask_len)
for i in range(min(self.num_fields, nvals)):
if not self.encoding_positions_are_trivial:
diff --git a/sdks/python/apache_beam/coders/row_coder.py
b/sdks/python/apache_beam/coders/row_coder.py
index 29f85ba9cbf..4c57de357b4 100644
--- a/sdks/python/apache_beam/coders/row_coder.py
+++ b/sdks/python/apache_beam/coders/row_coder.py
@@ -208,5 +208,13 @@ class LogicalTypeCoder(FastCoder):
def is_deterministic(self):
return self.representation_coder.is_deterministic()
+ def as_deterministic_coder(self, step_label, error_message=None):
+ if self.is_deterministic():
+ return self
+ return LogicalTypeCoder(
+ self.logical_type,
+ self.representation_coder.as_deterministic_coder(
+ step_label, error_message))
+
def to_type_hint(self):
return self.logical_type.language_type()
diff --git a/sdks/python/apache_beam/coders/row_coder_test.py
b/sdks/python/apache_beam/coders/row_coder_test.py
index 28170721c28..92b01241c14 100644
--- a/sdks/python/apache_beam/coders/row_coder_test.py
+++ b/sdks/python/apache_beam/coders/row_coder_test.py
@@ -33,6 +33,8 @@ from apache_beam.portability.api import schema_pb2
from apache_beam.testing.test_pipeline import TestPipeline
from apache_beam.testing.util import assert_that
from apache_beam.testing.util import equal_to
+from apache_beam.typehints.schemas import _SCHEMA_OPTION_STATIC_ENCODING
+from apache_beam.typehints.schemas import _static_encoding_option_pb2
from apache_beam.typehints.schemas import named_tuple_from_schema
from apache_beam.typehints.schemas import typing_to_runner_api
from apache_beam.utils.timestamp import Timestamp
@@ -510,6 +512,53 @@ class RowCoderTest(unittest.TestCase):
for field, a in columnar.items():
assert_array_equal(a[:n], dest[field][:n])
+ def test_row_coder_with_tuples(self):
+ class TupleRecord(typing.NamedTuple):
+ key: str
+ fixed_tuple: typing.Tuple[str, int]
+ var_tuple: typing.Tuple[int, ...]
+ homo_tuple: typing.Tuple[str, str]
+
+ coder = RowCoder(typing_to_runner_api(TupleRecord).row_type.schema)
+ record = TupleRecord("k1", ("hello", 42), (1, 2, 3), ("a", "b"))
+ encoded = coder.encode(record)
+ decoded = coder.decode(encoded)
+
+ self.assertEqual(record, decoded)
+ self.assertIsInstance(decoded.fixed_tuple, tuple)
+ self.assertIsInstance(decoded.var_tuple, tuple)
+ self.assertIsInstance(decoded.homo_tuple, tuple)
+ # Verify hashability as dict keys
+ d = {decoded.homo_tuple: "val1", decoded.fixed_tuple: "val2"}
+ self.assertEqual(d[("a", "b")], "val1")
+ self.assertEqual(d[("hello", 42)], "val2")
+ # Verify length mismatch raises ValueError instead of silent truncation
+ with self.assertRaises(ValueError):
+ coder.encode(
+ TupleRecord("k1", ("hello", 42, "extra"), (1, 2, 3), ("a", "b")))
+ with self.assertRaises(ValueError):
+ coder.encode(TupleRecord("k1", ("hello", ), (1, 2, 3), ("a", "b")))
+
+ def test_static_encoding(self):
+ schema = schema_pb2.Schema(
+ fields=[
+ schema_pb2.Field(
+ name="f_int32",
+ type=schema_pb2.FieldType(atomic_type=schema_pb2.INT32)),
+ schema_pb2.Field(
+ name="f_string",
+ type=schema_pb2.FieldType(atomic_type=schema_pb2.STRING)),
+ ],
+ options=[_static_encoding_option_pb2()])
+ RowType = named_tuple_from_schema(schema)
+ row = RowType(f_int32=42, f_string="hello world!")
+ coder = RowCoder(schema)
+ encoded = coder.encode(row)
+ # VarInt(42) = 1 byte, String("hello world!") = 1 byte len + 12 chars = 13
bytes.
+ # Total = 14 bytes (0 envelope overhead, matching TupleCoder)
+ self.assertEqual(14, len(encoded))
+ self.assertEqual(row, coder.decode(encoded))
+
if __name__ == "__main__":
logging.getLogger().setLevel(logging.INFO)
diff --git a/sdks/python/apache_beam/transforms/sql_test.py
b/sdks/python/apache_beam/transforms/sql_test.py
index 6649e210685..0d3ac699a05 100644
--- a/sdks/python/apache_beam/transforms/sql_test.py
+++ b/sdks/python/apache_beam/transforms/sql_test.py
@@ -57,6 +57,10 @@ UserTypeRow = typing.NamedTuple(
"UserTypeRow", [("id", int), ("arb", Aribitrary), ("complex", complex)])
coders.registry.register_coder(UserTypeRow, coders.RowCoder)
+TupleRow = typing.NamedTuple(
+ "TupleRow", [("id", int), ("coords", typing.Tuple[str, int])])
+coders.registry.register_coder(TupleRow, coders.RowCoder)
+
@pytest.mark.xlang_sql_expansion_service
@unittest.skipIf(
@@ -66,19 +70,12 @@ coders.registry.register_coder(UserTypeRow, coders.RowCoder)
class SqlTransformTest(unittest.TestCase):
"""Tests that exercise the cross-language SqlTransform (implemented in java).
- Note this test must be executed with pipeline options that run jobs on a
local
- job server. The easiest way to accomplish this is to run the
- `validatesCrossLanguageRunnerPythonUsingSql` gradle target for a particular
- job server, which will start the runner and job server for you. For example,
- `:runners:flink:1.13:job-server:validatesCrossLanguageRunnerPythonUsingSql`
to
- test on Flink 1.13.
-
- Alternatively, you may be able to iterate faster if you run the tests
directly
- using a runner like `FlinkRunner`, which can start a local Flink cluster and
- job server for you:
- $ pip install -e './sdks/python[gcp,test]'
+ To run these tests locally using PrismRunner, build the SQL expansion service
+ and prism binary first:
+ $ ./gradlew :sdks:java:extensions:sql:expansion-service:shadowJar
+ $ ./gradlew :runners:prism:build
$ pytest apache_beam/transforms/sql_test.py \\
- --test-pipeline-options="--runner=FlinkRunner"
+ --test-pipeline-options="--runner=PrismRunner"
"""
_multiprocess_can_split_ = True
@@ -229,6 +226,19 @@ class SqlTransformTest(unittest.TestCase):
# Verify the output matches the query (unaffected by the SET DDL)
assert_that(out, equal_to([(3, 30)]))
+ def test_tuple_field(self):
+ with TestPipeline() as p:
+ out = (
+ p
+ | beam.Create([
+ TupleRow(1, ("foo", 100)),
+ TupleRow(2, ("bar", 200)),
+ ])
+ | SqlTransform(
+ "SELECT t.id, t.coords.f0 AS `name`, t.coords.f1 AS `val`,
t.coords "
+ "FROM PCOLLECTION t WHERE t.coords.f1 > 150"))
+ assert_that(out, equal_to([(2, "bar", 200, ("bar", 200))]))
+
if __name__ == "__main__":
logging.getLogger().setLevel(logging.INFO)
diff --git a/sdks/python/apache_beam/typehints/native_type_compatibility.py
b/sdks/python/apache_beam/typehints/native_type_compatibility.py
index 7f28d57b2f8..e7e63eaf94e 100644
--- a/sdks/python/apache_beam/typehints/native_type_compatibility.py
+++ b/sdks/python/apache_beam/typehints/native_type_compatibility.py
@@ -124,15 +124,21 @@ def _safe_issubclass(derived, parent):
Returns:
issubclass(derived, parent), or False if a TypeError was raised.
"""
+ # Note: CPython behavior change between 3.12 and 3.13 in issubclass()
+ # with types.GenericAlias (e.g. tuple[int, str]):
+ # py<=3.12 return False; py>=3.13 throws TypeError. The following
+ # logic handles both scenario
try:
- return issubclass(derived, parent)
+ if issubclass(derived, parent):
+ return True
except (TypeError, AttributeError):
- if hasattr(derived, '__origin__'):
- try:
- return issubclass(derived.__origin__, parent)
- except TypeError:
- pass
- return False
+ pass
+ if hasattr(derived, '__origin__') and derived.__origin__ is not None:
+ try:
+ return issubclass(derived.__origin__, parent)
+ except (TypeError, AttributeError):
+ pass
+ return False
def _match_issubclass(match_against):
diff --git
a/sdks/python/apache_beam/typehints/native_type_compatibility_test.py
b/sdks/python/apache_beam/typehints/native_type_compatibility_test.py
index 33d6051afc7..8694cd08430 100644
--- a/sdks/python/apache_beam/typehints/native_type_compatibility_test.py
+++ b/sdks/python/apache_beam/typehints/native_type_compatibility_test.py
@@ -41,6 +41,7 @@ from apache_beam.typehints.native_type_compatibility import
convert_to_python_ty
from apache_beam.typehints.native_type_compatibility import
convert_typing_to_builtin
from apache_beam.typehints.native_type_compatibility import is_any
from apache_beam.typehints.native_type_compatibility import
match_dataclass_for_row
+from apache_beam.typehints.native_type_compatibility import _safe_issubclass
_TestNamedTuple = typing.NamedTuple(
'_TestNamedTuple', [('age', int), ('name', bytes)])
@@ -573,6 +574,19 @@ class NativeTypeCompatibilityTest(unittest.TestCase):
self.assertEqual(
compat_version == "2.73.0", match_dataclass_for_row(NonFrozenDC))
+ def test_safe_issubclass(self):
+ # In Python <= 3.12, issubclass(types.GenericAlias, tuple) returns False
+ # directly without raising TypeError, whereas in Python >= 3.13 it raises
+ # TypeError. _safe_issubclass inspects __origin__ to ensure
types.GenericAlias
+ # (e.g. tuple[...], list[...]) is recognized consistently across Python
versions.
+ self.assertTrue(_safe_issubclass(tuple[int, str], tuple))
+ self.assertTrue(_safe_issubclass(tuple[int, str], (str, tuple)))
+ self.assertTrue(_safe_issubclass(list[int], list))
+ self.assertTrue(_safe_issubclass(typing.Tuple[int, str], tuple))
+ self.assertTrue(_safe_issubclass(typing.List[int], list))
+ self.assertFalse(_safe_issubclass(int, tuple))
+ self.assertFalse(_safe_issubclass(typing.Union[int, str], tuple))
+
if __name__ == '__main__':
unittest.main()
diff --git a/sdks/python/apache_beam/typehints/row_type_test.py
b/sdks/python/apache_beam/typehints/row_type_test.py
index 30bda0cd98b..f8b1b01b492 100644
--- a/sdks/python/apache_beam/typehints/row_type_test.py
+++ b/sdks/python/apache_beam/typehints/row_type_test.py
@@ -193,6 +193,18 @@ class RowTypeTest(unittest.TestCase):
row_type.GeneratedClassRowTypeConstraint(
fields=[('id', int), ('name', str), ('name_hash', int)]))
+ def test_tuple_preserved_in_row_after_serialization(self):
+ with TestPipeline() as p:
+ res = (
+ p
+ | beam.Create([{"a": 1, "b": 2, "h": "h1"}])
+ | beam.GroupBy( # group_by with custom field generates a Beam Row
+ row_field=lambda x: (x["a"], x["b"]))
+ | beam.MapTuple(
+ lambda k, vs: (type(k.row_field), k.row_field)))
+
+ assert_that(res, equal_to([(tuple, (1, 2))]))
+
if __name__ == '__main__':
unittest.main()
diff --git a/sdks/python/apache_beam/typehints/schemas.py
b/sdks/python/apache_beam/typehints/schemas.py
index 80bee60ec95..71bd540c522 100644
--- a/sdks/python/apache_beam/typehints/schemas.py
+++ b/sdks/python/apache_beam/typehints/schemas.py
@@ -20,28 +20,31 @@
Imposes a mapping between common Python types and Beam portable schemas
(https://s.apache.org/beam-schemas)::
- Python Schema
- np.int8 <-----> BYTE
- np.int16 <-----> INT16
- np.int32 <-----> INT32
- np.int64 <-----> INT64
- int ------> INT64
- np.float32 <-----> FLOAT
- np.float64 <-----> DOUBLE
- float ------> DOUBLE
- bool <-----> BOOLEAN
- str <-----> STRING
- bytes <-----> BYTES
- ByteString ------> BYTES
- Timestamp <-----> LogicalType(urn="beam:logical_type:micros_instant:v1")
- datetime.date <---> LogicalType(urn="beam:logical_type:date:v1")
- Decimal <-----> LogicalType(urn="beam:logical_type:fixed_decimal:v1")
- Mapping <-----> MapType
- Sequence <-----> ArrayType
- NamedTuple <-----> RowType
- beam.Row ------> RowType
-
-One direction mapping of Python types from Beam portable schemas:
+ Python Schema
+ np.int8 <-----> BYTE
+ np.int16 <-----> INT16
+ np.int32 <-----> INT32
+ np.int64 <-----> INT64
+ int ------> INT64
+ np.float32 <-----> FLOAT
+ np.float64 <-----> DOUBLE
+ float ------> DOUBLE
+ bool <-----> BOOLEAN
+ str <-----> STRING
+ bytes <-----> BYTES
+ ByteString ------> BYTES
+ Timestamp <-----> LogicalType(urn="beam:logical_type:micros_instant:v1")
+ datetime.date <-----> LogicalType(urn="beam:logical_type:date:v1")
+ datetime.time <-----> LogicalType(urn="beam:logical_type:time:v1")
+ Decimal <-----> LogicalType(urn="beam:logical_type:fixed_decimal:v1")
+ Tuple[T, ...] <-----> LogicalType(urn="beam:logical_type:var_tuple:v1")
+ Tuple[T1, T2] <-----> LogicalType(urn="beam:logical_type:fixed_tuple:v1")
+ Mapping <-----> MapType
+ Sequence <-----> ArrayType
+ NamedTuple <-----> RowType
+ beam.Row ------> RowType
+
+One direction mapping of Python types from Beam portable schemas::
bytes
<------ LogicalType(urn="beam:logical_type:fixed_bytes:v1")
@@ -110,6 +113,8 @@ PYTHON_ANY_URN = "beam:logical:pythonsdk_any:v1"
_PYTHON_ANY_FIELD_TYPE_BYTE = "_pythonsdk_any_type_byte"
_PYTHON_ANY_FIELD_PAYLOAD = "payload"
_SCHEMA_OPTION_STATIC_ENCODING = "beam:option:row:static_encoding"
+FIXED_TUPLE_URN = "beam:logical_type:fixed_tuple:v1"
+VAR_TUPLE_URN = "beam:logical_type:var_tuple:v1"
# Bi-directional mappings
_PRIMITIVES = (
@@ -259,6 +264,14 @@ def schema_field(
description=description)
+def _static_encoding_option_pb2() -> schema_pb2.Option:
+ return schema_pb2.Option(
+ name=_SCHEMA_OPTION_STATIC_ENCODING,
+ type=schema_pb2.FieldType(atomic_type=schema_pb2.BOOLEAN),
+ value=schema_pb2.FieldValue(
+ atomic_value=schema_pb2.AtomicTypeValue(boolean=True)))
+
+
def _python_any_schema_pb2(has_repr):
# A portable schema matches FastPrimitivesCoder encoded values
if has_repr:
@@ -276,15 +289,7 @@ def _python_any_schema_pb2(has_repr):
type=schema_pb2.FieldType(
atomic_type=schema_pb2.BYTES, nullable=False))
],
- options=[
- schema_pb2.Option(
- name=_SCHEMA_OPTION_STATIC_ENCODING,
- type=schema_pb2.FieldType(
- atomic_type=schema_pb2.BOOLEAN),
- value=schema_pb2.FieldValue(
- atomic_value=schema_pb2.AtomicTypeValue(
- boolean=True)))
- ]))) if has_repr else None
+ options=[_static_encoding_option_pb2()]))) if has_repr else
None
else:
representation = None
@@ -374,7 +379,17 @@ class SchemaTranslation(object):
element_type=schema_pb2.FieldType(
atomic_type=PRIMITIVE_TO_ATOMIC_TYPE[int])))
- elif _safe_issubclass(type_, Sequence) and not _safe_issubclass(type_,
str):
+ elif _safe_issubclass(type_, tuple) and not match_is_named_tuple(type_):
+ arg_types = _get_args(type_)
+ if len(arg_types) == 2 and arg_types[1] is Ellipsis: # Tuple[typ, ...]
+ return self.typing_to_runner_api(VarTupleLogicalType(arg_types[0]))
+ elif len(arg_types) > 0: # Tuple[typ1, typ2, ...]
+ return self.typing_to_runner_api(FixedTupleLogicalType(arg_types))
+ else: # tuple of unknown type, just fallback to Any
+ return _python_any_schema_pb2(has_repr=True)
+
+ elif _safe_issubclass(
+ type_, Sequence) and not _safe_issubclass(type_, (str, tuple)):
arg_types = _get_args(type_)
if len(arg_types) > 0:
element_type = self.typing_to_runner_api(arg_types[0])
@@ -386,7 +401,8 @@ class SchemaTranslation(object):
return schema_pb2.FieldType(
map_type=schema_pb2.MapType(key_type=key_type,
value_type=value_type))
- elif _safe_issubclass(type_, Iterable) and not _safe_issubclass(type_,
str):
+ elif _safe_issubclass(
+ type_, Iterable) and not _safe_issubclass(type_, (str, tuple)):
arg_types = _get_args(type_)
if len(arg_types) > 0:
element_type = self.typing_to_runner_api(arg_types[0])
@@ -397,7 +413,9 @@ class SchemaTranslation(object):
return _python_any_schema_pb2(has_repr=False)
try:
- if LogicalType.is_known_logical_type(type_):
+ if isinstance(type_, LogicalType):
+ logical_type = type_
+ elif LogicalType.is_known_logical_type(type_):
logical_type = type_
else:
logical_type = LogicalType.from_typing(type_)
@@ -910,23 +928,15 @@ class LogicalType(Generic[LanguageT, RepresentationT,
ArgT]):
raise NotImplementedError()
@classmethod
- def from_runner_api(cls, logical_type_proto):
+ def _from_runner_api(cls, logical_type_proto):
# type: (schema_pb2.LogicalType) -> LogicalType
- """Construct an instance of a registered LogicalType implementation given a
- proto LogicalType.
-
- Raises ValueError if no LogicalType registered for the given URN.
+ """Construct an instance of this LogicalType implementation given a proto.
"""
- logical_type = cls._known_logical_types.get_logical_type_by_urn(
- logical_type_proto.urn)
- if logical_type is None:
- raise ValueError(
- "No logical type registered for URN '%s'" % logical_type_proto.urn)
if not logical_type_proto.HasField(
"argument_type") or not logical_type_proto.HasField("argument"):
# logical type_proto without argument
- return logical_type()
+ return cls()
else:
try:
argument = value_from_runner_api(
@@ -939,8 +949,24 @@ class LogicalType(Generic[LanguageT, RepresentationT,
ArgT]):
'Logical type %s with argument is currently unsupported. '
'Argument values are omitted',
logical_type_proto.urn)
- return logical_type()
- return logical_type(argument)
+ return cls()
+ return cls(argument)
+
+ @classmethod
+ def from_runner_api(cls, logical_type_proto):
+ # type: (schema_pb2.LogicalType) -> LogicalType
+
+ """Construct an instance of a registered LogicalType implementation given a
+ proto LogicalType.
+
+ Raises ValueError if no LogicalType registered for the given URN.
+ """
+ logical_type = cls._known_logical_types.get_logical_type_by_urn(
+ logical_type_proto.urn)
+ if logical_type is None:
+ raise ValueError(
+ "No logical type registered for URN '%s'" % logical_type_proto.urn)
+ return logical_type._from_runner_api(logical_type_proto)
@classmethod
def is_known_logical_type(cls, logical_type):
@@ -1511,3 +1537,92 @@ class JdbcTimeType(LogicalType[datetime.time,
MillisInstant, str]):
@classmethod
def _from_typing(cls, typ):
return cls()
+
+
+_TUPLE_NAMEDTUPLE_CACHE: Dict[int, type] = {}
+
+
+def _get_tuple_namedtuple(n: int) -> type:
+ cls = _TUPLE_NAMEDTUPLE_CACHE.get(n)
+ if cls is None:
+ cls = NamedTuple(f"_FixedTuple{n}", [(f"f{i}", object) for i in range(n)])
+ _TUPLE_NAMEDTUPLE_CACHE[n] = cls
+ return cls
+
+
+@LogicalType._register_internal
+class FixedTupleLogicalType(NoArgumentLogicalType[tuple, Any]):
+ """Logical type representing fixed-length Python tuples backed by a Row."""
+ def __init__(self, tuple_types: Sequence[type] = ()):
+ self._tuple_types = tuple(tuple_types)
+ self._namedtuple_cls = _get_tuple_namedtuple(len(self._tuple_types))
+
+ @classmethod
+ def urn(cls):
+ return FIXED_TUPLE_URN
+
+ def language_type(self=None):
+ if self is None or not self._tuple_types:
+ return tuple
+ return Tuple[self._tuple_types]
+
+ def representation_type(self):
+ if not self._tuple_types:
+ from apache_beam.pvalue import Row
+ return Row
+ fields = [(f"f{i}", t) for i, t in enumerate(self._tuple_types)]
+ options = []
+ st = SchemaTranslation(schema_registry=SCHEMA_REGISTRY)
+ if not any(st.typing_to_runner_api(t).nullable for t in self._tuple_types):
+ options.append((_SCHEMA_OPTION_STATIC_ENCODING, True))
+ return row_type.RowTypeConstraint.from_fields(
+ fields, schema_options=options)
+
+ def to_representation_type(self, value):
+ if len(value) != len(self._tuple_types):
+ raise ValueError(
+ f"Tuple length mismatch: expected {len(self._tuple_types)}, "
+ f"got {len(value)}")
+ return self._namedtuple_cls(*value)
+
+ def to_language_type(self, value):
+ return tuple(value)
+
+ @classmethod
+ def _from_runner_api(cls, logical_type_proto):
+ if logical_type_proto.HasField("representation"):
+ row_schema = logical_type_proto.representation.row_type.schema
+ return cls([typing_from_runner_api(f.type) for f in row_schema.fields])
+ return cls()
+
+
+@LogicalType._register_internal
+class VarTupleLogicalType(NoArgumentLogicalType[tuple, Sequence]):
+ """Logical type representing variable-length Python tuples backed by an
Array."""
+ def __init__(self, elem_type: type = object):
+ self._elem_type = elem_type
+
+ @classmethod
+ def urn(cls):
+ return VAR_TUPLE_URN
+
+ def language_type(self=None):
+ if self is None or self._elem_type is object:
+ return tuple
+ return Tuple[self._elem_type, ...]
+
+ def representation_type(self):
+ return Sequence[self._elem_type]
+
+ def to_representation_type(self, value):
+ return value
+
+ def to_language_type(self, value):
+ return tuple(value)
+
+ @classmethod
+ def _from_runner_api(cls, logical_type_proto):
+ if logical_type_proto.HasField("representation"):
+ elem_type = logical_type_proto.representation.array_type.element_type
+ return cls(typing_from_runner_api(elem_type))
+ return cls()
diff --git a/sdks/python/apache_beam/typehints/schemas_test.py
b/sdks/python/apache_beam/typehints/schemas_test.py
index 5e66a491090..5c202a3aa3d 100644
--- a/sdks/python/apache_beam/typehints/schemas_test.py
+++ b/sdks/python/apache_beam/typehints/schemas_test.py
@@ -32,6 +32,7 @@ from typing import Mapping
from typing import NamedTuple
from typing import Optional
from typing import Sequence
+from typing import Tuple
import numpy as np
import pytest
@@ -79,6 +80,11 @@ basic_map_types = [
all_primitives, all_primitives)
]
+basic_tuple_types = [
+ Tuple[str, np.int64],
+ Tuple[np.int64, ...],
+]
+
class AllPrimitives(NamedTuple):
field_int8: np.int8
@@ -109,6 +115,8 @@ class ComplexSchema(NamedTuple):
array_optional: Sequence[Optional[bool]]
timestamp: Timestamp
date: datetime.date
+ fixed_tuple: Tuple[str, np.int64]
+ var_tuple: Tuple[np.int64, ...]
def get_test_beam_fieldtype_protos():
@@ -390,7 +398,8 @@ class SchemaTest(unittest.TestCase):
@parameterized.expand([(user_type,) for user_type in
all_primitives + \
basic_array_types + \
- basic_map_types]
+ basic_map_types + \
+ basic_tuple_types]
)
def test_typing_survives_proto_roundtrip(self, user_type):
self.assertEqual(