[python] PY-79178/PY-80252 report invalid and unrequired casts

Merge-request: IJ-MR-171620
Merged-by: Morgan Bartholomew <morgan.bartholomew@jetbrains.com>

GitOrigin-RevId: d9b13a99bad56a112f04febfc939823397b55ddd
This commit is contained in:
Morgan Bartholomew
2025-08-26 01:51:36 +00:00
committed by intellij-monorepo-bot
parent ceb58f48e6
commit 51615780f9
9 changed files with 462 additions and 2 deletions
@@ -0,0 +1,29 @@
<html>
<body>
<p>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</p>
<p>This usually indicates a mistake. If the conversion is intentional, first convert the expression to the common parent type to make the intent explicit.</p>
<p><b>Example:</b></p>
<pre><code>
from typing import cast
# Non-overlapping types — likely a mistake
<b>cast(int, "a")</b> # 'str' -> 'int'
<b>cast(list[int], ["a"])</b> # 'list[str]' -> 'list[int]'
# Recommended explicit escape hatch is to use a "double cast"
cast(int, <b>cast(object, "a")</b>) # 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 <code>list</code> is invariant. It's not currently supported by this inspection
int_list = [1, 2, 3]
cast(list[object], int_list)
</code></pre>
<!-- tooltip end -->
<p>The inspection relies on static type information; when a type is unknown, no warning is reported.
Variance of generic types is not yet considered.</p>
</body>
</html>
@@ -0,0 +1,12 @@
<html>
<body>
<p>Reports unnecessary calls to typing.cast when the expression already has the specified target type.</p>
<p><b>Example:</b></p>
<pre><code>
from typing import cast
a: int
b = <b>cast(int, a)</b> # Unnecessary, a is already int
</code></pre>
</body>
</html>
@@ -238,6 +238,8 @@
<localInspection language="Python" shortName="PyMandatoryEncodingInspection" suppressId="PyMandatoryEncoding" bundle="messages.PyPsiBundle" key="INSP.NAME.mandatory.encoding" groupKey="INSP.GROUP.python" enabledByDefault="false" level="WARNING" implementationClass="com.jetbrains.python.inspections.PyMandatoryEncodingInspection"/>
<localInspection language="Python" shortName="PyClassHasNoInitInspection" suppressId="PyClassHasNoInit" bundle="messages.PyPsiBundle" key="INSP.NAME.class.has.no.init" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WEAK WARNING" implementationClass="com.jetbrains.python.inspections.PyClassHasNoInitInspection"/>
<localInspection language="Python" shortName="PyNoneFunctionAssignmentInspection" suppressId="PyNoneFunctionAssignment" bundle="messages.PyPsiBundle" key="INSP.NAME.none.function.assignment" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WEAK WARNING" implementationClass="com.jetbrains.python.inspections.PyNoneFunctionAssignmentInspection"/>
<localInspection language="Python" shortName="PyInvalidCastInspection" suppressId="PyInvalidCast" bundle="messages.PyPsiBundle" key="INSP.NAME.invalid.cast" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WARNING" implementationClass="com.jetbrains.python.inspections.PyInvalidCastInspection"/>
<localInspection language="Python" shortName="PyUnnecessaryCastInspection" suppressId="PyUnnecessaryCast" bundle="messages.PyPsiBundle" key="INSP.NAME.unnecessary.cast" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WARNING" implementationClass="com.jetbrains.python.inspections.PyUnnecessaryCastInspection"/>
<localInspection language="Python" shortName="PyProtectedMemberInspection" suppressId="PyProtectedMember" bundle="messages.PyPsiBundle" key="INSP.NAME.protected.member" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WEAK WARNING" implementationClass="com.jetbrains.python.inspections.PyProtectedMemberInspection"/>
<localInspection language="Python" shortName="PyMethodMayBeStaticInspection" suppressId="PyMethodMayBeStatic" bundle="messages.PyPsiBundle" key="INSP.NAME.method.may.be.static" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WEAK WARNING" implementationClass="com.jetbrains.python.inspections.PyMethodMayBeStaticInspection"/>
<localInspection language="Python" shortName="PyDocstringTypesInspection" suppressId="PyDocstringTypes" bundle="messages.PyPsiBundle" key="INSP.NAME.docstring.types" groupKey="INSP.GROUP.python" enabledByDefault="true" level="WEAK WARNING" implementationClass="com.jetbrains.python.inspections.PyDocstringTypesInspection"/>
@@ -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
@@ -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<PyType>? = 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<PyClassLikeType> {
val result = ArrayList<PyClassLikeType>()
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
}
@@ -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<PyType>? = 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())
}
}
@@ -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<PyType> actualFirstParamTypes = ContainerUtil.map(actualParameters.getParameters().subList(0, expectedPrefixSize),
List<PyType> 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);
@@ -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
<warning descr="Cast of type 'str' to type 'int' may be a mistake because no possible value of one is assignable with the other. If this was intentional, cast the expression to 'object' first.">cast(int, "a")</warning>
<warning descr="Cast of type 'list[str]' to type 'list[int]' may be a mistake because no possible value of one is assignable with the other. If this was intentional, cast the expression to 'object' first.">cast(list[int], ["a"])</warning>
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
<warning descr="Cast of type 'B1' to type 'B2' may be a mistake because no possible value of one is assignable with the other. If this was intentional, cast the expression to 'A' first.">cast(B2, B1())</warning>
""".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
<warning descr="Cast of type 'B1' to type 'B2' may be a mistake because no possible value of one is assignable with the other. If this was intentional, cast the expression to 'A' first."><caret>cast(B2, B1())</warning>
""".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
<warning descr="Cast of type 'A1 | None' to type 'int | str' may be a mistake because no possible value of one is assignable with the other. If this was intentional, cast the expression to 'object' first.">cast(int | str, a1)</warning>
""".trimIndent())
}
}
@@ -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<out PyInspection> = PyUnnecessaryCastInspection::class.java
fun `test basic`() {
doTestByText(
"""
from typing import cast
def f(a: int):
<warning descr="Unnecessary cast; type is already 'int'">cast(int,</warning>a)
""".trimIndent()
)
}
fun `test literal`() {
doTestByText(
"""
from typing import cast, Literal
one: Literal[1] = 1
cast(int, one)
<warning descr="Unnecessary cast; type is already 'Literal[1]'">cast(Literal[1],</warning> 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):
<warning descr="Unnecessary cast; type is already 'int'"><caret>cast(int,</warning> 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()
)
}
}