22from collections import defaultdict , deque
33from contextlib import contextmanager
44from datetime import datetime
5- from enum import Enum
5+ from enum import Enum as PyEnum
66from logging import getLogger
77from typing import (
88 TYPE_CHECKING ,
1414)
1515
1616import csp
17- from csp import ts
17+ from csp import Enum as CspEnum , ts
1818from csp .impl .genericpushadapter import GenericPushAdapter
1919from csp .impl .types .container_type_normalizer import ContainerTypeNormalizer
2020from csp .impl .types .tstype import TsType , isTsType
@@ -75,13 +75,18 @@ def _normalize_keyby(keyby: str | tuple[str, ...] | list) -> tuple[str, ...]:
7575 return (keyby ,)
7676
7777
78+ def _has_indexer (indexer ) -> bool :
79+ # Preserve the legacy empty-string sentinel while accepting zero-valued keys.
80+ return indexer not in (None , "" )
81+
82+
7883class _SnapshotModelBaseClass (BaseModel ):
7984 model_config = ConfigDict (arbitrary_types_allowed = True , extra = "forbid" , coerce_numbers_to_str = True )
8085
8186
8287def _recursive_remove_enums (vals_dict ):
8388 for k , v in list (vals_dict .items ()):
84- is_enum_key = isinstance (k , Enum )
89+ is_enum_key = isinstance (k , ( PyEnum , CspEnum ) )
8590 is_dict_value = isinstance (v , dict )
8691 if is_enum_key :
8792 v = vals_dict .pop (k )
@@ -375,9 +380,9 @@ def _add_field_to_graph(
375380 self ._modules_connections_graph [field ] = {"getters" : [], "setters" : []}
376381
377382 if isinstance (module , str ):
378- name = f"{ module } { f'<{ indexer } >' if indexer else '' } "
383+ name = f"{ module } { f'<{ indexer } >' if _has_indexer ( indexer ) else '' } "
379384 else :
380- name = f"{ module .__class__ .__name__ } { f'<{ indexer } >' if indexer else '' } "
385+ name = f"{ module .__class__ .__name__ } { f'<{ indexer } >' if _has_indexer ( indexer ) else '' } "
381386
382387 if setting :
383388 if name not in self ._modules_connections_graph [field ]["setters" ]:
@@ -400,7 +405,7 @@ def _keys_for_channel(self, field: str) -> dict[Any, _NONE_TYPE] | None:
400405 if is_dict_basket (tstype ):
401406 # get type of key in basket
402407 basket_key_type = get_dict_basket_key_type (tstype )
403- if issubclass (basket_key_type , Enum ):
408+ if issubclass (basket_key_type , ( PyEnum , CspEnum ) ):
404409 return {e : None for e in basket_key_type }
405410 else :
406411 return self ._dynamic_keys .get (field , {})
@@ -727,7 +732,7 @@ def get_channel(
727732 # now set field to the delayed edge
728733 setattr (self , field , self ._delayed_channels [field ])
729734
730- if is_dict_basket (tstype ) and indexer :
735+ if is_dict_basket (tstype ) and _has_indexer ( indexer ) :
731736 # if using an indexer, return that edge (raise if not recognized)
732737 if indexer not in self ._keys_for_channel (field ):
733738 raise GatewayException (
@@ -770,11 +775,11 @@ def set_channel(
770775 gateway_tstype = tstype
771776
772777 # validate arguments
773- if _is_dict_basket and isinstance (edge , Edge ) and not indexer :
778+ if _is_dict_basket and isinstance (edge , Edge ) and not _has_indexer ( indexer ) :
774779 # if its a dict basket and you set an edge, you need to provide an indexer
775780 raise GatewayException (f"Field `{ field } ` refers to a dict basket, and you have provided an edge { edge } but not an indexer" )
776781
777- if _is_dict_basket and isinstance (edge , dict ) and indexer :
782+ if _is_dict_basket and isinstance (edge , dict ) and _has_indexer ( indexer ) :
778783 # if its a dict basket and you set a dict, you should not provide an indexer
779784 raise GatewayException (f"Field `{ field } ` refers to a dict basket, and you have provided an edge basket but also an indexer { indexer } " )
780785
@@ -786,7 +791,7 @@ def set_channel(
786791 # must provide an edge type
787792 raise TypeError (f"Edge provided for field `{ field } ` is not an `Edge` instance { edge } " )
788793
789- if not _is_dict_basket and indexer :
794+ if not _is_dict_basket and _has_indexer ( indexer ) :
790795 # don't provide an indexer for non-dict basket
791796 raise GatewayException (f"Indexer provided for field `{ field } ` but it is not a basket instance" )
792797
@@ -967,7 +972,7 @@ def _set_last(self, field: str, indexer: str | int | None = None) -> None:
967972 tstype = self .get_outer_type (field )
968973
969974 # First ensure edge is constructed
970- if indexer :
975+ if _has_indexer ( indexer ) :
971976 edge = self .get_channel (field , indexer )
972977 else :
973978 edge = self .get_channel (field )
@@ -990,7 +995,7 @@ def _set_last(self, field: str, indexer: str | int | None = None) -> None:
990995 trigger = ConcurrentFutureAdapter (name = f"RequestLast<{ edge_type_name } >" )
991996
992997 if is_dict_basket (tstype ):
993- if indexer :
998+ if _has_indexer ( indexer ) :
994999 named_on_request_node_dict_basket (f"QueryLast<Basket<{ edge_type_name } >>" )({indexer : edge }, trigger .out ())
9951000 else :
9961001 named_on_request_node_dict_basket (f"QueryLast<Basket<{ edge_type_name } >>" )(edge , trigger .out ())
@@ -1009,7 +1014,7 @@ def _set_next(self, field: str, indexer: str | int | None = None) -> None:
10091014 tstype = self .get_outer_type (field )
10101015
10111016 # First ensure edge is constructed
1012- if indexer :
1017+ if _has_indexer ( indexer ) :
10131018 edge = self .get_channel (field , indexer )
10141019 else :
10151020 edge = self .get_channel (field )
@@ -1032,7 +1037,7 @@ def _set_next(self, field: str, indexer: str | int | None = None) -> None:
10321037 trigger = ConcurrentFutureAdapter (name = f"RequestLast<{ edge_type_name } >" )
10331038
10341039 if is_dict_basket (tstype ):
1035- if indexer :
1040+ if _has_indexer ( indexer ) :
10361041 named_wait_for_next_node_dict_basket (f"QueryNext<Basket<{ edge_type_name } >>" )({indexer : edge }, trigger .out ())
10371042 else :
10381043 named_wait_for_next_node_dict_basket (f"QueryNext<Basket<{ edge_type_name } >>" )(edge , trigger .out ())
@@ -1067,7 +1072,7 @@ def add_send_channel(self, field: str, indexer: str | int | None = None) -> None
10671072 tstype = self .get_outer_type (field )
10681073
10691074 if is_dict_basket (tstype ):
1070- if indexer :
1075+ if _has_indexer ( indexer ) :
10711076 # wire in now
10721077 tstype = get_dict_basket_value_type (tstype )
10731078 self ._send_channels [field , indexer ] = GenericPushAdapter (tstype , name = f"manual_{ field } [{ indexer } ]" )
@@ -1081,7 +1086,7 @@ def add_send_channel(self, field: str, indexer: str | int | None = None) -> None
10811086 tstype = tstype .typ
10821087 self ._send_channels [field , indexer ] = GenericPushAdapter (tstype , name = f"manual_{ field } " )
10831088
1084- def _add_send_channel_dict_basket (self , field : str , keys : list [str ] | Enum ) -> None :
1089+ def _add_send_channel_dict_basket (self , field : str , keys : list [str ] | type [ PyEnum ] | type [ CspEnum ] ) -> None :
10851090 # NOTE: Do not call this directly, it is used in the factory finalization
10861091
10871092 # get type of edge
@@ -1118,7 +1123,7 @@ def last(self, field: str, indexer: str | int | None = None, *, timeout=None) ->
11181123
11191124 # wait for result
11201125 result = future .result (timeout = timeout )
1121- if indexer :
1126+ if _has_indexer ( indexer ) :
11221127 return result .get (indexer )
11231128 return result
11241129
@@ -1134,7 +1139,7 @@ def next(self, field: str, indexer: str | int | None = None, *, timeout=None) ->
11341139
11351140 # wait for result
11361141 result = future .result (timeout = timeout )
1137- if indexer :
1142+ if _has_indexer ( indexer ) :
11381143 return result .get (indexer )
11391144 return result
11401145
@@ -1331,7 +1336,7 @@ def stage_lookup(
13311336 def _check (self , field : str , where : dict , kind : str , indexer : str | int | None = None ) -> None :
13321337 if (field , indexer ) not in where :
13331338 # TODO should only be called once the graph is started
1334- raise NoProviderException ("Nobody provides {}: {}{}" .format (kind , field , f"-{ indexer } " if indexer else "" ))
1339+ raise NoProviderException ("Nobody provides {}: {}{}" .format (kind , field , f"-{ indexer } " if _has_indexer ( indexer ) else "" ))
13351340
13361341 def override (self , field : str , value : Any ) -> None :
13371342 raise NotImplementedError ()
0 commit comments