diff --git a/python/python-psi-impl/resources/inspectionDescriptions/PyInvalidCastInspection.html b/python/python-psi-impl/resources/inspectionDescriptions/PyInvalidCastInspection.html new file mode 100644 index 000000000000..7c70e1a8d260 --- /dev/null +++ b/python/python-psi-impl/resources/inspectionDescriptions/PyInvalidCastInspection.html @@ -0,0 +1,29 @@ + + +

Reports calls to `typing.cast` where no possible value of the source type can be assignable to the target type. We can refer to this as "non overlapping" types

+

This usually indicates a mistake. If the conversion is intentional, first convert the expression to the common parent type to make the intent explicit.

+

Example:

+

+from typing import cast
+
+# Non-overlapping types — likely a mistake
+cast(int, "a")          # 'str' -> 'int'
+cast(list[int], ["a"])  # 'list[str]' -> 'list[int]'
+
+# Recommended explicit escape hatch is to use a "double cast"
+cast(int, cast(object, "a"))  # ok
+
+# Legitimate overlapping cases
+cast(int, object())    # a valid down cast
+cast(object, 1)        # a valid up cast
+
+# While the following is an invalid cast, as list is invariant. It's not currently supported by this inspection
+int_list = [1, 2, 3]
+cast(list[object], int_list)
+
+ +

The inspection relies on static type information; when a type is unknown, no warning is reported. + + Variance of generic types is not yet considered.

+ + \ No newline at end of file diff --git a/python/python-psi-impl/resources/inspectionDescriptions/PyUnnecessaryCastInspection.html b/python/python-psi-impl/resources/inspectionDescriptions/PyUnnecessaryCastInspection.html new file mode 100644 index 000000000000..2a6850000978 --- /dev/null +++ b/python/python-psi-impl/resources/inspectionDescriptions/PyUnnecessaryCastInspection.html @@ -0,0 +1,12 @@ + + +

Reports unnecessary calls to typing.cast when the expression already has the specified target type.

+

Example:

+

