Skip to content

Commit 9e50fec

Browse files
authored
Merge pull request #322 from Point72/investigate/csp-pr743-intenum-compat-v3
Support CSP IntEnum compatibility
2 parents 4f41f3e + d0f2ba6 commit 9e50fec

9 files changed

Lines changed: 86 additions & 29 deletions

File tree

csp_gateway/server/gateway/csp/channels.py

Lines changed: 24 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
from collections import defaultdict, deque
33
from contextlib import contextmanager
44
from datetime import datetime
5-
from enum import Enum
5+
from enum import Enum as PyEnum
66
from logging import getLogger
77
from typing import (
88
TYPE_CHECKING,
@@ -14,7 +14,7 @@
1414
)
1515

1616
import csp
17-
from csp import ts
17+
from csp import Enum as CspEnum, ts
1818
from csp.impl.genericpushadapter import GenericPushAdapter
1919
from csp.impl.types.container_type_normalizer import ContainerTypeNormalizer
2020
from 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+
7883
class _SnapshotModelBaseClass(BaseModel):
7984
model_config = ConfigDict(arbitrary_types_allowed=True, extra="forbid", coerce_numbers_to_str=True)
8085

8186

8287
def _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()

csp_gateway/server/gateway/csp/factory.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -103,7 +103,7 @@ def build(self, channels: ChannelsType) -> ChannelsType:
103103

104104
# second pass to finish connecting in wires
105105
for (field, indexer), push_adapter in channels._send_channels.items():
106-
with channels._connection_context(f"Send[{field}]{f'<{indexer}>' if indexer else ''}"):
106+
with channels._connection_context(f"Send[{field}]{f'<{indexer}>' if indexer not in (None, '') else ''}"):
107107
# Add it as an edge on the StreamGroup
108108
if isinstance(push_adapter, GenericPushAdapter):
109109
channels.set_channel(field, push_adapter.out(), indexer=indexer)

csp_gateway/server/shared/json_converter.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -93,7 +93,7 @@ def _create_snapshot_dict(all_data: list[ChannelValueModel]) -> dict[str, Any]:
9393
for data in all_data:
9494
channel = data.channel
9595
value = _convert_orjson_compatible(data.value)
96-
if key := data.dict_basket_key:
96+
if (key := data.dict_basket_key) not in (None, ""):
9797
res[channel][_convert_orjson_compatible(key)] = value
9898
else:
9999
res[channel] = value

csp_gateway/testing/harness.py

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -191,7 +191,12 @@ class GatewayAssertAttrUnsetEvent(BaseGatewayTestEvent):
191191
attr: str
192192

193193
def apply(self, now, values, tick_counts, *args, **kwargs):
194-
assert hasattr(values[self.channel], self.attr) == (not self.unset)
194+
value = values[self.channel]
195+
if hasattr(value, "_implicitly_unset"):
196+
is_set = self.attr not in value._implicitly_unset() and getattr(value, self.attr, None) is not None
197+
else:
198+
is_set = hasattr(value, self.attr)
199+
assert is_set == (not self.unset)
195200

196201

197202
class GatewayAssertLenEvent(BaseGatewayTestEvent):
@@ -256,7 +261,12 @@ class GatewayAssertIdxAttrUnsetEvent(BaseGatewayTestEvent):
256261
attr: str
257262

258263
def apply(self, now, values, tick_counts, *args, **kwargs):
259-
assert hasattr(values[self.channel][self.idx], self.attr) == (not self.unset)
264+
value = values[self.channel][self.idx]
265+
if hasattr(value, "_implicitly_unset"):
266+
is_set = self.attr not in value._implicitly_unset() and getattr(value, self.attr, None) is not None
267+
else:
268+
is_set = hasattr(value, self.attr)
269+
assert is_set == (not self.unset)
260270

261271

262272
class GatewayAssertTickedEvents(BaseGatewayTestEvent):

csp_gateway/tests/server/gateway/csp/test_channels.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,12 @@ class MyGatewayChannels(GatewayChannels):
6666
my_array_channel: ts[Numpy1DArray[float]] = None
6767

6868

69+
def test_csp_enum_dict_basket_has_static_keys():
70+
channels = MyGatewayChannels()
71+
72+
assert set(channels._keys_for_channel("my_enum_basket")) == set(MyEnum)
73+
74+
6975
class DerivedChannels(MyGatewayChannels):
7076
pass
7177

csp_gateway/tests/server/gateway/test_gateway.py

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
import time
44
from collections.abc import Callable
55
from datetime import datetime, timedelta
6-
from enum import Enum
6+
from enum import IntEnum
77
from io import StringIO
88
from typing import Annotated, Any
99

@@ -29,10 +29,14 @@
2929
from csp_gateway.utils import NoProviderException
3030

3131

32-
class MyEnum(Enum):
32+
class MyEnum(IntEnum):
33+
ZERO = 0
3334
ONE = 1
3435
TWO = 2
3536

37+
def __str__(self):
38+
return f"{type(self).__name__}.{self.name}"
39+
3640

3741
class MyStruct(GatewayStruct):
3842
foo: float
@@ -89,16 +93,19 @@ def connect(self, channels: MyGatewayChannels) -> None:
8993
channels.set_channel(MyGatewayChannels.my_array_channel, csp.const(np.array([1.0, 2.0])))
9094

