Skip to content

Commit 50bc1df

Browse files
committed
Add support for datetime, date, and frozenset to FastPrimitivesCoder
1 parent cf5cc27 commit 50bc1df

2 files changed

Lines changed: 114 additions & 3 deletions

File tree

sdks/python/apache_beam/coders/coder_impl.py

Lines changed: 73 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@
3333

3434
# ruff: noqa: UP006
3535
import dataclasses
36+
import datetime
3637
import decimal
3738
import enum
3839
import itertools
@@ -75,6 +76,11 @@
7576
except ImportError:
7677
dill = None
7778

79+
try:
80+
import zoneinfo
81+
except ImportError:
82+
zoneinfo = None
83+
7884
if TYPE_CHECKING:
7985
import proto
8086

@@ -342,10 +348,13 @@ def decode(self, value):
342348
BYTES_TYPE = 3
343349
UNICODE_TYPE = 4
344350
BOOL_TYPE = 9
351+
DATETIME_TYPE = 11
352+
DATE_TYPE = 12
345353
LIST_TYPE = 5
346354
TUPLE_TYPE = 6
347355
DICT_TYPE = 7
348356
SET_TYPE = 8
357+
FROZENSET_TYPE = 13
349358
ITERABLE_LIKE_TYPE = 10
350359

351360
PROTO_TYPE = 100
@@ -444,6 +453,22 @@ def encode_to_stream(self, value, stream, nested):
444453
elif t is bool:
445454
stream.write_byte(BOOL_TYPE)
446455
stream.write_byte(value)
456+
elif t is datetime.datetime:
457+
# We use RFC 9557 for lossless encoding of timezone info.
458+
stream.write_byte(DATETIME_TYPE)
459+
stream.write(value.isoformat().encode("utf-8"))
460+
if (zoneinfo is not None and value.tzinfo is not None and
461+
type(value.tzinfo) is not datetime.timezone):
462+
stream.write(f"[{value.tzinfo}]".encode("utf-8"))
463+
if type(
464+
value.tzinfo) is datetime.timezone and (tzname :=
465+
value.tzname()) is not None:
466+
stream.write(f"[tzn={tzname}]".encode("utf-8"))
467+
if value.fold != 0:
468+
stream.write(f"[f={value.fold}]".encode("utf-8"))
469+
elif t is datetime.date:
470+
stream.write_byte(DATE_TYPE)
471+
stream.write(value.isoformat().encode("utf-8"))
447472
elif t in _ITERABLE_LIKE_TYPES:
448473
stream.write_byte(ITERABLE_LIKE_TYPE)
449474
self.iterable_coder_impl.encode_to_stream(value, stream, nested)
@@ -466,8 +491,11 @@ def encode_to_stream(self, value, stream, nested):
466491
for k, v in dict_value.items():
467492
self.encode_to_stream(k, stream, True)
468493
self.encode_to_stream(v, stream, True)
469-
elif t is set:
470-
stream.write_byte(SET_TYPE)
494+
elif t is set or t is frozenset:
495+
if t is set:
496+
stream.write_byte(SET_TYPE)
497+
else:
498+
stream.write_byte(FROZENSET_TYPE)
471499
stream.write_var_int64(len(value))
472500
if self.requires_deterministic_step_label is not None:
473501
try:
@@ -602,13 +630,15 @@ def decode_from_stream(self, stream, nested):
602630
return stream.read_all(nested)
603631
elif t == UNICODE_TYPE:
604632
return stream.read_all(nested).decode("utf-8")
605-
elif t == LIST_TYPE or t == TUPLE_TYPE or t == SET_TYPE:
633+
elif t == LIST_TYPE or t == TUPLE_TYPE or t == SET_TYPE or t == FROZENSET_TYPE:
606634
vlen = stream.read_var_int64()
607635
vlist = [self.decode_from_stream(stream, True) for _ in range(vlen)]
608636
if t == LIST_TYPE:
609637
return vlist
610638
elif t == TUPLE_TYPE:
611639
return tuple(vlist)
640+
elif t == FROZENSET_TYPE:
641+
return frozenset(vlist)
612642
return set(vlist)
613643
elif t == DICT_TYPE:
614644
vlen = stream.read_var_int64()
@@ -619,6 +649,46 @@ def decode_from_stream(self, stream, nested):
619649
return v
620650
elif t == BOOL_TYPE:
621651
return not not stream.read_byte()
652+
elif t == DATETIME_TYPE:
653+
rfc_9557_str = stream.read_all(nested).decode("utf-8")
654+
first_tag_idx = rfc_9557_str.find("[")
655+
if first_tag_idx == -1:
656+
return datetime.datetime.fromisoformat(rfc_9557_str)
657+
658+
base_iso = rfc_9557_str[:first_tag_idx]
659+
tags_str = rfc_9557_str[first_tag_idx:]
660+
dt = datetime.datetime.fromisoformat(base_iso)
661+
662+
fold = 0
663+
zone_name = None
664+
tz_name = None
665+
666+
tags = tags_str.replace("]", "").split("[")
667+
for tag in tags:
668+
if not tag:
669+
continue
670+
if tag.startswith("f="):
671+
fold = int(tag[2:])
672+
elif tag.startswith("tzn="):
673+
tz_name = tag[4:]
674+
elif "=" in tag:
675+
# Skip unknown tags like [knort=blorgel]
676+
continue
677+
else:
678+
zone_name = tag
679+
680+
if tz_name and (offset := dt.utcoffset()) is not None:
681+
dt = dt.replace(tzinfo=datetime.timezone(offset=offset, name=tz_name))
682+
elif zoneinfo is not None and zone_name:
683+
dt = dt.replace(tzinfo=zoneinfo.ZoneInfo(zone_name))
684+
685+
if fold != dt.fold:
686+
dt = dt.replace(fold=fold)
687+
688+
return dt
689+
elif t == DATE_TYPE:
690+
return datetime.date.fromisoformat(
691+
stream.read_all(nested).decode("utf-8"))
622692
elif t == ITERABLE_LIKE_TYPE:
623693
return self.iterable_coder_impl.decode_from_stream(stream, nested)
624694
elif t == PROTO_TYPE:

