diff --git a/tests/test_typing.py b/tests/test_typing.py index 592955d7..a5c7c7b1 100644 --- a/tests/test_typing.py +++ b/tests/test_typing.py @@ -249,7 +249,7 @@ class T(HasTraits): @pytest.mark.mypy_testing def mypy_set_typing() -> None: class T(HasTraits): - remove_cell_tags = Set( + remove_cell_tags: Set[t.Any] = Set( Unicode(), default_value=[], help=( @@ -258,7 +258,7 @@ class T(HasTraits): ), ).tag(config=True) - safe_output_keys = Set( + safe_output_keys: Set[str] = Set( config=True, default_value={ "metadata", # Not a mimetype per-se, but expected and safe. @@ -272,14 +272,14 @@ class T(HasTraits): ) t = T() - reveal_type(Set("foo")) # R: traitlets.traitlets.Set - reveal_type(Set("").tag(sync=True)) # R: traitlets.traitlets.Set - reveal_type(Set(None, allow_none=True)) # R: traitlets.traitlets.Set - reveal_type(Set(None, allow_none=True).tag(sync=True)) # R: traitlets.traitlets.Set - reveal_type(T.remove_cell_tags) # R: traitlets.traitlets.Set + reveal_type(Set("foo")) # R: traitlets.traitlets.Set[Never] + reveal_type(Set("").tag(sync=True)) # R: traitlets.traitlets.Set[Never] + reveal_type(Set(None, allow_none=True)) # R: traitlets.traitlets.Set[Never] + reveal_type(Set(None, allow_none=True).tag(sync=True)) # R: traitlets.traitlets.Set[Never] + reveal_type(T.remove_cell_tags) # R: traitlets.traitlets.Set[Any] reveal_type(t.remove_cell_tags) # R: set[Any] - reveal_type(T.safe_output_keys) # R: traitlets.traitlets.Set - reveal_type(t.safe_output_keys) # R: set[Any] + reveal_type(T.safe_output_keys) # R: traitlets.traitlets.Set[str] + reveal_type(t.safe_output_keys) # R: set[str] @pytest.mark.mypy_testing @@ -452,3 +452,13 @@ class T(HasTraits): t.inst = "foo" # E: Incompatible types in assignment (expression has type "str", variable has type "Foo") [assignment] t.oinst = "foo" # E: Incompatible types in assignment (expression has type "str", variable has type "Foo | None") [assignment] t.inst = None # E: Incompatible types in assignment (expression has type "None", variable has type "Foo") [assignment] + + +@pytest.mark.mypy_testing +def mypy_generic_set_typing() -> None: + class T(HasTraits): + values: Set[str] = Set() + + t = T() + reveal_type(T.values) # R: traitlets.traitlets.Set[str] + reveal_type(t.values) # R: set[str] diff --git a/traitlets/traitlets.py b/traitlets/traitlets.py index 2cc239c1..2bfbd922 100644 --- a/traitlets/traitlets.py +++ b/traitlets/traitlets.py @@ -3676,10 +3676,10 @@ def set(self, obj: t.Any, value: t.Any) -> None: return super().set(obj, value) -class Set(Container[set[t.Any]]): +class Set(Container[set[T]]): """An instance of a Python set.""" - klass = set + klass = set # type: ignore[assignment] _cast_types = (tuple, list) _literal_from_string_pairs = ("[]", "()", "{}") @@ -3739,13 +3739,14 @@ def validate_elements(self, obj: t.Any, value: t.Any) -> t.Any: def set(self, obj: t.Any, value: t.Any) -> None: if isinstance(value, str): - return super().set(obj, {value}) + return super().set(obj, {t.cast(T, value)}) else: return super().set(obj, value) def default_value_repr(self) -> str: # Ensure default value is sorted for a reproducible build - list_repr = repr(sorted(self.make_dynamic_default() or [])) + default = t.cast(set[t.Any] | None, self.make_dynamic_default()) or set() + list_repr = repr(sorted(default)) if list_repr == "[]": return "set()" return "{" + list_repr[1:-1] + "}"