9195
if self.by_key:
96+
channels.set_channel(MyGatewayChannels.my_enum_basket, self.my_data, MyEnum.ZERO)
9297
channels.set_channel(MyGatewayChannels.my_enum_basket, self.my_data, MyEnum.ONE)
9398
channels.set_channel(MyGatewayChannels.my_enum_basket, self.my_data2, MyEnum.TWO)
9499
channels.set_channel(MyGatewayChannels.my_str_basket, self.my_data, "my_key")
95100
channels.set_channel(MyGatewayChannels.my_str_basket, self.my_data2, "my_key2")
101+
channels.set_channel(MyGatewayChannels.my_enum_basket_list, self.my_list_data, MyEnum.ZERO)
96102
channels.set_channel(MyGatewayChannels.my_enum_basket_list, self.my_list_data, MyEnum.ONE)
97103
channels.set_channel(MyGatewayChannels.my_enum_basket_list, self.my_list_data, MyEnum.TWO)
98104
else:
99105
channels.set_channel(
100106
MyGatewayChannels.my_enum_basket,
101107
{
108+
MyEnum.ZERO: self.my_data,
102109
MyEnum.ONE: self.my_data,
103110
MyEnum.TWO: self.my_data2,
104111
},
@@ -113,6 +120,7 @@ def connect(self, channels: MyGatewayChannels) -> None:
113120
channels.set_channel(
114121
MyGatewayChannels.my_enum_basket_list,
115122
{
123+
MyEnum.ZERO: self.my_list_data,
116124
MyEnum.ONE: self.my_list_data,
117125
MyEnum.TWO: self.my_list_data,
118126
},
@@ -185,6 +193,10 @@ def connect(self, channels: MyGatewayChannels) -> None:
185193
csp.add_graph_output(f"my_enum_basket_list[{k}]", v)
186194

187195
# Get by indexer
196+
csp.add_graph_output(
197+
"my_enum_basket_ZERO",
198+
channels.get_channel(MyGatewayChannels.my_enum_basket, MyEnum.ZERO),
199+
)
188200
csp.add_graph_output(
189201
"my_enum_basket_ONE",
190202
channels.get_channel(MyGatewayChannels.my_enum_basket, MyEnum.ONE),
@@ -585,7 +597,9 @@ def test_last(by_key):
585597

586598
output = gateway.channels.last("my_enum_basket")
587599
assert isinstance(output, dict)
588-
assert len(output) == 2
600+
assert len(output) == 3
601+
output = gateway.channels.last("my_enum_basket", MyEnum.ZERO)
602+
assert isinstance(output, MyStruct)
589603
output = gateway.channels.last("my_enum_basket", MyEnum.ONE)
590604
assert isinstance(output, MyStruct)
591605

@@ -598,7 +612,7 @@ def test_last(by_key):
598612

599613
output = gateway.channels.last("my_enum_basket_list")
600614
assert isinstance(output, dict)
601-
assert len(output) == 2
615+
assert len(output) == 3
602616
finally:
603617
gateway.stop()
604618

csp_gateway/tests/server/modules/io/test_json_converter.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
from datetime import datetime, timedelta, timezone
2+
from enum import IntEnum
23
from typing import Annotated, Any
34

45
import csp
@@ -290,6 +291,17 @@ def test_convert_orjson_compatible():
290291
assert _convert_orjson_compatible(my_str) == my_str
291292

292293

294+
def test_create_snapshot_dict_with_zero_valued_enum_key():
295+
class ZeroBasedEnum(IntEnum):
296+
ZERO = 0
297+
298+
value = MyStruct(foo=1.0)
299+
snapshot_dict = _create_snapshot_dict(
300+
[CVM(channel="enum_basket", value=value, dict_basket_key=ZeroBasedEnum.ZERO, timestamp=datetime(2020, 1, 1))]
301+
)
302+
assert snapshot_dict["enum_basket"] == {"ZERO": _convert_orjson_compatible(value)}
303+
304+
293305
def test_parse_snapshot_dict():
294306
timestamp = datetime(2020, 1, 1)
295307
dummy_id = "9"

csp_gateway/tests/server/web/test_webserver.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -283,7 +283,7 @@ def test_csp_last_basket(self, rest_client: TestClient):
283283
assert len(data) == 3
284284

285285
for channel_name, member in ExampleEnum.__members__.items():
286-
multiplier = member.value
286+
multiplier = getattr(member, "value", member)
287287
response = rest_client.get(f"/api/v1/last/basket/{channel_name}?token=test")
288288
assert response.status_code == 200
289289

@@ -345,7 +345,7 @@ def test_csp_next_basket(self, rest_client: TestClient):
345345
assert len(data) == 3
346346

347347
for channel_name, member in ExampleEnum.__members__.items():
348-
multiplier = member.value
348+
multiplier = getattr(member, "value", member)
349349
response = rest_client.get(f"/api/v1/next/basket/{channel_name}?token=test")
350350
assert response.status_code == 200
351351

csp_gateway/tests/test_harness.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,16 @@
88
from csp_gateway.testing.shared_helpful_classes import MyGateway, MyGatewayChannels, MyStruct
99

1010

11+
def test_assert_attr_unset_uses_gateway_struct_field_presence():
12+
h = GatewayTestHarness(test_channels=[MyGatewayChannels.my_channel])
13+
h.send(MyGatewayChannels.my_channel, MyStruct())
14+
h.assert_attr_unset(MyGatewayChannels.my_channel, "foo")
15+
h.assert_attr_unset(MyGatewayChannels.my_channel, "my_flag", unset=False)
16+
17+
gateway = MyGateway(modules=[h], channels=MyGatewayChannels())
18+
csp.run(gateway.graph, starttime=datetime(2020, 1, 1), endtime=timedelta(0))
19+
20+
1121
@pytest.mark.parametrize("make_invalid", (True, False))
1222
def test_delay(make_invalid):
1323
channels = [

0 commit comments

Comments
 (0)