Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions tests/test_traitlets_enum.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,18 @@ class CSColor(enum.Enum):
YeLLoW = 4


class StrColor(str, enum.Enum):
RED = "red"
GREEN = "green"
BLUE = "blue"


class ShadowedColor(str, enum.Enum):
# A member name that collides with a different member's value
red = "green"
green = "blue"


color_choices = ["red", "Green", "BLUE", "YeLLoW"]


Expand Down Expand Up @@ -129,6 +141,33 @@ def test_assign_bad_enum_value_number__raises_error(self):
with self.assertRaises(TraitError):
example.color = value

def test_assign_str_enum_value(self):
# -- CONVERT: value (as string) => Enum value (item)
class Example(HasTraits):
color = UseEnum(StrColor)

for enum_item in StrColor.__members__.values():
example = Example()
example.color = enum_item.value
self.assertIs(example.color, enum_item)

def test_assign_str_enum_name_wins_over_value(self):
# -- A name lookup must keep taking precedence over a value lookup
class Example(HasTraits):
color = UseEnum(ShadowedColor)

example = Example()
example.color = "green"
self.assertIs(example.color, ShadowedColor.green)

def test_assign_bad_str_enum_value__raises_error(self):
class Example(HasTraits):
color = UseEnum(StrColor)

example = Example()
with self.assertRaises(TraitError):
example.color = "purple"

def test_ctor_without_default_value(self):
# -- IMPLICIT: default_value = Color.red (first enum-value)
class Example2(HasTraits):
Expand Down
14 changes: 14 additions & 0 deletions traitlets/traitlets.py
Original file line number Diff line number Diff line change
Expand Up @@ -4275,6 +4275,15 @@ def select_by_name(self, value: str, default: t.Any = Undefined) -> t.Any:
value = value.replace(self.name_prefix, "", 1)
return self.enum_class.__members__.get(value, default)

def select_by_value(self, value: t.Any, default: t.Any = Undefined) -> t.Any:
"""Selects enum-value by using its value-constant."""
enum_members = self.enum_class.__members__
for enum_item in enum_members.values():
if enum_item.value == value:
return enum_item
# -- NOT FOUND:
return default

def validate(self, obj: t.Any, value: t.Any) -> t.Any:
if isinstance(value, self.enum_class):
return value
Expand All @@ -4288,6 +4297,11 @@ def validate(self, obj: t.Any, value: t.Any) -> t.Any:
value2 = self.select_by_name(value)
if value2 is not Undefined:
return value2
# -- CONVERT: value (as string) => enum_value (item), for str-based enums
# whose member names differ from their values
value2 = self.select_by_value(value)
if value2 is not Undefined:
return value2
elif value is None:
if self.allow_none:
return None
Expand Down
Loading