diff --git a/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaClassType.java b/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaClassType.java index f1bb6f86666d..bf1598d1c56f 100644 --- a/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaClassType.java +++ b/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaClassType.java @@ -1,4 +1,4 @@ -// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.intellij.python.community.plugin.java.psi.impl; import com.intellij.openapi.util.NlsSafe; @@ -10,7 +10,7 @@ import com.intellij.psi.ResolveState; import com.intellij.util.ProcessingContext; import com.intellij.util.Processor; import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.impl.ResolveResultList; import com.jetbrains.python.psi.resolve.CompletionVariantsProcessor; @@ -108,7 +108,7 @@ public class PyJavaClassType implements PyClassLikeType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { return getReturnType(context); } diff --git a/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaMethodType.java b/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaMethodType.java index 25bdb47d34d0..497f700224e3 100644 --- a/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaMethodType.java +++ b/python/pluginJava/src/com/intellij/python/community/plugin/java/psi/impl/PyJavaMethodType.java @@ -1,4 +1,4 @@ -// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.intellij.python.community.plugin.java.psi.impl; import com.intellij.psi.PsiClass; @@ -7,7 +7,7 @@ import com.intellij.psi.PsiMethod; import com.intellij.util.ArrayUtilRt; import com.intellij.util.ProcessingContext; import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.RatedResolveResult; @@ -34,7 +34,7 @@ public class PyJavaMethodType implements PyCallableType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { return getReturnType(context); } diff --git a/python/python-ast/src/com/jetbrains/python/ast/PyAstAugAssignmentStatement.kt b/python/python-ast/src/com/jetbrains/python/ast/PyAstAugAssignmentStatement.kt index a0783510fcb2..91d758c7fb86 100644 --- a/python/python-ast/src/com/jetbrains/python/ast/PyAstAugAssignmentStatement.kt +++ b/python/python-ast/src/com/jetbrains/python/ast/PyAstAugAssignmentStatement.kt @@ -1,14 +1,18 @@ // Copyright 2000-2021 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. package com.jetbrains.python.ast +import com.intellij.lang.ASTNode import com.intellij.psi.PsiElement +import com.intellij.psi.util.QualifiedName +import com.jetbrains.python.PyNames import com.jetbrains.python.PyTokenTypes import com.jetbrains.python.PythonDialectsTokenSetProvider import com.jetbrains.python.ast.impl.PyPsiUtilsCore +import com.jetbrains.python.psi.PyElementType import org.jetbrains.annotations.ApiStatus @ApiStatus.Experimental -interface PyAstAugAssignmentStatement : PyAstStatement { +interface PyAstAugAssignmentStatement : PyAstStatement, PyAstQualifiedExpression, PyAstCallSiteOwner, PyAstReferenceOwner { val target: PyAstExpression get() { return childToPsi(PythonDialectsTokenSetProvider.getInstance().expressionTokens, 0) @@ -21,6 +25,43 @@ interface PyAstAugAssignmentStatement : PyAstStatement { val operation: PsiElement? get() = PyPsiUtilsCore.getChildByFilter(this, PyTokenTypes.AUG_ASSIGN_OPERATIONS, 0) + fun isRightOperator(resolvedCallee: PyAstCallable?): Boolean { + return resolvedCallee != null && PyNames.isRightOperatorName(referencedName, resolvedCallee.getName()) + } + + fun isInplaceOperator(resolvedCallee: PyAstCallable?): Boolean { + return resolvedCallee != null && PyNames.isInplaceOperatorName(referencedName, resolvedCallee.getName()) + } + + override fun getQualifier(): PyAstExpression? { + return this.target + } + + override fun asQualifiedName(): QualifiedName? { + return PyPsiUtilsCore.asQualifiedName(this) + } + + override fun isQualified(): Boolean { + return qualifier != null + } + + override fun getReferencedName(): String? { + val t = this.operation?.node?.elementType as PyElementType + return t.specialMethodName + } + + override fun getNameElement(): ASTNode? { + return operation?.getNode() + } + + override fun getReceiver(resolvedCallee: PyAstCallable?): PyAstExpression? { + return if (isRightOperator(resolvedCallee)) this.value else this.target + } + + override fun getArguments(resolvedCallee: PyAstCallable?): List { + return listOf(if (isRightOperator(resolvedCallee)) this.target else this.value) + } + override fun acceptPyVisitor(pyVisitor: PyAstElementVisitor) { pyVisitor.visitPyAugAssignmentStatement(this) } diff --git a/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteExpression.java b/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteExpression.java index 0449d4aa7d34..9888907181e5 100644 --- a/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteExpression.java +++ b/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteExpression.java @@ -18,6 +18,7 @@ package com.jetbrains.python.ast; import org.jetbrains.annotations.ApiStatus; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; +import org.jetbrains.annotations.Unmodifiable; import java.util.List; @@ -26,7 +27,7 @@ import java.util.List; * */ @ApiStatus.Experimental -public interface PyAstCallSiteExpression extends PyAstExpression { +public interface PyAstCallSiteExpression extends PyAstCallSiteOwner, PyAstExpression { /** * Returns an expression that is treated as a receiver for this explicit or implicit (read, operator) call. @@ -38,9 +39,12 @@ public interface PyAstCallSiteExpression extends PyAstExpression { * * @param resolvedCallee optional callee corresponding to the call. Without it the receiver is deduced purely syntactically. */ + @Override @Nullable PyAstExpression getReceiver(@Nullable PyAstCallable resolvedCallee); + @Override @NotNull - List getArguments(@Nullable PyAstCallable resolvedCallee); + @Unmodifiable + List getArguments(@Nullable PyAstCallable resolvedCallee); } diff --git a/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteOwner.kt b/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteOwner.kt new file mode 100644 index 000000000000..e4f0fdcc68bc --- /dev/null +++ b/python/python-ast/src/com/jetbrains/python/ast/PyAstCallSiteOwner.kt @@ -0,0 +1,21 @@ +package com.jetbrains.python.ast + +import org.jetbrains.annotations.ApiStatus + +@ApiStatus.Experimental +interface PyAstCallSiteOwner { + /** + * Returns an expression that is treated as a receiver for this explicit or implicit (read, operator) call. + * + * + * For most operator expressions it returns the result of `getOperator()` since it naturally represents + * the object on which a special magic method is called. However for binary expressions that additionally + * can be reversible such as `__add__` and `__radd__` it also takes into account name of the + * actual callee method and chained comparisons order if any. + * + * @param resolvedCallee optional callee corresponding to the call. Without it the receiver is deduced purely syntactically. + */ + fun getReceiver(resolvedCallee: PyAstCallable?): PyAstExpression? + + fun getArguments(resolvedCallee: PyAstCallable?): List +} diff --git a/python/python-parser/src/com/jetbrains/python/PyNames.kt b/python/python-parser/src/com/jetbrains/python/PyNames.kt index c87eb1b17611..beaffdf8b5c0 100644 --- a/python/python-parser/src/com/jetbrains/python/PyNames.kt +++ b/python/python-parser/src/com/jetbrains/python/PyNames.kt @@ -678,6 +678,20 @@ object PyNames { return referencedName != null && calleeName != null && calleeName == leftToRightComparisonOperatorName(referencedName) } + private val INPLACE_OPERATOR_PATTERN = "__i([a-z]+)__".toRegex() + + @JvmStatic + private fun isInplaceOperatorName(name: String?): Boolean { + return name != null && (name.matches(INPLACE_OPERATOR_PATTERN)) + } + + @JvmStatic + fun isInplaceOperatorName(referencedName: String?, calleeName: String?): Boolean { + if (isInplaceOperatorName(calleeName)) return true + + return referencedName != null && calleeName != null && calleeName == leftToRightComparisonOperatorName(referencedName) + } + @JvmStatic fun leftToRightOperatorName(name: String?): String? { if (name == null) return null @@ -688,6 +702,20 @@ object PyNames { return name.replaceFirst("__([a-z]+)__".toRegex(), "__r$1__") } + @JvmStatic + fun inplaceToLeftOperatorName(name: String?): String? { + if (name == null) return null + + return name.replaceFirst(INPLACE_OPERATOR_PATTERN, "__$1__") + } + + @JvmStatic + fun inplaceToRightOperatorName(name: String?): String? { + if (name == null) return null + + return name.replaceFirst(INPLACE_OPERATOR_PATTERN, "__r$1__") + } + @Deprecated("use `PyAstElement.protectionLevel` instead") fun isProtected(name: @NonNls String): Boolean = ProtectionLevel.forName(name) == ProtectionLevel.PROTECTED diff --git a/python/python-parser/src/com/jetbrains/python/PyTokenTypes.java b/python/python-parser/src/com/jetbrains/python/PyTokenTypes.java index 6797c3fb5fd3..749db8f80087 100644 --- a/python/python-parser/src/com/jetbrains/python/PyTokenTypes.java +++ b/python/python-parser/src/com/jetbrains/python/PyTokenTypes.java @@ -121,19 +121,19 @@ public final class PyTokenTypes { public static final PyElementType TICK = new PyElementType("TICK");// ` public static final PyElementType EQ = new PyElementType("EQ");// = public static final PyElementType SEMICOLON = new PyElementType("SEMICOLON");// ; - public static final PyElementType PLUSEQ = new PyElementType("PLUSEQ");// += - public static final PyElementType MINUSEQ = new PyElementType("MINUSEQ");// -= - public static final PyElementType MULTEQ = new PyElementType("MULTEQ");// *= - public static final PyElementType ATEQ = new PyElementType("ATEQ"); // @= - public static final PyElementType DIVEQ = new PyElementType("DIVEQ"); // /= - public static final PyElementType FLOORDIVEQ = new PyElementType("FLOORDIVEQ"); // //= - public static final PyElementType PERCEQ = new PyElementType("PERCEQ");// %= - public static final PyElementType ANDEQ = new PyElementType("ANDEQ");// &= - public static final PyElementType OREQ = new PyElementType("OREQ");// |= - public static final PyElementType XOREQ = new PyElementType("XOREQ");// ^= - public static final PyElementType LTLTEQ = new PyElementType("LTLTEQ");// <<= - public static final PyElementType GTGTEQ = new PyElementType("GTGTEQ");// >>= - public static final PyElementType EXPEQ = new PyElementType("EXPEQ");// **= + public static final PyElementType PLUSEQ = new PyElementType("PLUSEQ", "__iadd__");// += + public static final PyElementType MINUSEQ = new PyElementType("MINUSEQ", "__isub__");// -= + public static final PyElementType MULTEQ = new PyElementType("MULTEQ", "__imul__");// *= + public static final PyElementType ATEQ = new PyElementType("ATEQ", "__imatmul__"); // @= + public static final PyElementType DIVEQ = new PyElementType("DIVEQ", "__itruediv__"); // /= + public static final PyElementType FLOORDIVEQ = new PyElementType("FLOORDIVEQ", "__ifloordiv__"); // //= + public static final PyElementType PERCEQ = new PyElementType("PERCEQ", "__imod__");// %= + public static final PyElementType ANDEQ = new PyElementType("ANDEQ", "__iand__");// &= + public static final PyElementType OREQ = new PyElementType("OREQ", "__ior__");// |= + public static final PyElementType XOREQ = new PyElementType("XOREQ", "__ixor__");// ^= + public static final PyElementType LTLTEQ = new PyElementType("LTLTEQ", "__ilshift__");// <<= + public static final PyElementType GTGTEQ = new PyElementType("GTGTEQ", "__irshift__");// >>= + public static final PyElementType EXPEQ = new PyElementType("EXPEQ", "__ipow__");// **= public static final PyElementType RARROW = new PyElementType("RARROW");// -> public static final PyElementType COLONEQ = new PyElementType("COLONEQ");// := diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyAssignmentExpression.kt b/python/python-psi-api/src/com/jetbrains/python/psi/PyAssignmentExpression.kt index df1280190e09..4e3a994795a3 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PyAssignmentExpression.kt +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyAssignmentExpression.kt @@ -13,11 +13,11 @@ interface PyAssignmentExpression : PyAstAssignmentExpression, PyExpression { /** * @return LHS of an expression (before :=), null if underlying target is not an identifier. */ - get() = super.target as PyTargetExpression? + get() = super.target as PyTargetExpression? override val assignedValue: PyExpression? /** * @return RHS of an expression (after :=), null if assigned value is omitted or not an expression. */ - get() = super.assignedValue as PyExpression? + get() = super.assignedValue as PyExpression? } diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyAugAssignmentStatement.kt b/python/python-psi-api/src/com/jetbrains/python/psi/PyAugAssignmentStatement.kt index 5ea4aa16f87e..2f40f8c5a3c0 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PyAugAssignmentStatement.kt +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyAugAssignmentStatement.kt @@ -3,10 +3,26 @@ package com.jetbrains.python.psi import com.jetbrains.python.ast.PyAstAugAssignmentStatement -interface PyAugAssignmentStatement : PyAstAugAssignmentStatement, PyStatement { +interface PyAugAssignmentStatement : PyAstAugAssignmentStatement, PyStatement, PyTypedElement, PyQualifiedExpression, PyCallSiteOwner, PyReferenceOwner { + /** + * this and [assignmentTarget] both refer to the same element, but for analysis, are split + * + * this refers to the reference before it has been assigned to + */ override val target: PyExpression - get() = super.value as PyExpression + get() = super.target as PyExpression + + /** + * see [target] + * + * this refers to the reference after it has been assigned to + */ + val assignmentTarget: PyTargetExpression override val value: PyExpression? get() = super.value as PyExpression? + + override fun getQualifier(): PyExpression? { + return super.getQualifier() as PyExpression? + } } diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallExpression.java b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallExpression.java index a6ade8907b2b..8291e5f50cc7 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallExpression.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallExpression.java @@ -1,4 +1,4 @@ -// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.jetbrains.python.psi; import com.intellij.psi.PsiElement; @@ -104,7 +104,7 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress List<@NotNull PyArgumentsMapping> multiMapArguments(@NotNull PyResolveContext resolveContext); class PyArgumentsMapping { - private final @NotNull PyCallSiteExpression myCallSiteExpression; + private final @NotNull PyCallSiteOwner myCallSiteOwner; private final @Nullable PyCallableType myCallableType; private final @NotNull List myImplicitParameters; private final @NotNull Map myMappedParameters; @@ -115,7 +115,7 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress private final @NotNull List myParametersMappedToVariadicKeywordArguments; private final @NotNull Map myMappedTupleParameters; - public PyArgumentsMapping(@NotNull PyCallSiteExpression callSiteExpression, + public PyArgumentsMapping(@NotNull PyCallSiteOwner callSiteOwner, @Nullable PyCallableType callableType, @NotNull List implicitParameters, @NotNull Map mappedParameters, @@ -125,7 +125,7 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress @NotNull List parametersMappedToVariadicPositionalArguments, @NotNull List parametersMappedToVariadicKeywordArguments, @NotNull Map tupleMappedParameters) { - myCallSiteExpression = callSiteExpression; + myCallSiteOwner = callSiteOwner; myCallableType = callableType; myImplicitParameters = implicitParameters; myMappedParameters = mappedParameters; @@ -137,7 +137,7 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress myMappedTupleParameters = tupleMappedParameters; } - public static @NotNull PyArgumentsMapping empty(@NotNull PyCallSiteExpression callSiteExpression) { + public static @NotNull PyArgumentsMapping empty(@NotNull PyCallSiteOwner callSiteExpression) { return new PyCallExpression.PyArgumentsMapping(callSiteExpression, null, Collections.emptyList(), @@ -150,8 +150,16 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress Collections.emptyMap()); } - public @NotNull PyCallSiteExpression getCallSiteExpression() { - return myCallSiteExpression; + public @NotNull PyCallSiteOwner getCallSiteOwner() { + return myCallSiteOwner; + } + + /** + * @deprecated use `getCallSiteOwner` instead + */ + @Deprecated(forRemoval = true) + public @NotNull PyCallSiteOwner getCallSiteExpression() { + return getCallSiteOwner(); } public @Nullable PyCallableType getCallableType() { diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteExpression.java b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteExpression.java index e1592a7208ba..8789ac460633 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteExpression.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteExpression.java @@ -1,4 +1,4 @@ -// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.jetbrains.python.psi; import com.jetbrains.python.ast.PyAstCallSiteExpression; @@ -12,7 +12,7 @@ import java.util.List; * Marker interface for Python expressions that are call sites for explicit or implicit function calls. * */ -public interface PyCallSiteExpression extends PyAstCallSiteExpression, PyExpression { +public interface PyCallSiteExpression extends PyAstCallSiteExpression, PyCallSiteOwner, PyExpression { /** * Returns an expression that is treated as a receiver for this explicit or implicit (read, operator) call. @@ -24,12 +24,16 @@ public interface PyCallSiteExpression extends PyAstCallSiteExpression, PyExpress * * @param resolvedCallee optional callee corresponding to the call. Without it the receiver is deduced purely syntactically. */ + @Override default @Nullable PyExpression getReceiver(@Nullable PyCallable resolvedCallee) { return (PyExpression)getReceiver((PyAstCallable)resolvedCallee); } + @Override default @NotNull List<@NotNull PyExpression> getArguments(@Nullable PyCallable resolvedCallee) { + // Safe raw cast: the returned list is @Unmodifiable, so no one can insert a non-PyExpression element, + // and all PyAstExpression instances in the PSI layer are also PyExpression at runtime. //noinspection unchecked - return (List<@NotNull PyExpression>)getArguments((PyAstCallable)resolvedCallee); + return (List)getArguments((PyAstCallable)resolvedCallee); } } diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteOwner.kt b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteOwner.kt new file mode 100644 index 000000000000..cbfa37da412b --- /dev/null +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallSiteOwner.kt @@ -0,0 +1,11 @@ +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.jetbrains.python.psi + +import org.jetbrains.annotations.ApiStatus + +@ApiStatus.Experimental +interface PyCallSiteOwner : PyElement { + fun getReceiver(resolvedCallee: PyCallable?): PyExpression? + + fun getArguments(resolvedCallee: PyCallable?): List +} diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallable.java b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallable.java index 134f0d36ba20..555f0cfa8643 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PyCallable.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyCallable.java @@ -50,7 +50,7 @@ public interface PyCallable extends PyAstCallable, PyTypedElement, PyQualifiedNa * Returns the type of the call to the callable. */ @Nullable - PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite); + PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite); /** @@ -59,7 +59,7 @@ public interface PyCallable extends PyAstCallable, PyTypedElement, PyQualifiedNa */ @Nullable PyType getCallType(@Nullable PyExpression receiver, - @Nullable PyCallSiteExpression pyCallSiteExpression, + @Nullable PyCallSiteOwner pyCallSiteExpression, @NotNull Map parameters, @NotNull TypeEvalContext context); diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java b/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java index 05746be09a8e..0d7f25dfb3d7 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java @@ -1,4 +1,4 @@ -// Copyright 2000-2021 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.jetbrains.python.psi.impl; import com.intellij.openapi.extensions.ExtensionPointName; @@ -7,6 +7,7 @@ import com.intellij.psi.PsiElement; import com.jetbrains.python.psi.AccessDirection; import com.jetbrains.python.psi.PyCallExpression; import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyExpression; @@ -41,6 +42,9 @@ public interface PyTypeProvider { @Nullable Ref getReturnType(@NotNull PyCallable callable, @NotNull TypeEvalContext context); + @Nullable + Ref getCallType(@NotNull PyFunction function, @NotNull PyCallSiteOwner callSite, @NotNull TypeEvalContext context); + @Nullable Ref getCallType(@NotNull PyFunction function, @NotNull PyCallSiteExpression callSite, @NotNull TypeEvalContext context); diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java index 7b56532b32de..1f60a56dc469 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java @@ -3,7 +3,7 @@ package com.jetbrains.python.psi.types; import com.intellij.openapi.util.text.StringUtil; import com.jetbrains.python.PyNames; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyFunction; import org.jetbrains.annotations.ApiStatus; @@ -41,7 +41,7 @@ public interface PyCallableType extends PyType { * Returns the type which is the result of calling an instance of this type. */ @Nullable - PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite); + PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite); /** * Returns the list of parameter types. diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java index 0a6acf8d1c94..690d07ca9355 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java @@ -1,4 +1,4 @@ -// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.jetbrains.python.psi.types; import com.intellij.openapi.util.Ref; @@ -6,6 +6,7 @@ import com.intellij.psi.PsiElement; import com.jetbrains.python.psi.AccessDirection; import com.jetbrains.python.psi.PyCallExpression; import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyExpression; @@ -46,6 +47,16 @@ public class PyTypeProviderBase implements PyTypeProvider { return null; } + @Override + public @Nullable Ref getCallType(@NotNull PyFunction function, + @NotNull PyCallSiteOwner callSite, + @NotNull TypeEvalContext context) { + if (callSite instanceof PyCallSiteExpression callSiteExpression) { + return getCallType(function, callSiteExpression, context); + } + return null; + } + @Override public @Nullable Ref getCallType(@NotNull PyFunction function, @NotNull PyCallSiteExpression callSite, diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/TypeEvalContext.kt b/python/python-psi-api/src/com/jetbrains/python/psi/types/TypeEvalContext.kt index e9fd21d408e8..e3162af5f8a5 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/TypeEvalContext.kt +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/TypeEvalContext.kt @@ -33,6 +33,8 @@ abstract class TypeEvalContext protected constructor() { abstract fun printTrace(): String abstract fun tracing(): Boolean + abstract fun traceWithIndent(message: String, block: () -> Unit) + @ApiStatus.Internal abstract fun assumeType(element: PyTypedElement, type: PyType?, func: (TypeEvalContext?) -> R): R? diff --git a/python/python-psi-impl/resources/messages/PyPsiBundle.properties b/python/python-psi-impl/resources/messages/PyPsiBundle.properties index b6dcb1df0008..01aa622034ad 100644 --- a/python/python-psi-impl/resources/messages/PyPsiBundle.properties +++ b/python/python-psi-impl/resources/messages/PyPsiBundle.properties @@ -1169,6 +1169,7 @@ INSP.type.checker.unexpected.argument.from.paramspec=Unexpected argument (from P INSP.type.checker.unfilled.parameter.for.paramspec=Parameter ''{0}'' unfilled (from ParamSpec ''{1}'') INSP.type.checker.unfilled.vararg=Parameter ''{0}'' unfilled, expected ''{1}'' INSP.type.checker.expected.type.from.dunder.set.got.type.instead=Expected type ''{0}'' (from ''__set__''), got ''{1}'' instead +INSP.type.checker.expected.type.from.aug.assignment.got.type.instead=Expected type ''{0}'' for augmented assignment, got ''{1}'' from operation instead INSP.type.checker.tuple.index.out.of.range=Tuple index out of range INSP.type.checker.access.to.generic.instance.variables.via.class.is.ambiguous=Access to generic instance variables via class is ambiguous INSP.type.checker.type.not.assignable=''{0}'' is not assignable to ''{1}'' diff --git a/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java b/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java index a204b0d8130b..d645851ec3da 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java +++ b/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java @@ -28,7 +28,7 @@ import com.intellij.util.Processor; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyElement; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyUtil; @@ -182,7 +182,7 @@ public final class PyCustomType implements PyClassLikeType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { return getReturnType(context); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyControlFlowBuilder.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyControlFlowBuilder.java index b9e8bc04c8f0..da935cacb294 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyControlFlowBuilder.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyControlFlowBuilder.java @@ -23,6 +23,7 @@ import com.intellij.codeInsight.controlflow.TransparentInstruction; import com.intellij.codeInsight.controlflow.impl.ConditionalInstructionImpl; import com.intellij.codeInsight.controlflow.impl.TransparentInstructionImpl; import com.intellij.openapi.util.Pair; +import com.intellij.openapi.util.Ref; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiNamedElement; import com.intellij.psi.util.PsiTreeUtil; @@ -108,6 +109,8 @@ import com.jetbrains.python.psi.impl.PyAugAssignmentStatementNavigator; import com.jetbrains.python.psi.impl.PyEvaluator; import com.jetbrains.python.psi.impl.PyImportStatementNavigator; import com.jetbrains.python.psi.impl.PyPsiUtils; +import com.jetbrains.python.psi.types.PyType; +import com.jetbrains.python.psi.types.PyTypeUtilKt; import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -255,10 +258,13 @@ public class PyControlFlowBuilder extends PyRecursiveElementVisitor { return; } - final ReadWriteInstruction.ACCESS access = PyAugAssignmentStatementNavigator.getStatementByTarget(node) != null - ? ReadWriteInstruction.ACCESS.READWRITE - : ReadWriteInstruction.ACCESS.READ; - final ReadWriteInstruction readWriteInstruction = ReadWriteInstruction.newInstruction(myBuilder, node, getName(node), access); + final ReadWriteInstruction readWriteInstruction; + if (PyAugAssignmentStatementNavigator.getStatementByTarget(node) != null) { + readWriteInstruction = ReadWriteInstruction.readWrite(myBuilder, node, getName(node), augAssignmentTypeCallback(node)); + } + else { + readWriteInstruction = ReadWriteInstruction.newInstruction(myBuilder, node, getName(node), ReadWriteInstruction.ACCESS.READ); + } myBuilder.addNodeAndCheckPending(readWriteInstruction); } @@ -1253,6 +1259,14 @@ public class PyControlFlowBuilder extends PyRecursiveElementVisitor { || element instanceof PyStatementList); } + private static @NotNull InstructionTypeCallback augAssignmentTypeCallback(@NotNull PyReferenceExpression target) { + final PyAugAssignmentStatement statement = PsiTreeUtil.getParentOfType(target, PyAugAssignmentStatement.class); + return context -> { + PyType assignmentType = statement != null ? context.getType(statement) : null; + return Ref.create(!PyTypeUtilKt.isUnknown(assignmentType) ? assignmentType : context.getType(target)); + }; + } + private void addTypeAssertionNodes(@NotNull PyElement condition, boolean positive) { addTypeAssertionNodes(condition, positive, null); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/ReadWriteInstruction.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/ReadWriteInstruction.java index 6a4b7519d477..4ee2e4b3aa35 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/ReadWriteInstruction.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/ReadWriteInstruction.java @@ -127,6 +127,13 @@ public final class ReadWriteInstruction extends InstructionImpl { return new ReadWriteInstruction(builder, element, name, ACCESS.ASSERTTYPE, getType); } + public static @NotNull ReadWriteInstruction readWrite(final @NotNull ControlFlowBuilder builder, + final @Nullable PsiElement element, + final @Nullable String name, + final @Nullable InstructionTypeCallback getType) { + return new ReadWriteInstruction(builder, element, name, ACCESS.READWRITE, getType); + } + public @Nullable Ref getType(TypeEvalContext context, @Nullable PsiElement anchor) { return myGetType.getType(context); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypeProviderWithCustomContext.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypeProviderWithCustomContext.java index 08e1fd2cb61a..17aeee4cedcf 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypeProviderWithCustomContext.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypeProviderWithCustomContext.java @@ -4,6 +4,7 @@ import com.intellij.openapi.util.Ref; import com.intellij.psi.PsiElement; import com.jetbrains.python.psi.PyCallExpression; import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyExpression; @@ -81,7 +82,7 @@ public abstract class PyTypeProviderWithCustomContext extends PyTypePro }); } - public Ref getCallType(@NotNull PyFunction function, @NotNull PyCallSiteExpression site, @NotNull Context context) { + public Ref getCallType(@NotNull PyFunction function, @NotNull PyCallSiteOwner site, @NotNull Context context) { return null; } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt index abad3931aa31..e892aa85ac2c 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt @@ -48,6 +48,7 @@ import com.jetbrains.python.psi.PyAssignmentStatement import com.jetbrains.python.psi.PyBinaryExpression import com.jetbrains.python.psi.PyCallExpression import com.jetbrains.python.psi.PyCallSiteExpression +import com.jetbrains.python.psi.PyCallSiteOwner import com.jetbrains.python.psi.PyCallable import com.jetbrains.python.psi.PyClass import com.jetbrains.python.psi.PyDecoratable @@ -248,7 +249,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext() { return null } - override fun getCallType(function: PyFunction, callSite: PyCallSiteExpression, context: Context): Ref? { + override fun getCallType(function: PyFunction, callSite: PyCallSiteOwner, context: Context): Ref? { val functionQName = function.qualifiedName if (CAST == functionQName || CAST_EXT == functionQName) { @@ -261,7 +262,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext() { } if (functionReturningCallSiteAsAType(function)) { - return callSite.getAsClassObjectType(context) + return if (callSite is PyCallSiteExpression) callSite.getAsClassObjectType(context) else null } return null diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index d61de6823244..d5d2525c832e 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -26,9 +26,10 @@ import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.inspections.quickfix.PyMakeFunctionReturnTypeQuickFix; import com.jetbrains.python.psi.PyAnnotation; import com.jetbrains.python.psi.PyAnnotationOwner; +import com.jetbrains.python.psi.PyAugAssignmentStatement; import com.jetbrains.python.psi.PyBinaryExpression; import com.jetbrains.python.psi.PyCallExpression; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyComprehensionElement; @@ -145,6 +146,12 @@ public class PyTypeCheckerInspection extends PyInspection { checkCallSite(node); } + @Override + public void visitPyAugAssignmentStatement(@NotNull PyAugAssignmentStatement node) { + checkCallSite(node); + visitPyTargetExpression(node.getAssignmentTarget()); + } + @Override public void visitPySubscriptionExpression(@NotNull PySubscriptionExpression node) { PyType operandType = myTypeEvalContext.getType(node.getOperand()); @@ -370,9 +377,21 @@ public class PyTypeCheckerInspection extends PyInspection { } if (!PyTypeChecker.match(expected, actual, myTypeEvalContext)) { + boolean isAugAssignment = node.getParent() instanceof PyAugAssignmentStatement; String message = - isDescriptor ? typeMismatchMessage(expected, actual, "INSP.type.checker.expected.type.from.dunder.set.got.type.instead") - : typeMismatchMessage(expected, actual); + isDescriptor + ? typeMismatchMessage( + expected, + actual, + "INSP.type.checker.expected.type.from.dunder.set.got.type.instead" + ) + : isAugAssignment + ? typeMismatchMessage( + expected, + actual, + "INSP.type.checker.expected.type.from.aug.assignment.got.type.instead" + ) + : typeMismatchMessage(expected, actual); registerProblem(assignedValue, message, effectiveHighlightType(ProblemHighlightType.GENERIC_ERROR_OR_WARNING)); @@ -421,7 +440,7 @@ public class PyTypeCheckerInspection extends PyInspection { ContainerUtil.exists(collectionType.getElementTypes(), Visitor::requiresTypeSpecialization); } - private @Nullable Ref getClassAttributeType(@NotNull PyTargetExpression attribute) { + private @Nullable Ref getClassAttributeType(@NotNull T attribute) { if (!attribute.isQualified()) return null; PsiElement definition = attribute.getReference(PyResolveContext.defaultContext(myTypeEvalContext)).resolve(); if (!(definition instanceof PyTargetExpression attrDefinition && PyUtil.isAttribute(attrDefinition))) return null; @@ -590,7 +609,7 @@ public class PyTypeCheckerInspection extends PyInspection { } } - private void checkCallSite(@NotNull PyCallSiteExpression callSite) { + private void checkCallSite(@NotNull PyCallSiteOwner callSite) { final List calleesResults = StreamEx .of(mapArguments(callSite, getResolveContext())) .filter(mapping -> mapping.isComplete()) @@ -641,7 +660,7 @@ public class PyTypeCheckerInspection extends PyInspection { } } - private @Nullable AnalyzeCalleeResults analyzeCallee(@NotNull PyCallSiteExpression callSite, + private @Nullable AnalyzeCalleeResults analyzeCallee(@NotNull PyCallSiteOwner callSite, @NotNull PyCallExpression.PyArgumentsMapping mapping) { final PyCallableType callableType = mapping.getCallableType(); if (callableType == null) return null; @@ -796,7 +815,7 @@ public class PyTypeCheckerInspection extends PyInspection { unfilledPositionalVarargs); } - private boolean isConstructorCall(@NotNull PyCallSiteExpression callSite) { + private boolean isConstructorCall(@NotNull PyCallSiteOwner callSite) { if (callSite instanceof PyCallExpression callExpression) { PyExpression callee = callExpression.getCallee(); if (callee != null && myTypeEvalContext.getType(callee) instanceof PyClassType calleeType && calleeType.isDefinition()) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java index 55259b6b1e08..2d7b32ac3cf1 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java @@ -33,7 +33,7 @@ import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.psi.PyArgumentList; import com.jetbrains.python.psi.PyBinaryExpression; import com.jetbrains.python.psi.PyCallExpression; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PySubscriptionExpression; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.psi.types.PyClassLikeType; @@ -52,7 +52,7 @@ import java.util.stream.Collectors; final class PyTypeCheckerInspectionProblemRegistrar { static void registerProblem(@NotNull PyInspectionVisitor visitor, - @NotNull PyCallSiteExpression callSite, + @NotNull PyCallSiteOwner callSite, @NotNull List argumentTypes, @NotNull List calleesResults, @NotNull TypeEvalContext context, @@ -66,7 +66,7 @@ final class PyTypeCheckerInspectionProblemRegistrar { } private static void registerSingleCalleeProblem(@NotNull PyInspectionVisitor visitor, - @NotNull PyCallSiteExpression callSite, + @NotNull PyCallSiteOwner callSite, @NotNull PyTypeCheckerInspection.AnalyzeCalleeResults calleeResults, @NotNull TypeEvalContext context, @Nullable ProblemHighlightType highlightOverride) { @@ -123,7 +123,7 @@ final class PyTypeCheckerInspectionProblemRegistrar { } private static void registerMultiCalleeProblem(@NotNull PyInspectionVisitor visitor, - @NotNull PyCallSiteExpression callSite, + @NotNull PyCallSiteOwner callSite, @NotNull List argumentTypes, @NotNull List calleesResults, @NotNull TypeEvalContext context, @@ -213,7 +213,7 @@ final class PyTypeCheckerInspectionProblemRegistrar { } } - private static @NotNull PsiElement getMultiCalleeElementToHighlight(@NotNull PyCallSiteExpression callSite) { + private static @NotNull PsiElement getMultiCalleeElementToHighlight(@NotNull PyCallSiteOwner callSite) { if (callSite instanceof PyCallExpression call) { final PyArgumentList argumentList = call.getArgumentList(); diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java b/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java index f44d227d0f94..70e3d38c08cb 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java @@ -52,6 +52,7 @@ import com.jetbrains.python.psi.LanguageLevel; import com.jetbrains.python.psi.PsiReferenceEx; import com.jetbrains.python.psi.PyAnnotation; import com.jetbrains.python.psi.PyAssignmentStatement; +import com.jetbrains.python.psi.PyAugAssignmentStatement; import com.jetbrains.python.psi.PyCallExpression; import com.jetbrains.python.psi.PyCallSiteExpression; import com.jetbrains.python.psi.PyCallable; @@ -145,6 +146,12 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor @Override public void visitPyTargetExpression(@NotNull PyTargetExpression node) { + // Augmented assignments (e.g., `x += 1`) do have a target expression, + // but for historical reasons it is not represented in the PSI, so delegate to the base visitor for general reference checks + if (node.getParent() instanceof PyAugAssignmentStatement) { + super.visitPyTargetExpression(node); + } + checkSlotsAndProperties(node); checkStrictClassAttributes(node); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyAugAssignmentStatementImpl.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyAugAssignmentStatementImpl.kt index 4d11650542ab..ec8003d639fe 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyAugAssignmentStatementImpl.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyAugAssignmentStatementImpl.kt @@ -2,11 +2,49 @@ package com.jetbrains.python.psi.impl import com.intellij.lang.ASTNode +import com.intellij.psi.PsiPolyVariantReference import com.jetbrains.python.psi.PyAugAssignmentStatement +import com.jetbrains.python.psi.PyCallable import com.jetbrains.python.psi.PyElementVisitor +import com.jetbrains.python.psi.PyExpression +import com.jetbrains.python.psi.PyTargetExpression +import com.jetbrains.python.psi.impl.references.PyOperatorReference +import com.jetbrains.python.psi.resolve.PyResolveContext +import com.jetbrains.python.psi.types.PyType +import com.jetbrains.python.psi.types.TypeEvalContext + +class PyAugAssignmentStatementImpl(astNode: ASTNode) : PyElementImpl(astNode), PyAugAssignmentStatement { + + override val assignmentTarget: PyTargetExpression = object : PyTargetExpressionImpl(firstChild.node) { + override fun findAssignedValue(): PyExpression { + return this@PyAugAssignmentStatementImpl + } + } -class PyAugAssignmentStatementImpl(astNode: ASTNode?) : PyElementImpl(astNode), PyAugAssignmentStatement { override fun acceptPyVisitor(pyVisitor: PyElementVisitor) { pyVisitor.visitPyAugAssignmentStatement(this) } + + /** + * the type after the operation has been applied + */ + override fun getType(context: TypeEvalContext, key: TypeEvalContext.Key): PyType? { + return PyCallExpressionHelper.getCallType(this, context, key) + } + + override fun getReference(): PsiPolyVariantReference { + return getReference(PyResolveContext.defaultContext(TypeEvalContext.codeInsightFallback(getProject()))) + } + + override fun getReference(context: PyResolveContext): PsiPolyVariantReference { + return PyOperatorReference(this, context) + } + + override fun getReceiver(resolvedCallee: PyCallable?): PyExpression? { + return if (isRightOperator(resolvedCallee)) value else target + } + + override fun getArguments(resolvedCallee: PyCallable?): List { + return listOf(if (isRightOperator(resolvedCallee)) target else value!!) + } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java index f0cdcb0abe9a..bbb7b95421cb 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java @@ -96,6 +96,21 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp return bothOperandsAreKnown ? callResultType : PyUnionType.createWeakType(callResultType); } + @Override + public PyExpression getLeftExpression() { + return PyBinaryExpression.super.getLeftExpression(); + } + + @Override + public @Nullable PyExpression getRightExpression() { + return PyBinaryExpression.super.getRightExpression(); + } + + @Override + public @Nullable PyExpression getQualifier() { + return PyBinaryExpression.super.getQualifier(); + } + private static boolean operandIsKnown(@Nullable PyExpression operand, @NotNull TypeEvalContext context) { if (operand == null) return false; diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt index d8c6f0d17e0e..85b5fc35cf6c 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt @@ -19,10 +19,12 @@ import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.psi.AccessDirection import com.jetbrains.python.psi.LanguageLevel import com.jetbrains.python.psi.PyArgumentList +import com.jetbrains.python.psi.PyAugAssignmentStatement import com.jetbrains.python.psi.PyBinaryExpression import com.jetbrains.python.psi.PyCallExpression import com.jetbrains.python.psi.PyCallExpression.PyArgumentsMapping import com.jetbrains.python.psi.PyCallSiteExpression +import com.jetbrains.python.psi.PyCallSiteOwner import com.jetbrains.python.psi.PyCallable import com.jetbrains.python.psi.PyClass import com.jetbrains.python.psi.PyDocStringOwner @@ -45,6 +47,8 @@ import com.jetbrains.python.psi.PySubscriptionExpression import com.jetbrains.python.psi.PyTupleParameter import com.jetbrains.python.psi.PyTypedElement import com.jetbrains.python.psi.PyUtil +import com.jetbrains.python.psi.impl.PyCallExpressionHelper.getCalleeType +import com.jetbrains.python.psi.impl.PyCallExpressionHelper.mapArguments import com.jetbrains.python.psi.resolve.PyResolveContext import com.jetbrains.python.psi.resolve.PyResolveUtil import com.jetbrains.python.psi.resolve.QualifiedRatedResolveResult @@ -560,7 +564,7 @@ object PyCallExpressionHelper { return getCallType(multipleResolveCallee(expression as PyReferenceOwner, resolveContext), expression, context) } - private fun getCallType(types: List, callSite: PyCallSiteExpression, context: TypeEvalContext): PyType? { + private fun getCallType(types: List, callSite: PyCallSiteOwner, context: TypeEvalContext): PyType? { return types.filter { it.isCallable } .groupBy { val callable = it.callable @@ -572,22 +576,49 @@ object PyCallExpressionHelper { .let(PyUnionType::union) } + /** + * CPython's runtime rule "prefer rhs.__r__ if rhs is a strict subclass of lhs" is not reproduced, + * as it relies on runtime classes rather than annotations, making annotation-based approximation unsound. + * + * @see #1154 + * @see #630 + */ @JvmStatic fun getCallType(expression: PyBinaryExpression, context: TypeEvalContext, @Suppress("unused") key: TypeEvalContext.Key): PyType? { + val leftExpr = expression.leftExpression ?: return null + val rightExpr = expression.rightExpression ?: return null + val resolveContext = PyResolveContext.defaultContext(context) val callableTypes = multipleResolveCallee(expression as PyReferenceOwner, resolveContext) - // TODO split normal and reflected operator methods and process them separately - // e.g. if there is matching __add__ of the left operand, don't consider signatures of __radd__ of the right operand, etc. - val matchingCallableTypes = callableTypes.filter { - val callable = it.callable + val matchingCallableTypes = callableTypes.filter { callableType -> + val callable = callableType.callable callable is PyFunction && matchesByArgumentTypes(callable, expression, context) } - return getCallType(matchingCallableTypes.ifEmpty { callableTypes }, expression, context) + + val leftType = context.getType(leftExpr) + val rightType = context.getType(rightExpr) + + val isTypeUnionSyntax = expression.isOperator("|") && + (leftType is PyClassType && leftType.isDefinition || + rightType is PyClassType && rightType.isDefinition) + + if (PyTypingTypeProvider.isInsideTypeHint(expression, context) || isTypeUnionSyntax) { + return getCallType(matchingCallableTypes.ifEmpty { callableTypes }, expression, context) + } + + val normalOperators = matchingCallableTypes.filter { !expression.isRightOperator(it.callable) } + + return if (normalOperators.isNotEmpty() && areAllTypesCoveredByCandidates(leftExpr, normalOperators, context)) { + getCallType(normalOperators, expression, context) + } + else { + getCallType(matchingCallableTypes.ifEmpty { callableTypes }, expression, context) + } } private fun getSameScopeCallablesCallTypes( types: List, - callSite: PyCallSiteExpression, + callSite: PyCallSiteOwner, context: TypeEvalContext, ): List { val firstCallable = types[0].callable @@ -597,7 +628,63 @@ object PyCallExpressionHelper { return types.map { it.getCallType(context, callSite) } } - private fun resolveOverloadsCallType(types: List, callSite: PyCallSiteExpression, context: TypeEvalContext): PyType? { + @JvmStatic + fun getCallType(statement: PyAugAssignmentStatement, context: TypeEvalContext, @Suppress("unused") key: TypeEvalContext.Key): PyType? { + val resolveContext = PyResolveContext.defaultContext(context) + val callableTypes = multipleResolveCallee(statement as PyReferenceOwner, resolveContext) + + val matchingCallableTypes = callableTypes.filter { callableType -> + val callable = callableType.callable + callable is PyFunction && matchesByArgumentTypes(callable, statement, context) + } + + val inplaceOperators = matchingCallableTypes.filter { callableType -> + val callable = callableType.callable + callable is PyFunction && statement.isInplaceOperator(callable) + } + if (inplaceOperators.isNotEmpty() && areAllTypesCoveredByCandidates(statement.target, inplaceOperators, context)) { + return getCallType(inplaceOperators, statement, context) + } + + val normalOperators = matchingCallableTypes.filter { callableType -> + val callable = callableType.callable + callable is PyFunction && + !statement.isInplaceOperator(callable) && + !statement.isRightOperator(callable) + } + if (normalOperators.isNotEmpty() && areAllTypesCoveredByCandidates(statement.target, normalOperators, context)) { + return getCallType(normalOperators, statement, context) + } + + val leftOperators = (inplaceOperators + normalOperators).distinct() + if (leftOperators.isNotEmpty() && areAllTypesCoveredByCandidates(statement.target, leftOperators, context)) { + return getCallType(leftOperators, statement, context) + } + + return getCallType(matchingCallableTypes.ifEmpty { callableTypes }, statement, context) + } + + private fun areAllTypesCoveredByCandidates( + operandExpression: PyExpression, + candidates: List, + context: TypeEvalContext, + ): Boolean { + val operandType = context.getType(operandExpression) + val operandMemberTypes = operandType.toStream().toList() + + val candidateOwners = candidates + .mapNotNull { ScopeUtil.getScopeOwner(it.callable) as? PyClass } + .mapNotNull { it.getType(context)?.toInstance() } + .toSet() + + return operandMemberTypes.all { member -> + member is PyClassType && candidateOwners.any { owner -> + PyTypeChecker.match(owner, member, context) + } + } + } + + private fun resolveOverloadsCallType(types: List, callSite: PyCallSiteOwner, context: TypeEvalContext): PyType? { val arguments = callSite.getArguments(types[0].callable) val matchingOverloads = types.filter { matchesByArgumentTypes(it.callable as PyFunction, callSite, context) } if (matchingOverloads.isEmpty()) { @@ -801,7 +888,7 @@ object PyCallExpressionHelper { * @see PyCallExpression.multiResolveCalleeFunction */ @JvmStatic - fun mapArguments(expression: PyCallSiteExpression, callableType: PyCallableType, context: TypeEvalContext): PyArgumentsMapping { + fun mapArguments(expression: PyCallSiteOwner, callableType: PyCallableType, context: TypeEvalContext): PyArgumentsMapping { val arguments = expression.getArguments(callableType.callable) val parameters = callableType.getParameters(context) ?.let { unpackParametersIfNeeded(it, arguments, context) } @@ -826,14 +913,14 @@ object PyCallExpressionHelper { } @JvmStatic - fun mapArguments(expression: PyCallSiteExpression, resolveContext: PyResolveContext): List { + fun mapArguments(expression: PyCallSiteOwner, resolveContext: PyResolveContext): List { val context = resolveContext.typeEvalContext return multiResolveCalleeFunction(expression, resolveContext).map { mapArguments(expression, it, context) } } - private fun multiResolveCalleeFunction(expression: PyCallSiteExpression, resolveContext: PyResolveContext): List { + private fun multiResolveCalleeFunction(expression: PyCallSiteOwner, resolveContext: PyResolveContext): List { when (expression) { is PyCallExpression -> { return expression.multiResolveCallee(resolveContext) @@ -866,7 +953,7 @@ object PyCallExpressionHelper { * @see mapArguments */ @JvmStatic - fun mapArguments(expression: PyCallSiteExpression, callable: PyCallable, context: TypeEvalContext): PyArgumentsMapping { + fun mapArguments(expression: PyCallSiteOwner, callable: PyCallable, context: TypeEvalContext): PyArgumentsMapping { val callableType = context.getType(callable) as? PyCallableType? ?: return PyArgumentsMapping.empty(expression) @@ -1353,7 +1440,7 @@ object PyCallExpressionHelper { } } - private fun matchesByArgumentTypes(function: PyFunction, callSite: PyCallSiteExpression, context: TypeEvalContext): Boolean { + private fun matchesByArgumentTypes(function: PyFunction, callSite: PyCallSiteOwner, context: TypeEvalContext): Boolean { val fullMapping = mapArguments(callSite, function, context) if (!fullMapping.isComplete) return false @@ -1515,7 +1602,7 @@ object PyCallExpressionHelper { private fun filterExplicitParameters( parameters: List, callable: PyCallable?, - callSite: PyCallSiteExpression, + callSite: PyCallSiteOwner, resolveContext: PyResolveContext, ): List { val implicitOffset: Int diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index f8691b485413..47fddb6933cf 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -44,7 +44,7 @@ import com.jetbrains.python.psi.PsiQuery; import com.jetbrains.python.psi.PyAnnotation; import com.jetbrains.python.psi.PyAssignmentStatement; import com.jetbrains.python.psi.PyCallExpression; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyDecorator; import com.jetbrains.python.psi.PyDecoratorList; @@ -241,7 +241,7 @@ public class PyFunctionImpl extends PyBaseElementImpl implements } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { for (PyTypeProvider typeProvider : PyTypeProvider.EP_NAME.getExtensionList()) { final Ref typeRef = typeProvider.getCallType(this, callSite, context); if (typeRef != null) { @@ -273,7 +273,7 @@ public class PyFunctionImpl extends PyBaseElementImpl implements @Override public @Nullable PyType getCallType(@Nullable PyExpression receiver, - @Nullable PyCallSiteExpression callSiteExpression, + @Nullable PyCallSiteOwner callSiteExpression, @NotNull Map parameters, @NotNull TypeEvalContext context) { @Nullable PyType type = context.getReturnType(this); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java index e71e90dbbd9c..280bde716c1e 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java @@ -5,7 +5,7 @@ import com.intellij.lang.ASTNode; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.codeInsight.controlflow.ControlFlowCache; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyElementVisitor; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyLambdaExpression; @@ -121,13 +121,13 @@ public class PyLambdaExpressionImpl extends PyElementImpl implements PyLambdaExp } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { return context.getReturnType(this); } @Override public @Nullable PyType getCallType(@Nullable PyExpression receiver, - @Nullable PyCallSiteExpression pyCallSiteExpression, + @Nullable PyCallSiteOwner pyCallSiteExpression, @NotNull Map parameters, @NotNull TypeEvalContext context) { return context.getReturnType(this); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java index 5cdf70da2139..5cd23f7470d6 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java @@ -331,7 +331,7 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere // 1. If WRITE instructions are found on all possible execution paths: // - Returns a union type combining the types from all getType() calls on those instructions // - // 2. If a WRITE instruction involving just the `qualifier` is found on any path + // 2. If a WRITE instruction involving just the `qualifier` is found on any path // (via PyTargetExpression or PyNamedParameter): // - The analysis stops and returns null, ignoring any other paths // @@ -536,8 +536,11 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere final ScopeOwner scopeOwner = ScopeUtil.getScopeOwner(anchor); final String name = ((PyElement)target).getName(); if (scopeOwner != null && name != null) { - if (!ScopeUtil.getElementsOfAccessType(name, scopeOwner, ReadWriteInstruction.ACCESS.ASSERTTYPE).isEmpty() || - (target instanceof PyTargetExpression || target instanceof PyNamedParameter) && ScopeUtil.getScopeOwner(target) == scopeOwner) { + if (!ScopeUtil.getElementsOfAccessType(name, scopeOwner, ReadWriteInstruction.ACCESS.ASSERTTYPE).isEmpty() + || (target instanceof PyTargetExpression + || target instanceof PyNamedParameter + || !ScopeUtil.getElementsOfAccessType(name, scopeOwner, ReadWriteInstruction.ACCESS.READWRITE).isEmpty()) + && ScopeUtil.getScopeOwner(target) == scopeOwner) { final PyType type = getTypeByControlFlow(name, context, anchor, scopeOwner).type(); if (!isUnknown(type)) { return Ref.create(type); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/references/PyOperatorReference.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/references/PyOperatorReference.java index 6bc5746726ab..f46b8649f6de 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/references/PyOperatorReference.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/references/PyOperatorReference.java @@ -21,6 +21,7 @@ import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.AccessDirection; import com.jetbrains.python.psi.LanguageLevel; +import com.jetbrains.python.psi.PyAugAssignmentStatement; import com.jetbrains.python.psi.PyBinaryExpression; import com.jetbrains.python.psi.PyCallSiteExpression; import com.jetbrains.python.psi.PyClass; @@ -37,6 +38,7 @@ import com.jetbrains.python.psi.types.PyClassType; import com.jetbrains.python.psi.types.PyType; import com.jetbrains.python.psi.types.PyTypeUtil; import com.jetbrains.python.psi.types.TypeEvalContext; +import kotlin.Unit; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -51,7 +53,10 @@ public class PyOperatorReference extends PyReferenceImpl { @Override protected @NotNull List resolveInner() { - if (myElement instanceof PyBinaryExpression expr) { + if (myElement instanceof PyAugAssignmentStatement stmt) { + return resolveInlineAndLeftAndRightOperators(stmt, stmt.getReferencedName()); + } + else if (myElement instanceof PyBinaryExpression expr) { final String name = expr.getReferencedName(); if (PyNames.CONTAINS.equals(name)) { return resolveMember(expr.getRightExpression(), name); @@ -103,23 +108,33 @@ public class PyOperatorReference extends PyReferenceImpl { final List result = new ArrayList<>(); final TypeEvalContext typeEvalContext = myContext.getTypeEvalContext(); - typeEvalContext.trace("Trying to resolve left operator"); - typeEvalContext.traceIndent(); - try { + typeEvalContext.traceWithIndent("Trying to resolve left operator", () -> { result.addAll(resolveMember(expr.getReceiver(null), name)); - } - finally { - typeEvalContext.traceUnindent(); - } - typeEvalContext.trace("Trying to resolve right operator"); - typeEvalContext.traceIndent(); - try { + return Unit.INSTANCE; + }); + typeEvalContext.traceWithIndent("Trying to resolve right operator", () -> { result.addAll(resolveMember(expr.getRightExpression(), PyNames.leftToRightOperatorName(name))); - } - finally { - typeEvalContext.traceUnindent(); - } + return Unit.INSTANCE; + }); + return result; + } + private @NotNull List resolveInlineAndLeftAndRightOperators(@NotNull PyAugAssignmentStatement stmt, @Nullable String name) { + final List result = new ArrayList<>(); + + final TypeEvalContext typeEvalContext = myContext.getTypeEvalContext(); + typeEvalContext.traceWithIndent("Trying to resolve inplace operator", () -> { + result.addAll(resolveMember(stmt.getReceiver(null), name)); + return Unit.INSTANCE; + }); + typeEvalContext.traceWithIndent("Trying to resolve left operator", () -> { + result.addAll(resolveMember(stmt.getReceiver(null), PyNames.inplaceToLeftOperatorName(name))); + return Unit.INSTANCE; + }); + typeEvalContext.traceWithIndent("Trying to resolve right operator", () -> { + result.addAll(resolveMember(stmt.getValue(), PyNames.inplaceToRightOperatorName(name))); + return Unit.INSTANCE; + }); return result; } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java index 2fc0c8febebd..f2d254f9c16d 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java @@ -6,7 +6,7 @@ import com.intellij.util.ArrayUtilRt; import com.intellij.util.ProcessingContext; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyFunction; @@ -89,7 +89,7 @@ public class PyCallableTypeImpl implements PyCallableType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { if (!PyTypeChecker.hasGenerics(myReturnType, context)) { return PyNarrowedType.Companion.bindIfNeeded(myReturnType, callSite); } @@ -105,7 +105,7 @@ public class PyCallableTypeImpl implements PyCallableType { @NotNull Map actualParameters, @NotNull Collection allParameters, @Nullable PyExpression receiver, - @NotNull PyCallSiteExpression callsite, + @NotNull PyCallSiteOwner callsite, @NotNull TypeEvalContext context) { final var substitutions = PyTypeChecker.unifyGenericCall(receiver, actualParameters, context); final var substitutionsWithUnresolvedReturnGenerics = diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java index 2687f4b860c2..1577481a0902 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java @@ -33,6 +33,7 @@ import com.jetbrains.python.psi.Property; import com.jetbrains.python.psi.PyAnnotationOwner; import com.jetbrains.python.psi.PyCallExpression; import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyElsePart; @@ -71,7 +72,6 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.Collections; import java.util.HashSet; -import java.util.Iterator; import java.util.LinkedHashSet; import java.util.List; import java.util.Map; @@ -398,13 +398,13 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { return getPossibleCallType(context, callSite); } - private @Nullable PyType getPossibleCallType(@NotNull TypeEvalContext context, @Nullable PyCallSiteExpression callSite) { + private @Nullable PyType getPossibleCallType(@NotNull TypeEvalContext context, @Nullable PyCallSiteOwner callSite) { if (!isDefinition()) { - return PyUtil.getReturnTypeOfMember(this, PyNames.CALL, callSite, context); + return PyUtil.getReturnTypeOfMember(this, PyNames.CALL, callSite instanceof PyCallSiteExpression callSiteExpression ? callSiteExpression : null , context); } else { return withUserDataCopy(new PyClassTypeImpl(getPyClass(), false)); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java index 3076b7ea34d4..52178f3e0cdb 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java @@ -20,7 +20,7 @@ import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.PsiElement; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyPsiFacade; import org.jetbrains.annotations.NotNull; @@ -56,7 +56,7 @@ public class PyCollectionTypeImpl extends PyClassTypeImpl implements PyCollectio } @Override - public @Nullable PyType getCallType(final @NotNull TypeEvalContext context, final @Nullable PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(final @NotNull TypeEvalContext context, final @Nullable PyCallSiteOwner callSite) { return getReturnType(context); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyDescriptorTypeUtil.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyDescriptorTypeUtil.java index eca066b5fc78..20f6162a5255 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyDescriptorTypeUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyDescriptorTypeUtil.java @@ -6,7 +6,6 @@ import com.jetbrains.python.psi.AccessDirection; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyFunction; import com.jetbrains.python.psi.PyQualifiedExpression; -import com.jetbrains.python.psi.PyTargetExpression; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.RatedResolveResult; @@ -39,7 +38,7 @@ public final class PyDescriptorTypeUtil { return getTypeFromSyntheticDunderGetCall(expression, attributeType, context); } - public static @Nullable Ref getExpectedValueTypeForDunderSet(@NotNull PyTargetExpression targetExpression, + public static @Nullable Ref getExpectedValueTypeForDunderSet(@NotNull PyQualifiedExpression targetExpression, @Nullable PyType attributeType, @NotNull TypeEvalContext context) { final PyClassLikeType targetType = as(attributeType, PyClassLikeType.class); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyFunctionTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyFunctionTypeImpl.java index bfb6bce39e44..ebc965469e73 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyFunctionTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyFunctionTypeImpl.java @@ -7,7 +7,7 @@ import com.intellij.util.ProcessingContext; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyReferenceExpression; @@ -61,7 +61,7 @@ public class PyFunctionTypeImpl implements PyFunctionType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { return myCallable.getCallType(context, callSite); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNamedTupleType.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNamedTupleType.java index a600155186dd..3b1fdfbd4b67 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNamedTupleType.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNamedTupleType.java @@ -12,7 +12,7 @@ import com.intellij.util.ObjectUtils; import com.intellij.util.ProcessingContext; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.psi.LanguageLevel; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyQualifiedNameOwner; @@ -110,7 +110,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType { } @Override - public @Nullable PyNamedTupleType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyNamedTupleType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { if (isDefinition()) { return getCallDefinitionType(callSite, context); } @@ -205,7 +205,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType { return this; } - private @NotNull PyNamedTupleType getCallDefinitionType(@NotNull PyCallSiteExpression callSite, @NotNull TypeEvalContext context) { + private @NotNull PyNamedTupleType getCallDefinitionType(@NotNull PyCallSiteOwner callSite, @NotNull TypeEvalContext context) { if (!myTyped) { final List arguments = callSite.getArguments(null); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNarrowedType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNarrowedType.kt index 17700e2c2b88..82bf8101bf3f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNarrowedType.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyNarrowedType.kt @@ -1,6 +1,6 @@ package com.jetbrains.python.psi.types -import com.jetbrains.python.psi.PyCallSiteExpression +import com.jetbrains.python.psi.PyCallSiteOwner import com.jetbrains.python.psi.PyClass import com.jetbrains.python.psi.PyElement import com.jetbrains.python.psi.PyReferenceExpression @@ -14,7 +14,7 @@ import org.jetbrains.annotations.ApiStatus class PyNarrowedType private constructor( pyClass: PyClass, val qname: String?, - val original: PyCallSiteExpression?, + val original: PyCallSiteOwner?, val negated: Boolean, val typeIs: Boolean, val narrowedType: PyType?, @@ -24,7 +24,7 @@ class PyNarrowedType private constructor( return PyNarrowedType(pyClass, qname, original, !negated, typeIs, narrowedType) } - fun bind(callExpression: PyCallSiteExpression, name: String): PyNarrowedType { + fun bind(callExpression: PyCallSiteOwner, name: String): PyNarrowedType { return PyNarrowedType(pyClass, name, callExpression, negated, typeIs, narrowedType) } @@ -73,7 +73,7 @@ class PyNarrowedType private constructor( return PyNarrowedType(pyClass, null, null, false, typeIs, returnType) } - fun bindIfNeeded(type: PyType?, callSiteExpression: PyCallSiteExpression?): PyType? { + fun bindIfNeeded(type: PyType?, callSiteExpression: PyCallSiteOwner?): PyType? { if (type is PyNarrowedType && callSiteExpression != null) { val arguments = callSiteExpression.getArguments(null) val pyReferenceExpression = arguments.firstOrNull() diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PySelfType.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PySelfType.java index 9be147a04f9f..1406eccba05f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PySelfType.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PySelfType.java @@ -5,7 +5,7 @@ import com.intellij.psi.PsiElement; import com.intellij.util.ProcessingContext; import com.intellij.util.Processor; import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyCallable; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyExpression; @@ -166,7 +166,7 @@ public final class PySelfType implements PyTypeParameterType, PyClassType { } @Override - public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) { + public @Nullable PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteOwner callSite) { if (isDefinition()) { return toInstance(); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeInferenceCspFactory.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeInferenceCspFactory.kt index 1f92dfaaa6d1..431f059501f6 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeInferenceCspFactory.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeInferenceCspFactory.kt @@ -4,6 +4,7 @@ import com.intellij.openapi.util.registry.Registry import com.jetbrains.python.psi.PyArgumentList import com.jetbrains.python.psi.PyCallExpression import com.jetbrains.python.psi.PyCallSiteExpression +import com.jetbrains.python.psi.PyCallSiteOwner import com.jetbrains.python.psi.PyDecoratorList import com.jetbrains.python.psi.PyExpression import com.jetbrains.python.psi.PyTypedElement @@ -24,7 +25,7 @@ object PyTypeInferenceCspFactory { @JvmStatic fun unifyReceiver(argsMapping: PyCallExpression.PyArgumentsMapping, context: TypeEvalContext): GenericSubstitutions { - val callSite = argsMapping.callSiteExpression + val callSite = argsMapping.callSiteOwner val callableType = argsMapping.callableType val receiver = callSite.getReceiver(callableType?.callable) if (!Registry.`is`("python.use.csp.type.inference")) { @@ -40,7 +41,7 @@ object PyTypeInferenceCspFactory { @JvmStatic fun unifyGenericCall( - callSite: PyCallSiteExpression?, + callSite: PyCallSiteOwner?, receiver: PyExpression?, callableType: PyCallableType?, mappedParameters: Map, @@ -60,7 +61,7 @@ object PyTypeInferenceCspFactory { // TODO: wrong parameter mapping passed by testExplicitlyParameterizedGenericConstructorCall: self missing? private fun doUnifyFunctionCall( - callSite: PyCallSiteExpression?, + callSite: PyCallSiteOwner?, receiver: PyExpression?, callableType: PyCallableType?, mappedParameters: Map, diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt index 17c39c44b30d..96307605c695 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt @@ -6,7 +6,7 @@ import com.jetbrains.python.PyNames import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.codeInsight.typing.isProtocol import com.jetbrains.python.psi.PyCallExpression -import com.jetbrains.python.psi.PyCallSiteExpression +import com.jetbrains.python.psi.PyCallSiteOwner import com.jetbrains.python.psi.PyClass import com.jetbrains.python.psi.PyDictLiteralExpression import com.jetbrains.python.psi.PyExpression @@ -33,7 +33,7 @@ class PyTypedDictType( return fields[key]?.type } - override fun getCallType(context: TypeEvalContext, callSite: PyCallSiteExpression): PyType? { + override fun getCallType(context: TypeEvalContext, callSite: PyCallSiteOwner): PyType? { return if (isDefinition) toInstance() else null } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypingNewType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypingNewType.kt index d2e65b8d3dcc..d85b1eb28f58 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypingNewType.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypingNewType.kt @@ -3,7 +3,7 @@ package com.jetbrains.python.psi.types import com.jetbrains.python.PyNames import com.jetbrains.python.psi.AccessDirection -import com.jetbrains.python.psi.PyCallSiteExpression +import com.jetbrains.python.psi.PyCallSiteOwner import com.jetbrains.python.psi.PyExpression import com.jetbrains.python.psi.PyTargetExpression import com.jetbrains.python.psi.resolve.PyResolveContext @@ -17,7 +17,7 @@ class PyTypingNewType( override val declarationElement: PyTargetExpression, ) : PyClassType by classType { - override fun getCallType(context: TypeEvalContext, callSite: PyCallSiteExpression): PyType { + override fun getCallType(context: TypeEvalContext, callSite: PyCallSiteOwner): PyType { return PyTypingNewType(classType.toInstance(), name, declarationElement) } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/TypeEvalContextImpl.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/TypeEvalContextImpl.kt index 6c6478dc5e73..0c79f065d036 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/TypeEvalContextImpl.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/TypeEvalContextImpl.kt @@ -115,6 +115,17 @@ open class TypeEvalContextImpl internal constructor( return myTrace != null } + override fun traceWithIndent(message: String, block: () -> Unit) { + trace(message) + traceIndent() + try { + block() + } + finally { + traceUnindent() + } + } + @ApiStatus.Internal override fun assumeType(element: PyTypedElement, type: PyType?, func: (TypeEvalContext?) -> R): R? { if (!Registry.`is`("python.use.better.control.flow.type.inference")) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyDefUseUtil.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyDefUseUtil.java index c477de34fc61..459226b623ea 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyDefUseUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyDefUseUtil.java @@ -33,7 +33,7 @@ import com.jetbrains.python.codeInsight.controlflow.ReadWriteInstruction; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; import com.jetbrains.python.psi.PyAugAssignmentStatement; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyElement; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyImplicitImportNameDefiner; @@ -98,7 +98,7 @@ public final class PyDefUseUtil { QualifiedName varQname = QualifiedName.fromDottedString(varName); final Collection result = new LinkedHashSet<>(); - final HashMap pendingTypeGuard = new HashMap<>(); + final HashMap pendingTypeGuard = new HashMap<>(); final Ref<@NotNull Boolean> foundPrefixWrite = Ref.create(false); final Ref<@NotNull Boolean> foundPrefixCall = Ref.create(false); iteratePrev(startNum, controlFlow, diff --git a/python/src/com/jetbrains/python/inspections/quickfix/PyChangeSignatureQuickFix.java b/python/src/com/jetbrains/python/inspections/quickfix/PyChangeSignatureQuickFix.java index 7f042e0f825b..3c740e0d941a 100644 --- a/python/src/com/jetbrains/python/inspections/quickfix/PyChangeSignatureQuickFix.java +++ b/python/src/com/jetbrains/python/inspections/quickfix/PyChangeSignatureQuickFix.java @@ -1,4 +1,4 @@ -// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.jetbrains.python.inspections.quickfix; import com.google.common.collect.Iterators; @@ -20,7 +20,7 @@ import com.intellij.xml.util.XmlStringUtil; import com.jetbrains.python.PyBundle; import com.jetbrains.python.psi.PyCallExpression; import com.jetbrains.python.psi.PyCallExpression.PyArgumentsMapping; -import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PyCallSiteOwner; import com.jetbrains.python.psi.PyElement; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyFunction; @@ -64,7 +64,7 @@ public final class PyChangeSignatureQuickFix extends LocalQuickFixOnPsiElement { final PyFunction function = as(mapping.getCallableType().getCallable(), PyFunction.class); assert function != null; Supplier>> extraParamsSupplier = () -> { - final PyCallSiteExpression callSiteExpression = mapping.getCallSiteExpression(); + final var callSiteExpression = mapping.getCallSiteOwner(); int positionalParamAnchor = -1; final PyParameter[] parameters = function.getParameterList().getParameters(); for (PyParameter parameter : parameters) { @@ -93,7 +93,7 @@ public final class PyChangeSignatureQuickFix extends LocalQuickFixOnPsiElement { } return newParameters; }; - return new PyChangeSignatureQuickFix(function, extraParamsSupplier, mapping.getCallSiteExpression()); + return new PyChangeSignatureQuickFix(function, extraParamsSupplier, mapping.getCallSiteOwner()); } public static @NotNull PyChangeSignatureQuickFix forMismatchingMethods(@NotNull PyFunction function, @NotNull PyFunction complementary) { @@ -113,7 +113,7 @@ public final class PyChangeSignatureQuickFix extends LocalQuickFixOnPsiElement { } private final @NotNull Supplier>> myExtraParametersSupplier; - private final @Nullable SmartPsiElementPointer myOriginalCallSiteExpression; + private final @Nullable SmartPsiElementPointer myOriginalCallSiteExpression; /** @@ -122,7 +122,7 @@ public final class PyChangeSignatureQuickFix extends LocalQuickFixOnPsiElement { */ private PyChangeSignatureQuickFix(@NotNull PyFunction function, @NotNull Supplier>> extraParametersSupplier, - @Nullable PyCallSiteExpression expression) { + @Nullable PyCallSiteOwner expression) { super(function); myExtraParametersSupplier = () -> ContainerUtil.sorted(extraParametersSupplier.get(), Comparator.comparingInt(p -> p.getFirst())); if (expression != null) { @@ -172,7 +172,7 @@ public final class PyChangeSignatureQuickFix extends LocalQuickFixOnPsiElement { } }; - final PyCallSiteExpression originalCallSite = myOriginalCallSiteExpression != null ? myOriginalCallSiteExpression.getElement() : null; + final var originalCallSite = myOriginalCallSiteExpression != null ? myOriginalCallSiteExpression.getElement() : null; try { if (originalCallSite != null) { originalCallSite.putUserData(CHANGE_SIGNATURE_ORIGINAL_CALL, true); diff --git a/python/src/com/jetbrains/python/testing/pyTestParametrized/PyTestParametrizedTypeProvider.kt b/python/src/com/jetbrains/python/testing/pyTestParametrized/PyTestParametrizedTypeProvider.kt index 512db62f1a7e..e02efb1cfe97 100644 --- a/python/src/com/jetbrains/python/testing/pyTestParametrized/PyTestParametrizedTypeProvider.kt +++ b/python/src/com/jetbrains/python/testing/pyTestParametrized/PyTestParametrizedTypeProvider.kt @@ -1,4 +1,4 @@ -// Copyright 2000-2025 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. package com.jetbrains.python.testing.pyTestParametrized import com.intellij.openapi.util.Ref diff --git a/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py b/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py index d5123b36b4a4..5608f73acbbb 100644 --- a/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py @@ -1,9 +1,7 @@ -from typing import Any - x = 42 # print('commented') -def func() -> int | Any: +def func() -> int: return x ** 2 diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index 227b07765920..7e29efbcee00 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -1236,7 +1236,7 @@ public class Py3TypeTest extends PyTestCase { } public void testNumpyResolveRaterDoesNotIncreaseRateForNotNdarrayRightOperatorFoundInStub() { - doMultiFileTest("D1 | D2", + doMultiFileTest("D1", """ class D1(object): pass @@ -6079,6 +6079,536 @@ public class Py3TypeTest extends PyTestCase { """); } + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentIAddSameType() { + doTest("MutableContainer", """ + class MutableContainer: + def __iadd__(self, other: int) -> MutableContainer: + return self + + m = MutableContainer() + m += 1 + expr = m + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentIAddSelf() { + doTest("MutableContainer", """ + from typing import Self + + class MutableContainer: + def __iadd__(self, other: int) -> Self: + return self + + m = MutableContainer() + m += 1 + expr = m + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentIAddDifferentType() { + doTest("str", """ + class IAddReturnsDifferent: + def __iadd__(self, other: int) -> str: + return "result" + + d = IAddReturnsDifferent() + d += 1 + expr = d + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentFallbackToAdd() { + doTest("int", """ + class AddOnly: + def __add__(self, other: int) -> int: + return 1 + + a = AddOnly() + a += 4 + expr = a + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentFallbackToRadd() { + doTest("float | int", """ + class NoOps: + pass + + class RightOperand: + def __radd__(self, other: NoOps) -> float: + return 1.0 + + n = NoOps() + n += RightOperand() + expr = n + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentIAddSignatureMismatchFallbackToAdd() { + doTest("float | int", """ + class IAddAnnotatedAdd: + def __iadd__(self, other: str) -> IAddAnnotatedAdd: + return self + def __add__(self, other: int) -> float: + return 1.0 + + ia = IAddAnnotatedAdd() + ia += 5 + expr = ia + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentBuiltinInt() { + doTest("int", """ + x: int = 1 + x += 1 + expr = x + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentBuiltinIntWidensToFloat() { + doTest("float | int", """ + y: int = 1 + y += 1.5 + expr = y + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentBuiltinList() { + doTest("list[int]", """ + lst: list[int] = [1, 2] + lst += [3, 4] + expr = lst + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentBuiltinStr() { + doTest("str", """ + s: str = "hello" + s += " world" + expr = s + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentGenericIAdd() { + doTest("MyList[int]", """ + class MyList[T]: + def __iadd__(self, other: list[T]) -> MyList[T]: + return self + + ml = MyList[int]() + ml += [1, 2, 3] + expr = ml + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentInLoop() { + doTest("Accumulator", """ + class Accumulator: + def __iadd__(self, other: int) -> Accumulator: + return self + + acc = Accumulator() + for i in range(10): + acc += i + expr = acc + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentTypeChangesInLoop() { + doTest("int", """ + class Counter: + def __add__(self, other: int) -> int: + return 0 + + c = Counter() + while True: + c += 1 + if bool(): + break + expr = c + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentSubOperator() { + doTest("str", """ + class SubOnly: + def __sub__(self, other: int) -> str: + return "" + + sub = SubOnly() + sub -= 1 + expr = sub + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentMulOperator() { + doTest("float | int", """ + class MulOnly: + def __mul__(self, other: int) -> float: + return 0.0 + + mul = MulOnly() + mul *= 3 + expr = mul + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentTruedivOperator() { + doTest("complex | float | int", """ + class DivOnly: + def __truediv__(self, other: int) -> complex: + return 0j + + div = DivOnly() + div /= 2 + expr = div + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentSubclassIAdd() { + doTest("Base", """ + class Base: + def __iadd__(self, other: int) -> Base: + return self + + class Child(Base): + def __iadd__(self, other: int) -> Child: + return self + + b: Base = Child() + b += 1 + expr = b + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentNoneReturn() { + doTest("None", """ + class BadIAdd: + def __iadd__(self, other: int) -> None: + pass + + bad = BadIAdd() + bad += 1 + expr = bad + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentIAddPrecedenceOverAdd() { + doTest("str", """ + class Multi: + def __iadd__(self, other: int) -> str: + return "" + def __add__(self, other: int) -> float: + return 0.0 + + p = Multi() + p += 1 + expr = p + """); + } + + @TestFor(issues = "PY-87051") + public void testAugmentedAssignmentTypeNarrowingLiteral() { + doTest("int", """ + from typing import Literal + + one: Literal[1] = 1 + x = one + x += 1 + expr = x + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionLeftPrecedenceOverRight() { + doTest("str", """ + class A: + def __add__(self, other: B) -> str: ... + + class B: + def __radd__(self, other: A) -> int: ... + + expr = A() + B() + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionFallbackToRadd() { + doTest("bool", """ + class E: + pass + + class F: + def __radd__(self, other: E) -> bool: ... + + expr = E() + F() + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionRtruediv() { + doTest("str", """ + class D1: + pass + + class D2: + def __rtruediv__(self, other: D1) -> str: ... + + expr = D1() / D2() + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionUnionLeftAllHaveAdd() { + doTest("float | int | str", """ + class A: + def __add__(self, other: int) -> float: ... + + class B: + def __add__(self, other: int) -> str: ... + + x: A | B = A() + expr = x + 1 + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionDifferentReturnTypes() { + doTest("str", """ + class A: + def __sub__(self, other: int) -> str: ... + + expr = A() - 1 + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionMulPrecedence() { + doTest("str", """ + class A: + def __mul__(self, other: B) -> str: ... + + class B: + def __rmul__(self, other: A) -> int: ... + + expr = A() * B() + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionDoNotPreferReflectedIfItDoesNotMatchArguments() { + doTest("str", """ + class A: + def __add__(self, other: B) -> str: ... + + class B(A): + def __radd__(self, other: int) -> int: ... + + expr = A() + B() + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionDoNotPreferReflectedForUnrelatedTypes() { + doTest("str", """ + class A: + def __mul__(self, other: B) -> str: ... + + class B: + def __rmul__(self, other: A) -> int: ... + + expr = A() * B() + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentUnionIAddAndAdd() { + doTest("P | str", """ + class P: + def __iadd__(self, other: int) -> P: ... + + class Q: + def __add__(self, other: int) -> str: ... + + u: P | Q = P() + u += 1 + expr = u + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentUnionAllIAdd() { + doTest("str | bool", """ + class A: + def __iadd__(self, other: int) -> str: ... + + class B: + def __iadd__(self, other: int) -> bool: ... + + x: A | B = A() + x += 1 + expr = x + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentUnionSubOperator() { + doTest("int | str", """ + class A: + def __isub__(self, other: int) -> int: ... + + class B: + def __sub__(self, other: int) -> str: ... + + x: A | B = A() + x -= 1 + expr = x + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentUnionInplacePrecedencePerClass() { + doTest("str | bool", """ + class A: + def __iadd__(self, other: int) -> str: ... + def __add__(self, other: int) -> float: ... + + class B: + def __iadd__(self, other: int) -> bool: ... + def __add__(self, other: int) -> complex: ... + + x: A | B = A() + x += 1 + expr = x + """); + } + + @TestFor(issues = "PY-80622") + public void testAugmentedAssignmentInheritedLeftOperatorMatches() { + doTest("str", """ + from typing import Any + + class Super: + def __iadd__(self, other: Any) -> str: ... + + class Sub(Super): + pass + + class Operand: + def __radd__(self, other: Super) -> int: ... + + x = Sub() + x += Operand() + expr = x + """); + } + + @TestFor(issues = "PY-80622") + public void testBinaryExpressionInheritedLeftOperatorMatches() { + doTest("int", """ + class Right: + def __radd__(self, other: 'Super') -> str: ... + + class Super: + def __add__(self, other: Right) -> int: ... + + class Sub(Super): + pass + + expr = Sub() + Right() + """); + } + + /** + * CPython's runtime rule "prefer rhs.__r__ if rhs is a strict subclass of lhs" is not reproduced, + * as it relies on runtime classes rather than annotations, making annotation-based approximation unsound. + * + * @see #1154 + * @see #630 + */ + @TestFor(issues = "PY-80622") + public void testBinaryExpressionWhenRightOperandIsSubtypeOfLeft() { + doTest("str", """ + class A: + def __add__(self, other: B) -> str: ... + + class B(A): + def __radd__(self, other: A) -> int: ... + + expr = A() + B() + """); + } + + /** + * CPython's runtime rule "prefer rhs.__r__ if rhs is a strict subclass of lhs" is not reproduced, + * as it relies on runtime classes rather than annotations, making annotation-based approximation unsound. + * + * @see #1154 + * @see #630 + */ + @TestFor(issues = "PY-80622") + public void testBinaryExpressionWhenRightOperandIsInheritedSubtypeOfLeft() { + doTest("str", """ + class A: + def __add__(self, other: BBase) -> str: ... + + class BBase(A): + def __radd__(self, other: A) -> int: ... + + class B(BBase): + pass + + expr = A() + B() + """); + } + + /** + * CPython's runtime rule "prefer rhs.__r__ if rhs is a strict subclass of lhs" is not reproduced, + * as it relies on runtime classes rather than annotations, making annotation-based approximation unsound. + * + * @see #1154 + * @see #630 + */ + @TestFor(issues = "PY-80622") + public void testBinaryExpressionWhenRightOperandIsUnionSubtypeOfLeft() { + doTest("str", """ + class A: + def __add__(self, other: BBase) -> str: ... + + class BBase(A): + def __radd__(self, other: A) -> int: ... + + class B1(BBase): + pass + + class B2(BBase): + pass + + x: B1 | B2 = B1() + expr = A() + x + """); + } + private void doTest(final String expectedType, final String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class); diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index 0816679efba3..0fda70a5f72f 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -5340,4 +5340,120 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { foo(1, "hello", name=42) """); } + + @TestFor(issues = "PY-6426") + public void testAugmentedAssignmentArguments() { + doTestByText(""" + class A: + def __iadd__(self, other: int) -> str: ... + + a = A() + a += 1 + + a = A() + a += "a" + """); + doTestByText(""" + class A: + def __add__(self, other: int) -> str: ... + + a = A() + a += 1 + + a = A() + a += "a" + """); + doTestByText(""" + class A: pass + + class B: + def __radd__(self, other: A) -> str: ... + + a = A() + a += B() + """); + } + + @TestFor(issues = "PY-6426") + public void testAugmentedAssignmentAssignment() { + doTestByText(""" + class A: + def __iadd__(self, other: int) -> str: ... + + a: A = A() + a += 1 + """); + } + + @TestFor(issues = "PY-6426") + public void testAugmentedAssignmentQualified() { + doTestByText(""" + class A: + i: int + + a: A = A() + a.i += 1 + a.i += "s" + """); + } + + @TestFor(issues = "PY-6426") + // test case regarding name resolution + public void testAugmentedAssignmentQualifiedCollision() { + doTestByText(""" + class A: + a: int + + a: A = A() + a.a += 1 + a.a += "s" + """); + } + + @TestFor(issues = "PY-6426") + public void testAugmentedAssignmentGenericAttribute() { + doTestByText(""" + class A[T]: + attr: T + + a: A[int] = A() + a.attr += 1 + a.attr += "s" + + class B: + def __add__(self, other) -> int: ... + + a: A[B] + a.attr += 1 + + class C: + def __iadd__(self, other) -> int: ... + + a: A[C] + a.attr += 1 + """); + } + + @TestFor(issues = "PY-6426") + public void testAugmentedAssignmentDescriptorAttribute() { + doTestByText(""" + class B: + def __add__(self, other: int) -> int: ... + + class C: + def __radd__(self, other: B) -> str: ... + + class Desc: + def __get__(self, instance, owner) -> B: ... + def __set__(self, instance, value: int) -> None: ... + + class A: + attr: Desc + + a: A = A() + a.attr += 1 + a.attr += "s" + a.attr += C() + """); + } } diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3UnresolvedReferencesInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3UnresolvedReferencesInspectionTest.java index adbf1d6bb064..dc1fa8be0079 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3UnresolvedReferencesInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3UnresolvedReferencesInspectionTest.java @@ -1,18 +1,4 @@ -/* - * Copyright 2000-2017 JetBrains s.r.o. - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ +// Copyright 2000-2026 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.intellij.idea.TestFor; @@ -643,4 +629,29 @@ public class Py3UnresolvedReferencesInspectionTest extends PyInspectionTestCase """ ); } + + @TestFor(issues = "PY-80622") + public void testAugAssignmentRaddDefinedButIaddMissingOnTarget() { + doTestByText(""" + class A: pass + class B: + def __radd__(self, other: A) -> str: ... + + a = A() + a += B() # ok + + b = B() + b += A() + """); + } + + @TestFor(issues = "PY-80622") + public void testAugAssignmentIaddNotDefinedOnClass(){ + doTestByText(""" + class A: pass + + a = A() + a += a + """); + } }