diff --git a/mypy/subtypes.py b/mypy/subtypes.py index d9c1b47c317f..cbab19f6247f 100644 --- a/mypy/subtypes.py +++ b/mypy/subtypes.py @@ -1390,7 +1390,9 @@ def f(self) -> A: ... subtype_context=SubtypeContext(ignore_pos_arg_names=ignore_names), proper_subtype=proper_subtype, ) - type_state.record_subtype_cache_entry(subtype_kind, left, right) + # A partial protocol check does not prove full subtype compatibility. + if not skip: + type_state.record_subtype_cache_entry(subtype_kind, left, right) return True diff --git a/test-data/unit/check-protocols.test b/test-data/unit/check-protocols.test index 75a0dfe2f5f4..9be565136723 100644 --- a/test-data/unit/check-protocols.test +++ b/test-data/unit/check-protocols.test @@ -2730,6 +2730,46 @@ def test(func: A[T, S]) -> Tuple[T, S]: ... reveal_type(test(f)) # N: Revealed type is "tuple[builtins.str, builtins.int]" [builtins fixtures/tuple.pyi] +[case testUnionOfCallbackProtocolsInference] +from typing import Protocol, TypeVar + +T = TypeVar("T") +PT = TypeVar("PT", contravariant=True) + +class OneArg(Protocol[PT]): + def __call__(self, x: PT, /) -> None: ... + +class TwoArgs(Protocol[PT]): + def __call__(self, x: PT, y: int, /) -> None: ... + +def first_arg(fn: OneArg[T] | TwoArgs[T]) -> T: ... + +def one(x: int) -> None: ... +def two(x: str, y: int) -> None: ... + +reveal_type(first_arg(one)) # N: Revealed type is "builtins.int" +reveal_type(first_arg(two)) # N: Revealed type is "builtins.str" + +[case testUnionOfCallbackProtocolsInferenceReversed] +from typing import Protocol, TypeVar + +T = TypeVar("T") +PT = TypeVar("PT", contravariant=True) + +class OneArg(Protocol[PT]): + def __call__(self, x: PT, /) -> None: ... + +class TwoArgs(Protocol[PT]): + def __call__(self, x: PT, y: int, /) -> None: ... + +def first_arg(fn: OneArg[T] | TwoArgs[T]) -> T: ... + +def one(x: int) -> None: ... +def two(x: str, y: int) -> None: ... + +reveal_type(first_arg(two)) # N: Revealed type is "builtins.str" +reveal_type(first_arg(one)) # N: Revealed type is "builtins.int" + [case testProtocolsAlwaysABCs] from typing import Protocol