7676except ImportError :
7777 dill = None
7878
79+ try :
80+ import zoneinfo
81+ except ImportError :
82+ zoneinfo = None
83+
7984if TYPE_CHECKING :
8085 import proto
8186
@@ -343,10 +348,13 @@ def decode(self, value):
343348BYTES_TYPE = 3
344349UNICODE_TYPE = 4
345350BOOL_TYPE = 9
351+ DATETIME_TYPE = 11
352+ DATE_TYPE = 12
346353LIST_TYPE = 5
347354TUPLE_TYPE = 6
348355DICT_TYPE = 7
349356SET_TYPE = 8
357+ FROZENSET_TYPE = 13
350358ITERABLE_LIKE_TYPE = 10
351359
352360PROTO_TYPE = 100
@@ -445,6 +453,22 @@ def encode_to_stream(self, value, stream, nested):
445453 elif t is bool :
446454 stream .write_byte (BOOL_TYPE )
447455 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" ))
448472 elif t in _ITERABLE_LIKE_TYPES :
449473 stream .write_byte (ITERABLE_LIKE_TYPE )
450474 self .iterable_coder_impl .encode_to_stream (value , stream , nested )
@@ -467,8 +491,11 @@ def encode_to_stream(self, value, stream, nested):
467491 for k , v in dict_value .items ():
468492 self .encode_to_stream (k , stream , True )
469493 self .encode_to_stream (v , stream , True )
470- elif t is set :
471- 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 )
472499 stream .write_var_int64 (len (value ))
473500 if self .requires_deterministic_step_label is not None :
474501 try :
@@ -603,13 +630,15 @@ def decode_from_stream(self, stream, nested):
603630 return stream .read_all (nested )
604631 elif t == UNICODE_TYPE :
605632 return stream .read_all (nested ).decode ("utf-8" )
606- 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 :
607634 vlen = stream .read_var_int64 ()
608635 vlist = [self .decode_from_stream (stream , True ) for _ in range (vlen )]
609636 if t == LIST_TYPE :
610637 return vlist
611638 elif t == TUPLE_TYPE :
612639 return tuple (vlist )
640+ elif t == FROZENSET_TYPE :
641+ return frozenset (vlist )
613642 return set (vlist )
614643 elif t == DICT_TYPE :
615644 vlen = stream .read_var_int64 ()
@@ -620,6 +649,46 @@ def decode_from_stream(self, stream, nested):
620649 return v
621650 elif t == BOOL_TYPE :
622651 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" ))
623692 elif t == ITERABLE_LIKE_TYPE :
624693 return self .iterable_coder_impl .decode_from_stream (stream , nested )
625694 elif t == PROTO_TYPE :
0 commit comments