Skip to content

Commit 3b324c3

Browse files
committed
Improve _EnumType and _Enum
1 parent 432b902 commit 3b324c3

3 files changed

Lines changed: 51 additions & 10 deletions

File tree

scripts/update-declarations.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -36,10 +36,7 @@ def __new__(metacls, name, bases, namespace):
3636
members = namespace["_members_"]
3737
3838
namespace["_reverse_map_"] = {v: k for k, v in members.items()}
39-
cls = type(ctypes.c_int32).__new__(metacls, name, bases, namespace)
40-
for key, value in cls._members_.items():
41-
globals()[key] = value
42-
return cls
39+
return type(ctypes.c_int32).__new__(metacls, name, bases, namespace)
4340
4441
def __repr__(self):
4542
return f"<Enum {self.__name__}>"
@@ -55,8 +52,12 @@ def __repr__(self):
5552
def __eq__(self, other):
5653
if isinstance(other, int):
5754
return self.value == other
55+
if type(self) is type(other):
56+
return self.value == other.value
57+
return NotImplemented
5858
59-
return type(self) is type(other) and self.value == other.value
59+
def __hash__(self):
60+
return hash(self.value)
6061
"""
6162

6263

src/ctypes_dlpack/_c_api.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -20,10 +20,7 @@ def __new__(metacls, name, bases, namespace):
2020
members = namespace["_members_"]
2121

2222
namespace["_reverse_map_"] = {v: k for k, v in members.items()}
23-
cls = type(ctypes.c_int32).__new__(metacls, name, bases, namespace)
24-
for key, value in cls._members_.items():
25-
globals()[key] = value
26-
return cls
23+
return type(ctypes.c_int32).__new__(metacls, name, bases, namespace)
2724

2825
def __repr__(self):
2926
return f"<Enum {self.__name__}>"
@@ -39,8 +36,12 @@ def __repr__(self):
3936
def __eq__(self, other):
4037
if isinstance(other, int):
4138
return self.value == other
39+
if type(self) is type(other):
40+
return self.value == other.value
41+
return NotImplemented
4242

43-
return type(self) is type(other) and self.value == other.value
43+
def __hash__(self):
44+
return hash(self.value)
4445

4546

4647
class DLDeviceType(_Enum):

tests/test_enum.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
from ctypes_dlpack import DLDataTypeCode, DLDeviceType
2+
3+
4+
def test_enum_eq():
5+
assert DLDeviceType.kDLCPU == 1
6+
assert 1 == DLDeviceType.kDLCPU
7+
assert DLDeviceType.kDLCUDA == 2
8+
assert DLDeviceType.kDLCPU != 2
9+
10+
assert DLDeviceType.kDLCPU == DLDeviceType.kDLCPU
11+
assert DLDeviceType.kDLCPU != DLDeviceType.kDLCUDA
12+
13+
# Both have value 1, but different enum types — must not compare equal.
14+
assert DLDeviceType.kDLCPU != DLDataTypeCode.kDLInt
15+
assert DLDataTypeCode.kDLInt != DLDeviceType.kDLCPU
16+
17+
18+
def test_enum_hashable():
19+
table = {DLDeviceType.kDLCPU: "cpu", DLDeviceType.kDLCUDA: "cuda"}
20+
assert table[DLDeviceType.kDLCPU] == "cpu"
21+
assert table[1] == "cpu"
22+
assert table[DLDeviceType.kDLCUDA] == "cuda"
23+
24+
s = {DLDeviceType.kDLCPU, DLDeviceType.kDLCUDA, DLDeviceType.kDLCPU}
25+
assert len(s) == 2
26+
27+
28+
def test_enum_repr():
29+
# ``DLDeviceType.kDLCPU`` is a plain ``int`` (class-body assignment); the
30+
# custom ``__repr__`` only applies to instances, as produced by ctypes
31+
# when reading a structure field typed as ``DLDeviceType``.
32+
assert repr(DLDeviceType(1)) == "DLDeviceType.kDLCPU"
33+
assert repr(DLDataTypeCode(2)) == "DLDataTypeCode.kDLFloat"
34+
35+
unknown = DLDeviceType(99)
36+
assert repr(unknown) == "DLDeviceType.99"
37+
38+
assert repr(DLDeviceType) == "<Enum DLDeviceType>"
39+
assert repr(DLDataTypeCode) == "<Enum DLDataTypeCode>"

0 commit comments

Comments
 (0)