diff --git a/tests/test_traitlets_enum.py b/tests/test_traitlets_enum.py index b72b1fd1..fc78097c 100644 --- a/tests/test_traitlets_enum.py +++ b/tests/test_traitlets_enum.py @@ -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"] @@ -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): diff --git a/traitlets/traitlets.py b/traitlets/traitlets.py index 0989ea98..ab04e303 100644 --- a/traitlets/traitlets.py +++ b/traitlets/traitlets.py @@ -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 @@ -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