diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java index b644b9b956cc..cce82541127c 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -17,7 +17,6 @@ import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; -import com.jetbrains.python.psi.impl.PyLambdaExpressionImpl; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.psi.impl.PyTypeProvider; import com.jetbrains.python.psi.resolve.PyResolveContext; @@ -281,10 +280,8 @@ public final class PyTypeChecker { substitutedRef = null; } } - final PyType substitution = Ref.deref(substitutedRef); PyType bound = expected.getBound(); List<@Nullable PyType> constraints = expected.getConstraints(); - int matchedConstraintIndex = -1; // Promote int in Type[TypeVar('T', int)] to Type[int] before checking that bounds match if (expected.isDefinition()) { bound = toClass(bound); @@ -294,6 +291,25 @@ public final class PyTypeChecker { // Remove value-specific components from the actual type to make it safe to propagate PyType safeActual = constraints.isEmpty() && bound instanceof PyLiteralStringType ? actual : replaceLiteralStringWithStr(actual); + if (substitutedRef != null) { + final PyType substitution = substitutedRef.get(); + if (expected.equals(safeActual) || expected.equals(substitution)) { + return true; + } + + MyMatchHelper matchHelper = new MyMatchHelper(expected, context); + if (matchHelper.match(substitution, safeActual)) { + return true; + } + if (context.mySubstitutions.frozenTypeVars.contains(expected)) { + return false; + } + if (!matchHelper.match(safeActual, substitution)) { + safeActual = PyUnionType.union(safeActual, substitution); + } + } + + int matchedConstraintIndex = -1; if (constraints.isEmpty()) { Optional match = match(bound, safeActual, context); if (match.isPresent() && !match.get()) { @@ -301,25 +317,13 @@ public final class PyTypeChecker { } } else { - matchedConstraintIndex = ContainerUtil.indexOf(constraints, constraint -> match(constraint, safeActual, context).orElse(true)); + final PyType finalSafeActual = safeActual; + matchedConstraintIndex = ContainerUtil.indexOf(constraints, constraint -> match(constraint, finalSafeActual, context).orElse(true)); if (matchedConstraintIndex == -1) { return false; } } - if (substitutedRef != null) { - if (expected.equals(safeActual) || expected.equals(substitution)) { - return true; - } - - Optional recursiveMatch = RecursionManager.doPreventingRecursion( - expected, false, context.reversedSubstitutions - ? () -> match(safeActual, substitution, context) - : () -> match(substitution, safeActual, context) - ); - return recursiveMatch != null ? recursiveMatch.orElse(false) : false; - } - if (safeActual != null) { PyType type = constraints.isEmpty() ? safeActual : constraints.get(matchedConstraintIndex); context.mySubstitutions.typeVars.put(expected, Ref.create(type)); @@ -334,6 +338,25 @@ public final class PyTypeChecker { return true; } + private static final class MyMatchHelper { + private final @NotNull Object myKey; + private final @NotNull MatchContext myContext; + + MyMatchHelper(@NotNull Object key, @NotNull MatchContext context) { + myKey = key; + myContext = context; + } + + boolean match(@Nullable PyType expected, @Nullable PyType actual) { + Optional result = RecursionManager.doPreventingRecursion( + myKey, false, myContext.reversedSubstitutions + ? () -> PyTypeChecker.match(actual, expected, myContext) + : () -> PyTypeChecker.match(expected, actual, myContext) + ); + return result != null && result.orElse(false); + } + } + private static @Nullable PyType toClass(@Nullable PyType type) { return PyTypeUtil.toStream(type) .map(t -> t instanceof PyInstantiableType instantiableType ? instantiableType.toClass() : t) @@ -1587,6 +1610,7 @@ public final class PyTypeChecker { } }); } + substitutions.frozenTypeVars = new HashSet<>(substitutions.typeVars.keySet()); return substitutions; } @@ -1896,8 +1920,6 @@ public final class PyTypeChecker { @ApiStatus.Experimental public static class GenericSubstitutions { - - // Nullable-Nullable because of com.jetbrains.python.psi.types.PyTypeChecker.collectTypeSubstitutions private final @NotNull Map> typeVars; private final @NotNull Map typeVarTuples; @@ -1906,6 +1928,8 @@ public final class PyTypeChecker { private @Nullable PyType qualifierType; + private @NotNull Set frozenTypeVars = Collections.emptySet(); + public GenericSubstitutions(@NotNull Map typeParameters) { this( EntryStream.of(typeParameters) diff --git a/python/testData/inspections/PyTypeCheckerInspection/GenericUserFunctions.py b/python/testData/inspections/PyTypeCheckerInspection/GenericUserFunctions.py index 12fa804905c6..b34a24519a89 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/GenericUserFunctions.py +++ b/python/testData/inspections/PyTypeCheckerInspection/GenericUserFunctions.py @@ -40,7 +40,11 @@ def test(): print(result) print(result + 'foo') - f2(1, ['foo'], 'bar') + # Bug: Expected error. + # Generics are considered to be covariant. + # I.e. `list[str]` is assignable to `list[int | str]`. + # Thus, substitution `T` -> `int | str` is considered valid. + f2(1, ['foo'], 'bar') result = f3(1, 'foo', True) f4(result) diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index 4499932bf7fc..43322f7cfa89 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -3359,4 +3359,33 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { var: Template = Concrete() """); } + + // PY-25989 PY-84544 + public void testTypeVarWidening() { + myFixture.enableInspections(PyAssertTypeInspection.class); + doTestByText(""" + from collections.abc import Iterable + from typing import assert_type + + + # PY-84544 + def foo(iterable: Iterable[int] | Iterable[str]) -> None: + assert_type(next(iter(iterable)), int | str) + + + # PY-25989 + assert_type(max(1, 2.6), float) + assert_type(max(2.6, 1), float) + max(1, object()) + + + def bar[T: int, str](v1: T, v2: T) -> T: + if (bool(input())): + return v1 + return v2 + + + _ = bar(1, "a") + """); + } } diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java index 95256beee45e..82313f268d30 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java @@ -1347,7 +1347,10 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { def accepts_anything(x: str) -> None: pass - func(42, accepts_anything)""") + # Bug: Expected error. + # `Callable[[str], None]` is assignable to `Callable[[int | str], None]`. + # Thus, substitution `T` -> `int | str` is considered valid. + func(42, accepts_anything)""") ); }