PY-6426, PY-80622 type inference for augmented assignments and binary operators

Also fixes PY-87051, PY-36969.
Introduces PyCallSiteOwner to support augmented assignments as call sites
and corrects operator precedence for
binary expressions.

Co-authored-by: Morgan Bartholomew <morgan.bartholomew@jetbrains.com>
Co-authored-by: Mikhail Golubev <mikhail.golubev@jetbrains.com>

GitOrigin-RevId: bdada648f21b2291c2fcd9aae46c78787b452db4
This commit is contained in:
Azim Akhmadjonov
2026-05-14 22:29:49 +00:00
committed by intellij-monorepo-bot
co-authored by Morgan Bartholomew Mikhail Golubev
parent 7b9bf36bd7
commit 176830d757
52 changed files with 1184 additions and 160 deletions
@@ -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);
}
@@ -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);
}
@@ -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<PyAstExpression?> {
return listOf(if (isRightOperator(resolvedCallee)) this.target else this.value)
}
override fun acceptPyVisitor(pyVisitor: PyAstElementVisitor) {
pyVisitor.visitPyAugAssignmentStatement(this)
}
@@ -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<? extends @NotNull PyAstExpression> getArguments(@Nullable PyAstCallable resolvedCallee);
@Unmodifiable
List<PyAstExpression> getArguments(@Nullable PyAstCallable resolvedCallee);
}
@@ -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<PyAstExpression?>
}
@@ -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
@@ -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");// :=
@@ -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<PyAstAssignmentExpression>.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<PyAstAssignmentExpression>.assignedValue as PyExpression?
}
@@ -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?
}
}
@@ -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<PyCallableParameter> myImplicitParameters;
private final @NotNull Map<PyExpression, PyCallableParameter> myMappedParameters;
@@ -115,7 +115,7 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress
private final @NotNull List<PyCallableParameter> myParametersMappedToVariadicKeywordArguments;
private final @NotNull Map<PyExpression, PyCallableParameter> myMappedTupleParameters;
public PyArgumentsMapping(@NotNull PyCallSiteExpression callSiteExpression,
public PyArgumentsMapping(@NotNull PyCallSiteOwner callSiteOwner,
@Nullable PyCallableType callableType,
@NotNull List<PyCallableParameter> implicitParameters,
@NotNull Map<PyExpression, PyCallableParameter> mappedParameters,
@@ -125,7 +125,7 @@ public interface PyCallExpression extends PyAstCallExpression, PyCallSiteExpress
@NotNull List<PyCallableParameter> parametersMappedToVariadicPositionalArguments,
@NotNull List<PyCallableParameter> parametersMappedToVariadicKeywordArguments,
@NotNull Map<PyExpression, PyCallableParameter> 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() {
@@ -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);
}
}
@@ -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<PyExpression>
}
@@ -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<PyExpression, PyCallableParameter> parameters,
@NotNull TypeEvalContext context);
@@ -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<PyType> getReturnType(@NotNull PyCallable callable, @NotNull TypeEvalContext context);
@Nullable
Ref<PyType> getCallType(@NotNull PyFunction function, @NotNull PyCallSiteOwner callSite, @NotNull TypeEvalContext context);
@Nullable
Ref<PyType> getCallType(@NotNull PyFunction function, @NotNull PyCallSiteExpression callSite, @NotNull TypeEvalContext context);
@@ -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.
@@ -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<PyType> 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<PyType> getCallType(@NotNull PyFunction function,
@NotNull PyCallSiteExpression callSite,
@@ -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 <R> assumeType(element: PyTypedElement, type: PyType?, func: (TypeEvalContext?) -> R): R?
@@ -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}''
@@ -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);
}
@@ -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);
}
@@ -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<PyType> getType(TypeEvalContext context, @Nullable PsiElement anchor) {
return myGetType.getType(context);
}
@@ -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<Context> extends PyTypePro
});
}
public Ref<PyType> getCallType(@NotNull PyFunction function, @NotNull PyCallSiteExpression site, @NotNull Context context) {
public Ref<PyType> getCallType(@NotNull PyFunction function, @NotNull PyCallSiteOwner site, @NotNull Context context) {
return null;
}
@@ -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<Context?>() {
return null
}
override fun getCallType(function: PyFunction, callSite: PyCallSiteExpression, context: Context): Ref<PyType?>? {
override fun getCallType(function: PyFunction, callSite: PyCallSiteOwner, context: Context): Ref<PyType?>? {
val functionQName = function.qualifiedName
if (CAST == functionQName || CAST_EXT == functionQName) {
@@ -261,7 +262,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext<Context?>() {
}
if (functionReturningCallSiteAsAType(function)) {
return callSite.getAsClassObjectType(context)
return if (callSite is PyCallSiteExpression) callSite.getAsClassObjectType(context) else null
}
return null
@@ -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<PyType> getClassAttributeType(@NotNull PyTargetExpression attribute) {
private <T extends PyQualifiedExpression & PyReferenceOwner> @Nullable Ref<PyType> 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<AnalyzeCalleeResults> 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()) {
@@ -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<PyType> argumentTypes,
@NotNull List<PyTypeCheckerInspection.AnalyzeCalleeResults> 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<PyType> argumentTypes,
@NotNull List<PyTypeCheckerInspection.AnalyzeCalleeResults> 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();
@@ -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);
}
@@ -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<PyExpression> {
return listOf(if (isRightOperator(resolvedCallee)) target else value!!)
}
}
@@ -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;
@@ -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<PyCallableType>, callSite: PyCallSiteExpression, context: TypeEvalContext): PyType? {
private fun getCallType(types: List<PyCallableType>, 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<op>__ 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 <a href="https://github.com/astral-sh/ty/issues/1154">#1154</a>
* @see <a href="https://github.com/astral-sh/ty/issues/630">#630</a>
*/
@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<PyCallableType>,
callSite: PyCallSiteExpression,
callSite: PyCallSiteOwner,
context: TypeEvalContext,
): List<PyType?> {
val firstCallable = types[0].callable
@@ -597,7 +628,63 @@ object PyCallExpressionHelper {
return types.map { it.getCallType(context, callSite) }
}
private fun resolveOverloadsCallType(types: List<PyCallableType>, 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<PyCallableType>,
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<PyCallableType>, 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<PyArgumentsMapping> {
fun mapArguments(expression: PyCallSiteOwner, resolveContext: PyResolveContext): List<PyArgumentsMapping> {
val context = resolveContext.typeEvalContext
return multiResolveCalleeFunction(expression, resolveContext).map {
mapArguments(expression, it, context)
}
}
private fun multiResolveCalleeFunction(expression: PyCallSiteExpression, resolveContext: PyResolveContext): List<PyCallableType> {
private fun multiResolveCalleeFunction(expression: PyCallSiteOwner, resolveContext: PyResolveContext): List<PyCallableType> {
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<PyCallableParameter>,
callable: PyCallable?,
callSite: PyCallSiteExpression,
callSite: PyCallSiteOwner,
resolveContext: PyResolveContext,
): List<PyCallableParameter> {
val implicitOffset: Int
@@ -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<PyFunctionStub> 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<PyType> typeRef = typeProvider.getCallType(this, callSite, context);
if (typeRef != null) {
@@ -273,7 +273,7 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
@Override
public @Nullable PyType getCallType(@Nullable PyExpression receiver,
@Nullable PyCallSiteExpression callSiteExpression,
@Nullable PyCallSiteOwner callSiteExpression,
@NotNull Map<PyExpression, PyCallableParameter> parameters,
@NotNull TypeEvalContext context) {
@Nullable PyType type = context.getReturnType(this);
@@ -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<PyExpression, PyCallableParameter> parameters,
@NotNull TypeEvalContext context) {
return context.getReturnType(this);
@@ -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);
@@ -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<RatedResolveResult> 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<RatedResolveResult> 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<RatedResolveResult> resolveInlineAndLeftAndRightOperators(@NotNull PyAugAssignmentStatement stmt, @Nullable String name) {
final List<RatedResolveResult> 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;
}
@@ -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<PyExpression, PyCallableParameter> actualParameters,
@NotNull Collection<PyCallableParameter> allParameters,
@Nullable PyExpression receiver,
@NotNull PyCallSiteExpression callsite,
@NotNull PyCallSiteOwner callsite,
@NotNull TypeEvalContext context) {
final var substitutions = PyTypeChecker.unifyGenericCall(receiver, actualParameters, context);
final var substitutionsWithUnresolvedReturnGenerics =
@@ -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));
@@ -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);
}
@@ -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<PyType> getExpectedValueTypeForDunderSet(@NotNull PyTargetExpression targetExpression,
public static @Nullable Ref<PyType> getExpectedValueTypeForDunderSet(@NotNull PyQualifiedExpression targetExpression,
@Nullable PyType attributeType,
@NotNull TypeEvalContext context) {
final PyClassLikeType targetType = as(attributeType, PyClassLikeType.class);
@@ -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);
}
@@ -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<PyExpression> arguments = callSite.getArguments(null);
@@ -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()
@@ -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();
}
@@ -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<PyExpression, PyCallableParameter>,
@@ -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<PyExpression, PyCallableParameter>,
@@ -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
}
@@ -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)
}
@@ -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 <R> assumeType(element: PyTypedElement, type: PyType?, func: (TypeEvalContext?) -> R): R? {
if (!Registry.`is`("python.use.better.control.flow.type.inference")) {
@@ -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<Instruction> result = new LinkedHashSet<>();
final HashMap<PyCallSiteExpression, ConditionalInstruction> pendingTypeGuard = new HashMap<>();
final HashMap<PyCallSiteOwner, ConditionalInstruction> pendingTypeGuard = new HashMap<>();
final Ref<@NotNull Boolean> foundPrefixWrite = Ref.create(false);
final Ref<@NotNull Boolean> foundPrefixCall = Ref.create(false);
iteratePrev(startNum, controlFlow,
@@ -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<List<Pair<Integer, PyParameterInfo>>> 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<List<Pair<Integer, PyParameterInfo>>> myExtraParametersSupplier;
private final @Nullable SmartPsiElementPointer<PyCallSiteExpression> myOriginalCallSiteExpression;
private final @Nullable SmartPsiElementPointer<PyCallSiteOwner> myOriginalCallSiteExpression;
/**
@@ -122,7 +122,7 @@ public final class PyChangeSignatureQuickFix extends LocalQuickFixOnPsiElement {
*/
private PyChangeSignatureQuickFix(@NotNull PyFunction function,
@NotNull Supplier<List<Pair<Integer, PyParameterInfo>>> 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);
@@ -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
@@ -1,9 +1,7 @@
from typing import Any
x = 42
# print('commented')
def func() -> int | Any:
def func() -> int:
return x ** 2
@@ -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<op>__ 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 <a href="https://github.com/astral-sh/ty/issues/1154">#1154</a>
* @see <a href="https://github.com/astral-sh/ty/issues/630">#630</a>
*/
@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<op>__ 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 <a href="https://github.com/astral-sh/ty/issues/1154">#1154</a>
* @see <a href="https://github.com/astral-sh/ty/issues/630">#630</a>
*/
@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<op>__ 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 <a href="https://github.com/astral-sh/ty/issues/1154">#1154</a>
* @see <a href="https://github.com/astral-sh/ty/issues/630">#630</a>
*/
@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);
@@ -5340,4 +5340,120 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
foo(1, "hello", <warning descr="Expected type 'str', got 'int' instead">name=42</warning>)
""");
}
@TestFor(issues = "PY-6426")
public void testAugmentedAssignmentArguments() {
doTestByText("""
class A:
def __iadd__(self, other: int) -> str: ...
a = A()
a += 1
a = A()
a += <warning descr="Expected type 'int', got 'str' instead">"a"</warning>
""");
doTestByText("""
class A:
def __add__(self, other: int) -> str: ...
a = A()
a += 1
a = A()
a += <warning descr="Expected type 'int', got 'str' instead">"a"</warning>
""");
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()
<warning descr="Expected type 'A' for augmented assignment, got 'str' from operation instead">a += 1</warning>
""");
}
@TestFor(issues = "PY-6426")
public void testAugmentedAssignmentQualified() {
doTestByText("""
class A:
i: int
a: A = A()
a.i += 1
a.i += <warning descr="Expected type 'int', got 'str' instead">"s"</warning>
""");
}
@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 += <warning descr="Expected type 'int', got 'str' instead">"s"</warning>
""");
}
@TestFor(issues = "PY-6426")
public void testAugmentedAssignmentGenericAttribute() {
doTestByText("""
class A[T]:
attr: T
a: A[int] = A()
a.attr += 1
a.attr += <warning descr="Expected type 'int', got 'str' instead">"s"</warning>
class B:
def __add__(self, other) -> int: ...
a: A[B]
<warning descr="Expected type 'B' for augmented assignment, got 'int' from operation instead">a.attr += 1</warning>
class C:
def __iadd__(self, other) -> int: ...
a: A[C]
<warning descr="Expected type 'C' for augmented assignment, got 'int' from operation instead">a.attr += 1</warning>
""");
}
@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 += <warning descr="Expected type 'int', got 'str' instead">"s"</warning>
<warning descr="Expected type 'int' (from '__set__'), got 'str' instead">a.attr += C()</warning>
""");
}
}
@@ -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 <warning descr="Class 'B' does not define '__iadd__', so the '+=' operator cannot be used on its instances">+=</warning> A()
""");
}
@TestFor(issues = "PY-80622")
public void testAugAssignmentIaddNotDefinedOnClass(){
doTestByText("""
class A: pass
a = A()
a <warning descr="Class 'A' does not define '__iadd__', so the '+=' operator cannot be used on its instances">+=</warning> a
""");
}
}