+from typing import cast
+
+a: int
+b = cast(int, a)  # Unnecessary, a is already int
+
+ + diff --git a/python/python-psi-impl/resources/intellij.python.psi.impl.xml b/python/python-psi-impl/resources/intellij.python.psi.impl.xml index 6c2cfaf66721..f5f7ac99cbdb 100644 --- a/python/python-psi-impl/resources/intellij.python.psi.impl.xml +++ b/python/python-psi-impl/resources/intellij.python.psi.impl.xml @@ -238,6 +238,8 @@ + + diff --git a/python/python-psi-impl/resources/messages/PyPsiBundle.properties b/python/python-psi-impl/resources/messages/PyPsiBundle.properties index 614e53d48ca5..ffbf0b1affad 100644 --- a/python/python-psi-impl/resources/messages/PyPsiBundle.properties +++ b/python/python-psi-impl/resources/messages/PyPsiBundle.properties @@ -1332,3 +1332,17 @@ INSP.NAME.new.type.new.type.cannot.be.used.with=NewType cannot be used with ''{0 INSP.NAME.new.type.new.type.cannot.be.generic=NewType cannot be generic packaging.could.not.parse.relation=Could not parse relation from: {0} + +# PyInvalidCastInspection +INSP.NAME.invalid.cast=Type cast with impossible types +INSP.invalid.cast.message=Cast of type ''{0}'' to type ''{1}'' may be a mistake because no possible value of one is assignable with the other. If this was intentional, cast the expression to ''{2}'' first. + +# Quick fixes for PyInvalidCastInspection +QFIX.add.intermediate.cast=Add cast({0}, ...) + +# PyUnnecessaryCastInspection +INSP.NAME.unnecessary.cast=Unnecessary type cast +INSP.unnecessary.cast.message=Unnecessary cast; type is already ''{0}'' + +# Quick fixes for PyUnnecessaryCastInspection +QFIX.remove.cast.call=Remove 'cast' call diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyInvalidCastInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyInvalidCastInspection.kt new file mode 100644 index 000000000000..87aa350f661e --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyInvalidCastInspection.kt @@ -0,0 +1,106 @@ +package com.jetbrains.python.inspections + +import com.intellij.codeInspection.LocalInspectionToolSession +import com.intellij.codeInspection.ProblemsHolder +import com.intellij.modcommand.ModPsiUpdater +import com.intellij.modcommand.PsiUpdateModCommandQuickFix +import com.intellij.openapi.project.Project +import com.intellij.openapi.util.Ref +import com.intellij.psi.PsiElement +import com.intellij.psi.PsiElementVisitor +import com.jetbrains.python.PyPsiBundle +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.documentation.PythonDocumentationProvider +import com.jetbrains.python.psi.* +import com.jetbrains.python.psi.types.* + +class PyInvalidCastInspection : PyInspection() { + override fun buildVisitor(holder: ProblemsHolder, isOnTheFly: Boolean, session: LocalInspectionToolSession): PsiElementVisitor { + return object : PyInspectionVisitor(holder, getContext(session)) { + override fun visitPyCallExpression(callExpression: PyCallExpression) { + val callees = callExpression.multiResolveCalleeFunction(resolveContext) + val isCastCall = callees.any { (it as? PyFunction)?.qualifiedName == PyTypingTypeProvider.CAST || + (it as? PyFunction)?.qualifiedName == PyTypingTypeProvider.CAST_EXT } + if (!isCastCall) return + + val args = callExpression.getArguments() + if (args.size != 2) return + val targetTypeRef: Ref? = PyTypingTypeProvider.getType(args[0], myTypeEvalContext) + val targetType = Ref.deref(targetTypeRef) + val actualType = myTypeEvalContext.getType(args[1]) + + if (PyTypeChecker.overlappingTypes(targetType, actualType, myTypeEvalContext)) return + val fromName = PythonDocumentationProvider.getTypeName(actualType, myTypeEvalContext) + val toName = PythonDocumentationProvider.getVerboseTypeName(targetType, myTypeEvalContext) + + val suggestedName = computeSuggestedIntermediateTypeName(targetType, actualType, myTypeEvalContext) + + registerProblem( + callExpression, + PyPsiBundle.message( + "INSP.invalid.cast.message", + fromName, + toName, + suggestedName + ), + AddIntermediateCastQuickFix(suggestedName) + ) + } + } + } +} + +private class AddIntermediateCastQuickFix(private val typeText: String) : PsiUpdateModCommandQuickFix() { + override fun getFamilyName(): String = PyPsiBundle.message("QFIX.add.intermediate.cast", typeText) + + override fun applyFix(project: Project, element: PsiElement, updater: ModPsiUpdater) { + val call = element as? PyCallExpression ?: return + val args = call.arguments + if (args.size != 2) return + val expr = args[1] ?: return + + val calleeText = call.callee?.text ?: "cast" + val langLevel = LanguageLevel.forElement(call) + val generator = PyElementGenerator.getInstance(project) + val castExprText = "$calleeText($typeText, ${expr.text})" + val newExpr = generator.createExpressionFromText(langLevel, castExprText) + expr.replace(newExpr) + } +} + +private fun computeSuggestedIntermediateTypeName(targetType: PyType?, actualType: PyType?, context: TypeEvalContext): String { + val objectName = "object" + + fun toNonCollectionClassLike(t: PyType?): PyClassLikeType? = when (t) { + is PyCollectionType -> null // avoid suggesting collection classes like 'list' as an intermediate type + is PyClassLikeType -> t + else -> null + } + + val left = toNonCollectionClassLike(actualType) + val right = toNonCollectionClassLike(targetType) + if (left != null && right != null) { + fun mro(t: PyClassLikeType): List { + val result = ArrayList() + result.add(t) + result.addAll(t.getAncestorTypes(context)) + return result + } + + val leftMro = mro(left) + val rightQNames = mro(right).mapNotNull { it.classQName } + val rightSet = rightQNames.toSet() + + for (t in leftMro) { + val qn = t.classQName + if (qn != null && rightSet.contains(qn)) { + val name = t.name + if (name != null && name != objectName) { + return name + } + } + } + } + + return objectName +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyUnnecessaryCastInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyUnnecessaryCastInspection.kt new file mode 100644 index 000000000000..2d39968ef427 --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyUnnecessaryCastInspection.kt @@ -0,0 +1,70 @@ +package com.jetbrains.python.inspections + +import com.intellij.codeInspection.LocalInspectionToolSession +import com.intellij.codeInspection.ProblemHighlightType +import com.intellij.codeInspection.ProblemsHolder +import com.intellij.codeInspection.ex.ProblemDescriptorImpl +import com.intellij.modcommand.ModPsiUpdater +import com.intellij.modcommand.PsiUpdateModCommandQuickFix +import com.intellij.openapi.editor.colors.TextAttributesKey +import com.intellij.openapi.project.Project +import com.intellij.openapi.util.Ref +import com.intellij.openapi.util.TextRange +import com.intellij.psi.PsiElement +import com.intellij.psi.PsiElementVisitor +import com.intellij.psi.util.endOffset +import com.intellij.psi.util.startOffset +import com.jetbrains.python.PyPsiBundle +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.documentation.PythonDocumentationProvider +import com.jetbrains.python.psi.PyCallExpression +import com.jetbrains.python.psi.PyFunction +import com.jetbrains.python.psi.types.PyType +import com.jetbrains.python.psi.types.PyTypeChecker + +class PyUnnecessaryCastInspection : PyInspection() { + override fun buildVisitor(holder: ProblemsHolder, isOnTheFly: Boolean, session: LocalInspectionToolSession): PsiElementVisitor { + return object : PyInspectionVisitor(holder, getContext(session)) { + override fun visitPyCallExpression(callExpression: PyCallExpression) { + val callees = callExpression.multiResolveCalleeFunction(resolveContext) + val isCastCall = callees.any { + (it as? PyFunction)?.qualifiedName == PyTypingTypeProvider.CAST || + (it as? PyFunction)?.qualifiedName == PyTypingTypeProvider.CAST_EXT + } + if (!isCastCall) return + + val args = callExpression.getArguments() + if (args.size != 2) return + val targetTypeRef: Ref? = PyTypingTypeProvider.getType(args[0], myTypeEvalContext) + val targetType = Ref.deref(targetTypeRef) + val actualType: PyType? = myTypeEvalContext.getType(args[1]) + + if (!PyTypeChecker.sameType(targetType, actualType, myTypeEvalContext)) return + val toName = PythonDocumentationProvider.getTypeName(targetType, myTypeEvalContext) + registerProblem( + callExpression, + PyPsiBundle.message( + "INSP.unnecessary.cast.message", + toName + ), + ProblemHighlightType.LIKE_UNUSED_SYMBOL, + null, + TextRange(0, callExpression.arguments[0].nextSibling.endOffset - callExpression.startOffset), + RemoveUnnecessaryCastQuickFix(), + ) + } + } + } +} + +private class RemoveUnnecessaryCastQuickFix : PsiUpdateModCommandQuickFix() { + override fun getFamilyName(): String = PyPsiBundle.message("QFIX.remove.cast.call") + + override fun applyFix(project: Project, element: PsiElement, updater: ModPsiUpdater) { + val call = element as? PyCallExpression ?: return + val args = call.getArguments() + if (args.size != 2) return + val expr = args[1] ?: return + call.replace(expr.copy()) + } +} 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 d9635ec902a9..66813af1284f 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 @@ -156,7 +156,7 @@ public final class PyTypeChecker { if (expected instanceof PyConcatenateType concatenateType) { return Optional.of(match(concatenateType, actual, context)); } - + if (expected == null || actual == null || isUnknown(actual, context.context)) { return Optional.of(true); } @@ -445,7 +445,7 @@ public final class PyTypeChecker { if (expectedPrefixSize > actualParameters.getParameters().size()) { return false; } - List actualFirstParamTypes = ContainerUtil.map(actualParameters.getParameters().subList(0, expectedPrefixSize), + List actualFirstParamTypes = ContainerUtil.map(actualParameters.getParameters().subList(0, expectedPrefixSize), it -> it.getType(context.context)); if (!match(expectedFirstTypes, actualFirstParamTypes, context)) { return false; @@ -571,6 +571,27 @@ public final class PyTypeChecker { return Optional.empty(); } + public static boolean sameType(@Nullable PyType type1, @Nullable PyType type2, @NotNull TypeEvalContext context) { + if ((type1 == null || type2 == null) && type1 != type2) return false; + + return match(type1, type2, context) + && match(type2, type1, context); + } + + /** + * if some possible value of one type is assignable to the other type + */ + public static boolean overlappingTypes(@Nullable PyType type1, @Nullable PyType type2, @NotNull TypeEvalContext context) { + if (type1 instanceof PyUnionType unionType1) { + return ContainerUtil.exists(unionType1.getMembers(), t -> overlappingTypes(t, type2, context)); + } + if (type2 instanceof PyUnionType unionType2) { + return ContainerUtil.exists(unionType2.getMembers(), t -> overlappingTypes(type1, t, context)); + } + return match(type1, type2, context) + || match(type2, type1, context); + } + private static boolean matchProtocols(@NotNull PyClassType expected, @NotNull PyClassType actual, @NotNull MatchContext matchContext) { GenericSubstitutions substitutions = collectTypeSubstitutions(actual, matchContext.context); diff --git a/python/testSrc/com/jetbrains/python/inspections/PyInvalidCastInspectionTest.kt b/python/testSrc/com/jetbrains/python/inspections/PyInvalidCastInspectionTest.kt new file mode 100644 index 000000000000..6cc426a46b1e --- /dev/null +++ b/python/testSrc/com/jetbrains/python/inspections/PyInvalidCastInspectionTest.kt @@ -0,0 +1,125 @@ +// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.jetbrains.python.inspections + +import com.jetbrains.python.PyPsiBundle +import com.jetbrains.python.PythonFileType +import com.jetbrains.python.fixtures.PyInspectionTestCase + +class PyInvalidCastInspectionTest : PyInspectionTestCase() { + override fun getInspectionClass() = PyInvalidCastInspection::class.java + + fun `test basic`() { + doTestByText( + """ + from typing import cast + + cast(int, "a") + cast(list[int], ["a"]) + + cast(int, object()) # ok + cast(object, 1) # ok + + lint = [1, 2, 3] + cast(list[object], lint) # ok + """.trimIndent() + ) + } + + fun `test Any`() { + doTestByText( + """ + from typing import cast + + cast(Any, 1) # ok + any: Any = 1 + cast(int, any) # ok + """.trimIndent() + ) + } + + /** + * test that the common super type is shown in the error message + */ + fun `test common super type`() { + doTestByText( + """ + from typing import cast + + class A: pass + + class B1(A): pass + class B2(A): pass + + cast(B2, B1()) + """.trimIndent() + ) + } + + /** + * test that normally castable generics will report an error if they are invariant + */ + fun `test generic variance`() { + doTestByText( + """ + from typing import cast, Sequence + + lint = [1, 2, 3] + # should actually fail because a `list[int]` can never be a `list[object]` + cast(list[object], lint) + + cast(Sequence[object], lint) # ok + """.trimIndent() + ) + } + + fun `test quickfix add intermediate cast`() { + val text = """ + from typing import cast + + class A: pass + + class B1(A): pass + class B2(A): pass + + cast(B2, B1()) + """.trimIndent() + myFixture.configureByText(PythonFileType.INSTANCE, text) + configureInspection() + val hint = PyPsiBundle.message("QFIX.add.intermediate.cast", "A") + val action = myFixture.findSingleIntention(hint) + myFixture.launchAction(action) + myFixture.checkResult( + """ + from typing import cast + + class A: pass + + class B1(A): pass + class B2(A): pass + + cast(B2, cast(A, B1())) + """.trimIndent() + ) + } + + fun `test overlapping unions`() { + doTestByText(""" + from typing import cast, Literal + + type AB = Literal["a", "b"] + type BC = Literal["b", "c"] + + def foo(x: AB): + cast(BC, x) # ok + + class A1: ... + class A2(A1): ... + + def bar(a1: A1 | None, x2: a2 | None): + cast(A2 | None, a1) # ok + cast(A1 | None, a2) # ok + + cast(int | str, a1) + """.trimIndent()) + } +} diff --git a/python/testSrc/com/jetbrains/python/inspections/PyUnnecessaryCastInspectionTest.kt b/python/testSrc/com/jetbrains/python/inspections/PyUnnecessaryCastInspectionTest.kt new file mode 100644 index 000000000000..56194131441e --- /dev/null +++ b/python/testSrc/com/jetbrains/python/inspections/PyUnnecessaryCastInspectionTest.kt @@ -0,0 +1,81 @@ +// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.jetbrains.python.inspections + +import com.jetbrains.python.PyPsiBundle +import com.jetbrains.python.PythonFileType +import com.jetbrains.python.fixtures.PyInspectionTestCase + +class PyUnnecessaryCastInspectionTest : PyInspectionTestCase() { + override fun getInspectionClass(): Class = PyUnnecessaryCastInspection::class.java + + fun `test basic`() { + doTestByText( + """ + from typing import cast + + def f(a: int): + cast(int,a) + """.trimIndent() + ) + } + + fun `test literal`() { + doTestByText( + """ + from typing import cast, Literal + + one: Literal[1] = 1 + cast(int, one) + cast(Literal[1], one) + """.trimIndent() + ) + } + + fun `test union`() { + doTestByText( + """ +from typing import cast + + + """.trimIndent() + ) + } + + fun `test okay`(){ + doTestByText( + """ + from typing import cast + + class B: ... + class C(B): ... + + cast(B, C()) # ok + + a: int | str + b = cast(str, a) # ok + """.trimIndent() + ) + } + + fun `test quickfix remove`() { + val text = """ + from typing import cast + + def f(a: int): + cast(int, a) + """.trimIndent() + myFixture.configureByText(PythonFileType.INSTANCE, text) + configureInspection() + val hint = PyPsiBundle.message("QFIX.remove.cast.call") + val action = myFixture.findSingleIntention(hint) + myFixture.launchAction(action) + myFixture.checkResult( + """ + from typing import cast + + def f(a: int): + a + """.trimIndent() + ) + } +}