sdks/python/apache_beam/coders/coders_test_common.py

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121

2222
import base64
2323
import collections
24+
import datetime
2425
import enum
2526
import logging
2627
import math
@@ -67,6 +68,11 @@
6768
except ImportError:
6869
dill = None
6970

71+
try:
72+
import zoneinfo
73+
except ImportError:
74+
zoneinfo = None
75+
7076
MyNamedTuple = collections.namedtuple("A", ["x", "y"]) # type: ignore[name-match]
7177
AnotherNamedTuple = collections.namedtuple("AnotherNamedTuple", ["x", "y"])
7278
MyTypedNamedTuple = NamedTuple("MyTypedNamedTuple", [("f1", int), ("f2", str)])
@@ -431,6 +437,41 @@ def test_bytes_coder(self):
431437
def test_bool_coder(self):
432438
self.check_coder(coders.BooleanCoder(), True, False)
433439

440+
def test_fast_primitives_coder_datetime(self):
441+
self.check_coder(
442+
coders.FastPrimitivesCoder(),
443+
datetime.datetime(2026, 1, 1),
444+
datetime.datetime(
445+
2025,
446+
2,
447+
3,
448+
tzinfo=datetime.timezone(datetime.timedelta(hours=3, minutes=30))),
449+
datetime.datetime(
450+
2025,
451+
2,
452+
3,
453+
tzinfo=datetime.timezone(datetime.timedelta(hours=3), name="Foo")),
454+
# Nonsense tznaive fold is still preserved.
455+
datetime.datetime(2026, 11, 1, 1, 30, fold=1),
456+
)
457+
if zoneinfo is not None:
458+
tz = zoneinfo.ZoneInfo("America/New_York")
459+
self.check_coder(
460+
coders.FastPrimitivesCoder(),
461+
datetime.datetime(2026, 11, 1, 1, 30, tzinfo=tz, fold=0),
462+
datetime.datetime(2026, 11, 1, 1, 30, tzinfo=tz, fold=1),
463+
)
464+
465+
def test_fast_primitives_coder_date(self):
466+
self.check_coder(
467+
coders.FastPrimitivesCoder(),
468+
datetime.date(2026, 1, 1),
469+
)
470+
471+
def test_fast_primitives_coder_frozenset(self):
472+
self.check_coder(
473+
coders.FastPrimitivesCoder(), frozenset(), frozenset(["a", "b", "c"]))
474+
434475
def test_varint_coder(self):
435476
# Small ints.
436477
self.check_coder(coders.VarIntCoder(), *range(-10, 10))

0 commit comments

Comments
 (0)