From a6ba6900ff90574a408d334fe5a65fae8370689e Mon Sep 17 00:00:00 2001 From: Mikhail Golubev Date: Tue, 30 Sep 2025 16:37:29 +0300 Subject: [PATCH] PY-76886, PY-76860 Make PySelfType extend PyClassType to be able to access its scope class through getPyClass() and infer it for `self` parameters, not having to replace it everywhere with PyClassType. This is similar to how other types with a "backing" PyClass are implemented, e.g. PyTypedDictType delegates to PyClass(dict), PyNarrowedType delegates to PyClass(bool), and so on. It allowed removing special-casing for PySelfType and in many places treat it identically to its scope class (that was previously inferred for `self`), without converting back and forth between the two. The only exceptions are places where we actually need to access the scope class for error messages and code generation, e.g. in PyUnresolvedReferencesInspection and its quick fixes. Namely, extractScopeClassTypeIfNeeded and mutable `matchingScope` were removed from PySelfType. PyTypeChecker.match was simplified to consider PySelfType a subtype of its scope PyClassType, itself compatible only with another PySelfType (if a Self type is expected, only another Self type can be passed). Since PySelf type now propagates through type inference freely, not replaced with PyClassType when unnecessary, there is no need to track where it's used. In particular, situations like the following are properly handled because calling a method returning Self on `self` still returns Self, but calling it on an instance of the scope class returns PyClass for this class. ``` class C: def copy(self): return self def m(self): x: Self = self.copy() # ok x: Self = C().copy() # expected Self@C, actual C ``` Few places in the type checker unification logic were updated to take into account that the type of `self` for generic classes is now just PySelfType, not PyCollectionType with unsubstituted type arguments. Among other benefits, this change allowed removing special-casing for `object.__class__`. Now its type is just `type[Self]` according to the type hints in Typeshed. One unexpected place that required an update was resolveMember and getVariants for PyCallableType. Because the type `self` parameter is now PySelfType, which implements PyTypeParameter, types of all unbound methods are now considered generic types and hence go through the substitution process, that replaces PyFunctionType with PyCallableTypeImpl (to be revised separately). Surprisingly, PyCallableType didn't provide any of the attributes of `types.FunctionType` or `types.UnboundMethodType` as PyFunctionType did. Now the two callable types are in sync. GitOrigin-RevId: eddb56f2fe1f2694554943fc64df288da6d84670 --- .../inspections/PyTypeCheckerInspection.java | 28 +--- .../inspections/PyTypeHintsInspection.kt | 3 +- .../PyUnresolvedReferencesVisitor.java | 15 +- .../src/com/jetbrains/python/psi/PyUtil.java | 75 +++++---- .../python/psi/impl/PyCallExpressionHelper.kt | 17 ++- .../psi/impl/PyReferenceExpressionImpl.java | 21 +-- .../psi/impl/PyTargetExpressionImpl.java | 9 +- .../impl/references/PyQualifiedReference.java | 3 +- .../python/psi/types/PyCallableTypeImpl.java | 14 +- .../psi/types/PyDescriptorTypeUtil.java | 8 +- .../python/psi/types/PyFunctionTypeImpl.java | 75 ++------- .../python/psi/types/PySelfType.java | 54 +++---- .../python/psi/types/PyTypeChecker.java | 142 ++++++++++-------- .../b.py | 2 +- .../PyArgumentListInspection/dictFromKeys.py | 2 +- ...olveWhenAllResultsHaveUnmappedArguments.py | 2 +- ...lveWhenAllResultsHaveUnmappedParameters.py | 2 +- .../overloadsAndImplementationInClass.py | 2 +- .../PyArgumentListInspection/slice.py | 4 +- .../unicodeConstructor.py | 2 +- .../PyArgumentListInspection/xRange.py | 4 +- .../methodSpecialAttributes.py | 2 +- .../com/jetbrains/python/PyTypeTest.java | 4 +- 23 files changed, 218 insertions(+), 272 deletions(-) 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 a6afc81a98f8..48041b97ce6a 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 @@ -127,10 +127,6 @@ public class PyTypeCheckerInspection extends PyInspection { } PyType actual = returnExpr != null ? tryPromotingType(returnExpr, expected) : PyBuiltinCache.getInstance(node).getNoneType(); - if (returnExpr != null) { - expected = trySetSelfMatchingScope(expected, returnExpr, myTypeEvalContext); - } - if (!PyTypeChecker.match(expected, actual, myTypeEvalContext)) { final String expectedName = PythonDocumentationProvider.getVerboseTypeName(expected, myTypeEvalContext); final String actualName = PythonDocumentationProvider.getTypeName(actual, myTypeEvalContext); @@ -280,9 +276,6 @@ public class PyTypeCheckerInspection extends PyInspection { } final PyType actual = tryPromotingType(value, expected); - - expected = trySetSelfMatchingScope(expected, value, myTypeEvalContext); - if (!PyTypeChecker.match(expected, actual, myTypeEvalContext)) { String expectedName = PythonDocumentationProvider.getVerboseTypeName(expected, myTypeEvalContext); String actualName = PythonDocumentationProvider.getTypeName(actual, myTypeEvalContext); @@ -356,25 +349,6 @@ public class PyTypeCheckerInspection extends PyInspection { return myTypeEvalContext.getType(value); } - private @Nullable PyType trySetSelfMatchingScope(@Nullable PyType type, @NotNull PyExpression anchor, @NotNull TypeEvalContext context) { - if (type != null) { - ScopeOwner scopeOwner = ScopeUtil.getScopeOwner(anchor); - if (scopeOwner instanceof PyFunction function) { - scopeOwner = function.getContainingClass(); - } - if (scopeOwner instanceof PyClass pyClass) { - return PyCloningTypeVisitor.clone(type, new PyCloningTypeVisitor(context) { - @Override - public PyType visitPySelfType(@NotNull PySelfType selfType) { - selfType.setMatchingScope(pyClass); - return super.visitPySelfType(selfType); - } - }); - } - } - return type; - } - @Override public void visitPyFunction(@NotNull PyFunction node) { final PyAnnotation annotation = node.getAnnotation(); @@ -505,7 +479,7 @@ public class PyTypeCheckerInspection extends PyInspection { for (Map.Entry entry : regularMappedParameters.entrySet()) { final PyExpression argument = entry.getKey(); final PyCallableParameter parameter = entry.getValue(); - final PyType expected = trySetSelfMatchingScope(parameter.getArgumentType(myTypeEvalContext), argument, myTypeEvalContext); + final PyType expected = parameter.getArgumentType(myTypeEvalContext); final PyType promotedToLiteral = PyLiteralType.Companion.promoteToLiteral(argument, expected, myTypeEvalContext, substitutions); final var actual = promotedToLiteral != null ? promotedToLiteral : myTypeEvalContext.getType(argument); diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt index f134cef17af4..b329f0090088 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt @@ -209,7 +209,8 @@ class PyTypeHintsInspection : PyInspection() { } if (node.referencedName == PyNames.CANONICAL_SELF) { - val typeName = myTypeEvalContext.getType(node)?.name + val refType = myTypeEvalContext.getType(node) + val typeName = (if (refType is PySelfType) refType.scopeClassType else refType)?.name if (typeName != null && typeName != PyNames.CANONICAL_SELF) { registerProblem(node, PyPsiBundle.message("INSP.type.hints.invalid.type.self"), ProblemHighlightType.GENERIC_ERROR, null, ReplaceWithTypeNameQuickFix(typeName)) 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 accaae3432b2..347f6902ea8d 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 @@ -79,8 +79,8 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor final PyExpression qualifier = node.getQualifier(); final String attrName = node.getReferencedName(); if (qualifier != null && attrName != null) { - final PyType type = myTypeEvalContext.getType(qualifier); - if (PySelfType.extractScopeClassTypeIfNeeded(type) instanceof PyClassType classType && + final PyType type = replaceSelfWithItsScopeClass(myTypeEvalContext.getType(qualifier)); + if (type instanceof PyClassType classType && !classType.isAttributeWritable(attrName, myTypeEvalContext)) { final ASTNode nameNode = node.getNameElement(); final PsiElement e = nameNode != null ? nameNode.getPsi() : node; @@ -89,6 +89,10 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor } } + private static @Nullable PyType replaceSelfWithItsScopeClass(@Nullable PyType type) { + return type instanceof PySelfType selfType ? selfType.getScopeClassType() : type; + } + @Override public void visitPyElement(final @NotNull PyElement node) { super.visitPyElement(node); @@ -286,7 +290,7 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor } final PyExpression qualifier = getReferenceQualifier(reference); if (qualifier != null) { - final PyType type = myTypeEvalContext.getType(qualifier); + PyType type = replaceSelfWithItsScopeClass(myTypeEvalContext.getType(qualifier)); if (type != null) { if (ignoreUnresolvedMemberForType(type, reference, refName) || isDeclaredInSlots(type, refName)) { return; @@ -306,6 +310,7 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor ((PyOperatorReference)reference).getReadableOperatorName()); } else { + // TODO use proper type rendering here description = PyPsiBundle.message("INSP.unresolved.refs.unresolved.attribute.for.class", refText, type.getName()); } } @@ -792,8 +797,8 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor @NotNull String refText) { List result = new ArrayList<>(); PsiElement element = reference.getElement(); - if (type instanceof PyClassTypeImpl) { - PyClass cls = ((PyClassType)type).getPyClass(); + if (type instanceof PyClassType pyClassType) { + PyClass cls = pyClassType.getPyClass(); if (!PyBuiltinCache.getInstance(element).isBuiltin(cls)) { if (element.getParent() instanceof PyCallExpression) { result.add(new AddMethodQuickFix(refText, cls.getName(), true)); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java b/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java index 7fe028921b0c..ce77b50f0ec0 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java @@ -182,37 +182,6 @@ public final class PyUtil { // TODO: move to a more proper place? - /** - * Determine the type of a special attribute. Currently supported: {@code __class__} and {@code __dict__}. - * - * @param ref reference to a possible attribute; only qualified references make sense. - * @return type, or null (if type cannot be determined, reference is not to a known attribute, etc.) - */ - public static @Nullable PyType getSpecialAttributeType(@Nullable PyReferenceExpression ref, TypeEvalContext context) { - if (ref != null) { - PyExpression qualifier = ref.getQualifier(); - if (qualifier != null) { - String attr_name = ref.getReferencedName(); - if (PyNames.__CLASS__.equals(attr_name)) { - PyType qualifierType = context.getType(qualifier); - if (qualifierType instanceof PyClassType qClassType) { - return new PyClassTypeImpl(qClassType.getPyClass(), true); // always as class, never instance - } - if (qualifierType instanceof PySelfType selfType) { - return selfType.toClass(); - } - } - else if (PyNames.DUNDER_DICT.equals(attr_name)) { - PyType qualifierType = context.getType(qualifier); - if (qualifierType instanceof PyClassType qClassType && qClassType.isDefinition()) { - return PyBuiltinCache.getInstance(ref).getDictType(); - } - } - } - } - return null; - } - /** * Makes sure that 'thing' is not null; else throws an {@link IncorrectOperationException}. * @@ -1189,6 +1158,50 @@ public final class PyUtil { } } + @ApiStatus.Internal + public static @Nullable PyClassType selectCallableTypeRuntimeClass(@NotNull PyCallableType callableType, + @Nullable PyExpression location, + @NotNull TypeEvalContext context) { + String className; + if (location instanceof PyReferenceExpression re && isBoundMethodReference(callableType, re, context)) { + className = PyNames.TYPES_METHOD_TYPE; + } + else { + className = PyNames.TYPES_FUNCTION_TYPE; + } + PyCallable callable = callableType.getCallable(); + PyClass cls = callable != null ? PyPsiFacade.getInstance(callable.getProject()).createClassByQName(className, callable) : null; + return cls != null ? new PyClassTypeImpl(cls, false) : null; + } + + private static boolean isBoundMethodReference(@NotNull PyCallableType callableType, + @NotNull PyReferenceExpression location, + @NotNull TypeEvalContext context) { + final PyFunction function = as(callableType.getCallable(), PyFunction.class); + final boolean isNonStaticMethod = function != null && function.getContainingClass() != null && function.getModifier() != STATICMETHOD; + if (isNonStaticMethod) { + // In Python 2 unbound methods have __method fake type + if (LanguageLevel.forElement(location).isPython2()) { + return true; + } + final PyExpression qualifier; + if (location.isQualified()) { + qualifier = location.getQualifier(); + } + else { + final PyResolveContext resolveContext = PyResolveContext.defaultContext(context); + qualifier = ContainerUtil.getLastItem(location.followAssignmentsChain(resolveContext).getQualifiers()); + } + if (qualifier != null) { + final PyType qualifierType = context.getType(qualifier); + if (PyTypeUtil.toStream(qualifierType).select(PyClassType.class).anyMatch(it -> !it.isDefinition())) { + return true; + } + } + } + return false; + } + public static final class MethodFlags { private final boolean myIsStaticMethod; 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 9b64296e5f5a..cc0022bbfaad 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 @@ -162,7 +162,14 @@ private fun PyCallExpression.getExplicitResolveResults(resolveContext: PyResolve val result = mutableListOf() for (type in PyTypeUtil.toStream(calleeType)) { - if (type is PyClassType) { + // When invoking cls(), turn type[Self] into Self. + // Otherwise, we will delegate to __init__() of its scope class and return a concrete type class + // as a call result, losing Self. + // See e.g. Py3TypeCheckerInspectionTest.testSelfInClassMethods + if (type is PySelfType) { + result.add(type) + } + else if (type is PyClassType) { val implicitlyInvokedMethods = type .resolveImplicitlyInvokedMethods(this, resolveContext) .forEveryScopeTakeOverloadsOtherwiseImplementations(context) @@ -194,7 +201,7 @@ private fun PyCallExpression.getImplicitResolveResults(resolveContext: PyResolve val qualifier = callee.qualifier if (qualifier == null || !qualifier.canQualifyAnImplicitName()) return listOf() - val qualifierType = PySelfType.extractScopeClassTypeIfNeeded(context.getType(qualifier)) + val qualifierType = context.getType(qualifier) if (PyTypeChecker.isUnknown(qualifierType, context) || qualifierType is PyStructuralType && qualifierType.isInferredFromUsages ) { @@ -407,7 +414,7 @@ private fun PyCallable?.isQualifiedByInstance(qualifier: PyExpression, context: } private fun PyCallable?.isQualifiedByClass(qualifier: PyExpression, context: TypeEvalContext): Boolean { - val qualifierType = PySelfType.extractScopeClassTypeIfNeeded(context.getType(qualifier)) + val qualifierType = context.getType(qualifier) if (qualifierType is PyClassType) { return qualifierType.isDefinition() && belongsToSpecifiedClassHierarchy(qualifierType.pyClass, context) @@ -594,7 +601,7 @@ private fun PyCallExpression.getSuperCallType(context: TypeEvalContext): Maybe

typeOfProperty = getTypeOfProperty(qualifierType, attrName, context); if (typeOfProperty != null) { return typeOfProperty; @@ -272,10 +269,9 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere final PyExpression qualifier = getQualifier(); assert qualifier != null; - final PyType qType = PySelfType.extractScopeClassTypeIfNeeded(context.getType(qualifier)); - if (qType instanceof PyClassType classType) { + if (context.getType(qualifier) instanceof PyClassLikeType classLikeType) { final ResolveResult getattr = ContainerUtil.getFirstItem( - classType.resolveMember(PyNames.GETATTR, qualifier, AccessDirection.READ, PyResolveContext.defaultContext(context))); + classLikeType.resolveMember(PyNames.GETATTR, qualifier, AccessDirection.READ, PyResolveContext.defaultContext(context))); if (getattr != null && getattr.getElement() instanceof PyCallable method) { return context.getReturnType(method); } @@ -323,7 +319,7 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere } private @Nullable Ref getTypeOfProperty(@Nullable PyType qualifierType, @NotNull String name, @NotNull TypeEvalContext context) { - if (PySelfType.extractScopeClassTypeIfNeeded(qualifierType) instanceof PyClassType classType) { + if (qualifierType instanceof PyClassType classType) { final PyClass pyClass = classType.getPyClass(); final Property property = pyClass.findProperty(name, true, context); @@ -360,14 +356,6 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere return null; } - private static @Nullable PyType getDunderClassType(@Nullable PyType qualifierType, @NotNull String attrName) { - if (qualifierType instanceof PyClassType classType && PyNames.__CLASS__.equals(attrName)) { - // PyInstantiableType#toClass() does not work here, as we also need to remove generic parameters - return new PyClassTypeImpl(classType.getPyClass(), true); - } - return null; - } - private @Nullable PyType getTypeFromProviders(@NotNull TypeEvalContext context) { for (PyTypeProvider provider : PyTypeProvider.EP_NAME.getExtensionList()) { try { @@ -489,8 +477,7 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere @NotNull TypeEvalContext context, @NotNull PyReferenceExpression anchor) { if (type instanceof PyFunctionType functionType && context.maySwitchToAST(anchor) && anchor.getQualifier() != null) { - PyType qualifierType = PySelfType.extractScopeClassTypeIfNeeded(context.getType(anchor.getQualifier())); - if (qualifierType instanceof PyClassLikeType classLikeType && classLikeType.isDefinition() && + if (context.getType(anchor.getQualifier()) instanceof PyClassLikeType classLikeType && classLikeType.isDefinition() && functionType.getModifier() != PyAstFunction.Modifier.CLASSMETHOD) { return type; } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java index fe54496f9bef..bb0a13074c6d 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java @@ -348,6 +348,7 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl namesAlready = new HashSet<>(); ctx.put(PyType.CTX_NAMES, namesAlready); 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 d7ee347dda0d..10478c42f051 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 @@ -14,10 +14,7 @@ import com.jetbrains.python.psi.resolve.RatedResolveResult; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.util.Collection; -import java.util.List; -import java.util.Map; -import java.util.Objects; +import java.util.*; public class PyCallableTypeImpl implements PyCallableType { private final @Nullable List myParameters; @@ -83,12 +80,17 @@ public class PyCallableTypeImpl implements PyCallableType { @Nullable PyExpression location, @NotNull AccessDirection direction, @NotNull PyResolveContext resolveContext) { - return null; + PyClassType delegate = PyUtil.selectCallableTypeRuntimeClass(this, location, resolveContext.getTypeEvalContext()); + return delegate != null ? delegate.resolveMember(name, location, direction, resolveContext) : Collections.emptyList(); } + @SuppressWarnings("DuplicatedCode") @Override public Object[] getCompletionVariants(String completionPrefix, PsiElement location, ProcessingContext context) { - return ArrayUtilRt.EMPTY_OBJECT_ARRAY; + TypeEvalContext typeEvalContext = TypeEvalContext.codeCompletion(location.getProject(), location.getContainingFile()); + PyExpression callee = location instanceof PyReferenceExpression re ? re.getQualifier() : null; + PyClassType delegate = PyUtil.selectCallableTypeRuntimeClass(this, callee, typeEvalContext); + return delegate != null ? delegate.getCompletionVariants(completionPrefix, location, context) : ArrayUtilRt.EMPTY_OBJECT_ARRAY; } @Override 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 d6ecfbed4af3..e362ab57f4ee 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 @@ -24,7 +24,7 @@ public final class PyDescriptorTypeUtil { @Nullable PyType attributeType, @NotNull TypeEvalContext context) { if (!expression.isQualified()) return null; - final PyClassLikeType targetType = as(PySelfType.extractScopeClassTypeIfNeeded(attributeType), PyClassLikeType.class); + final PyClassLikeType targetType = as(attributeType, PyClassLikeType.class); if (targetType == null || targetType.isDefinition()) return null; final PyResolveContext resolveContext = PyResolveContext.noProperties(context); @@ -38,7 +38,7 @@ public final class PyDescriptorTypeUtil { public static @Nullable Ref getExpectedValueTypeForDunderSet(@NotNull PyTargetExpression targetExpression, @Nullable PyType attributeType, @NotNull TypeEvalContext context) { - final PyClassLikeType targetType = as(PySelfType.extractScopeClassTypeIfNeeded(attributeType), PyClassLikeType.class); + final PyClassLikeType targetType = as(attributeType, PyClassLikeType.class); if (targetType == null || targetType.isDefinition()) return null; final PyResolveContext resolveContext = PyResolveContext.noProperties(context); @@ -54,8 +54,8 @@ public final class PyDescriptorTypeUtil { @NotNull TypeEvalContext context) { PyExpression qualifier = expression.getQualifier(); if (qualifier != null && attributeType instanceof PyCallableType receiverType) { - PyType qualifierType = PySelfType.extractScopeClassTypeIfNeeded(context.getType(qualifier)); - if (qualifierType instanceof PyClassType classType) { + PyType qualifierType = context.getType(qualifier); + if (qualifierType instanceof PyClassLikeType classType) { PyType instanceArgumentType; PyType instanceTypeArgument; final var noneType = PyBuiltinCache.getInstance(expression).getNoneType(); 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 fcbeafa6dc38..4de0637894d4 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 @@ -5,7 +5,6 @@ import com.intellij.psi.PsiElement; import com.intellij.util.ArrayUtilRt; import com.intellij.util.ProcessingContext; import com.intellij.util.containers.ContainerUtil; -import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.RatedResolveResult; @@ -16,9 +15,6 @@ import java.util.ArrayList; import java.util.Collections; import java.util.List; -import static com.jetbrains.python.ast.PyAstFunction.Modifier.STATICMETHOD; -import static com.jetbrains.python.psi.PyUtil.as; - /** * Type of a particular function that is represented as a {@link PyCallable} in the PSI tree. */ @@ -71,70 +67,21 @@ public class PyFunctionTypeImpl implements PyFunctionType { } @Override - public List resolveMember(@NotNull String name, - @Nullable PyExpression location, - @NotNull AccessDirection direction, - @NotNull PyResolveContext resolveContext) { - final PyClassType delegate = selectCallableType(location, resolveContext.getTypeEvalContext()); - if (delegate == null) { - return Collections.emptyList(); - } - return delegate.resolveMember(name, location, direction, resolveContext); + public @Nullable List resolveMember(@NotNull String name, + @Nullable PyExpression location, + @NotNull AccessDirection direction, + @NotNull PyResolveContext resolveContext) { + PyClassType delegate = PyUtil.selectCallableTypeRuntimeClass(this, location, resolveContext.getTypeEvalContext()); + return delegate != null ? delegate.resolveMember(name, location, direction, resolveContext) : Collections.emptyList(); } + @SuppressWarnings("DuplicatedCode") @Override public Object[] getCompletionVariants(String completionPrefix, PsiElement location, ProcessingContext context) { - final TypeEvalContext typeEvalContext = TypeEvalContext.codeCompletion(location.getProject(), location.getContainingFile()); - final PyClassType delegate; - if (location instanceof PyReferenceExpression) { - delegate = selectCallableType(((PyReferenceExpression)location).getQualifier(), typeEvalContext); - } - else { - final PyClass cls = PyPsiFacade.getInstance(myCallable.getProject()).createClassByQName(PyNames.TYPES_FUNCTION_TYPE, myCallable); - delegate = cls != null ? new PyClassTypeImpl(cls, false) : null; - } - if (delegate == null) { - return ArrayUtilRt.EMPTY_OBJECT_ARRAY; - } - return delegate.getCompletionVariants(completionPrefix, location, context); - } - - private @Nullable PyClassType selectCallableType(@Nullable PyExpression location, @NotNull TypeEvalContext context) { - final String className; - if (location instanceof PyReferenceExpression && isBoundMethodReference((PyReferenceExpression)location, context)) { - className = PyNames.TYPES_METHOD_TYPE; - } - else { - className = PyNames.TYPES_FUNCTION_TYPE; - } - final PyClass cls = PyPsiFacade.getInstance(myCallable.getProject()).createClassByQName(className, myCallable); - return cls != null ? new PyClassTypeImpl(cls, false) : null; - } - - private boolean isBoundMethodReference(@NotNull PyReferenceExpression location, @NotNull TypeEvalContext context) { - final PyFunction function = as(getCallable(), PyFunction.class); - final boolean isNonStaticMethod = function != null && function.getContainingClass() != null && function.getModifier() != STATICMETHOD; - if (isNonStaticMethod) { - // In Python 2 unbound methods have __method fake type - if (LanguageLevel.forElement(location).isPython2()) { - return true; - } - final PyExpression qualifier; - if (location.isQualified()) { - qualifier = location.getQualifier(); - } - else { - final PyResolveContext resolveContext = PyResolveContext.defaultContext(context); - qualifier = ContainerUtil.getLastItem(location.followAssignmentsChain(resolveContext).getQualifiers()); - } - if (qualifier != null) { - final PyType qualifierType = PySelfType.extractScopeClassTypeIfNeeded(context.getType(qualifier)); - if (PyTypeUtil.toStream(qualifierType).select(PyClassType.class).anyMatch(it -> !it.isDefinition())) { - return true; - } - } - } - return false; + TypeEvalContext typeEvalContext = TypeEvalContext.codeCompletion(location.getProject(), location.getContainingFile()); + PyExpression callee = location instanceof PyReferenceExpression re ? re.getQualifier() : null; + PyClassType delegate = PyUtil.selectCallableTypeRuntimeClass(this, callee, typeEvalContext); + return delegate != null ? delegate.getCompletionVariants(completionPrefix, location, context) : ArrayUtilRt.EMPTY_OBJECT_ARRAY; } @Override 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 34274327c280..8898fafbef33 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 @@ -1,5 +1,6 @@ package com.jetbrains.python.psi.types; +import com.intellij.openapi.util.Key; import com.intellij.psi.PsiElement; import com.intellij.util.ProcessingContext; import com.intellij.util.Processor; @@ -13,9 +14,8 @@ import java.util.List; import java.util.Objects; import java.util.Set; -public final class PySelfType implements PyTypeParameterType, PyClassLikeType { +public final class PySelfType implements PyTypeParameterType, PyClassType { private final @NotNull PyClassType myScopeClassType; - private @Nullable PyClass matchingScope = null; public PySelfType(@NotNull PyClassType scopeClassType) { myScopeClassType = scopeClassType; @@ -50,13 +50,6 @@ public final class PySelfType implements PyTypeParameterType, PyClassLikeType { return myScopeClassType; } - public static @Nullable PyType extractScopeClassTypeIfNeeded(@Nullable PyType type) { - if (type instanceof PySelfType selfType) { - return selfType.myScopeClassType; - } - return type; - } - /** * @return true if type[Self], false otherwise */ @@ -75,29 +68,6 @@ public final class PySelfType implements PyTypeParameterType, PyClassLikeType { return new PySelfType(myScopeClassType.toClass()); } - /** - * The presense of this field indicates that we perform type-check inside the class that this type belongs to. - * It means that some special rules should be applied, e.g.: - *

-   * {@code
-   *   class Example:
-   *      def returns_instance(self) -> Self:
-   *          return Example() # Error
-   *
-   *       def returns_self(self) -> Self:
-   *          return self # OK
-   * }
-   * 
- */ - @Nullable - public PyClass getMatchingScope() { - return matchingScope; - } - - public void setMatchingScope(@Nullable PyClass matchingScopeClass) { - this.matchingScope = matchingScopeClass; - } - @Override public boolean equals(Object o) { if (this == o) return true; @@ -209,6 +179,26 @@ public final class PySelfType implements PyTypeParameterType, PyClassLikeType { return myScopeClassType.getImplicitOffset(); } + @Override + public @NotNull PyClass getPyClass() { + return myScopeClassType.getPyClass(); + } + + @Override + public @Nullable T getUserData(@NotNull Key key) { + return null; + } + + @Override + public void putUserData(@NotNull Key key, @Nullable T value) { + + } + + @Override + public boolean isAttributeWritable(@NotNull String name, @NotNull TypeEvalContext context) { + return myScopeClassType.isAttributeWritable(name, context); + } + @Override public T acceptTypeVisitor(@NotNull PyTypeVisitor visitor) { if (visitor instanceof PyTypeVisitorExt visitorExt) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java index c5daf26bab7c..c9cf1c679450 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -143,15 +143,11 @@ public final class PyTypeChecker { } if (expected instanceof PySelfType selfType) { - return Optional.of(matchSelf(selfType, actual, context)); + return Optional.of(match(selfType, actual, context)); } if (actual instanceof PySelfType selfType && context.reversedSubstitutions) { - return Optional.of(matchSelf(selfType, expected, context)); - } - - if (actual instanceof PySelfType selfType) { - return Optional.of(matchSelf(expected, selfType, context)); + return Optional.of(match(selfType, expected, context)); } if (expected instanceof PyParamSpecType paramSpecType) { @@ -362,59 +358,15 @@ public final class PyTypeChecker { } } - private static boolean matchSelf(@Nullable PyType expected, @Nullable PyType actual, @NotNull MatchContext context) { - if (expected == null || actual == null) { - return true; + private static boolean match(@NotNull PySelfType expected, @Nullable PyType actual, @NotNull MatchContext context) { + if (actual == null) return true; + PyType substitution = context.mySubstitutions.getQualifierType(); + if (substitution != null && !(substitution instanceof PySelfType)) { + return match(substitution, actual, context).orElse(false); } - if (expected instanceof PySelfType expectedSelf) { - var substitutedExpected = substitute(expectedSelf, context.mySubstitutions, context.context); - - if (substitutedExpected instanceof PySelfType expectedSelfAfterSubst) { // no qualifier - // The presence of the matching scope indicates that we are within the function Self belongs to - if (expectedSelf.getMatchingScope() != null) { - if (actual instanceof PySelfType actualSelf) { - /* - x: Self = cls # Error - y: Self = cls() # OK - */ - return expectedSelf.getMatchingScope().equals(actualSelf.getScopeClassType().getPyClass()) && - expectedSelf.isDefinition() == actualSelf.isDefinition(); - } - /* - Other cases should fail, especially important for: - class MyClass: - def method(self) -> Self: - return MyClass() <- Error, forbidden - */ - return false; - } - return match(expectedSelfAfterSubst.getScopeClassType(), actual, context).orElse(false); - } - else { // qualifier is present, match the qualifier and the passed type - /* - class Z: - def method(self): ... - - class A: - def method(self): - Z.method(self) # E, Self@A - Z.method(A()) # E, A - */ - if (substitutedExpected instanceof PyClassType expectedClassType && - PySelfType.extractScopeClassTypeIfNeeded(actual) instanceof PyClassType actualClassType) { - return match(expectedClassType, actualClassType, context).orElse(false); - } - } - return true; - } - - if (actual instanceof PySelfType actualSelf) { - if (expected instanceof PyClassType expectedClassType) { - return match(expectedClassType, actualSelf.getScopeClassType(), context).orElse(false); - } - } - // fallback - return match(expected, actual, context).orElse(true); + if (!(actual instanceof PySelfType actualSelf)) return false; + return expected.isDefinition() == actualSelf.isDefinition() && + match(expected.getScopeClassType(), actualSelf.getScopeClassType(), context).orElse(false); } private static @Nullable PyType toClass(@Nullable PyType type) { @@ -664,7 +616,15 @@ public final class PyTypeChecker { private static boolean matchProtocols(@NotNull PyClassType expected, @NotNull PyClassType actual, @NotNull MatchContext matchContext) { GenericSubstitutions substitutions = collectTypeSubstitutions(actual, matchContext.context); - MatchContext protocolContext = new MatchContext(matchContext.context, new GenericSubstitutions(), matchContext.reversedSubstitutions); + + // See https://typing.python.org/en/latest/spec/generics.html#use-in-protocols + // > If a protocol uses Self in methods or attribute annotations, then a class Foo is assignable to the protocol + // if its corresponding methods and attribute annotations use either Self or Foo or any of Foo’s subclasses. + // + // It should be equivalent to replacing Self in the protocol with the Foo class we're matching it with. + GenericSubstitutions protocolSubstitutions = new GenericSubstitutions(); + protocolSubstitutions.qualifierType = actual.toInstance(); + MatchContext protocolContext = new MatchContext(matchContext.context, protocolSubstitutions, matchContext.reversedSubstitutions); for (kotlin.Pair> pair : PyProtocolsKt.inspectProtocolSubclass(expected, actual, matchContext.context)) { final PyTypeMember protocolMember = pair.getFirst(); @@ -997,9 +957,8 @@ public final class PyTypeChecker { } private static @Nullable PyType getActualReturnType(@NotNull PyCallableType actual, @NotNull TypeEvalContext context) { - PyCallable callable = actual.getCallable(); - if (callable instanceof PyFunction) { - return getReturnTypeToAnalyzeAsCallType((PyFunction)callable, context); + if (actual instanceof PyFunctionType functionType && functionType.getCallable() instanceof PyFunction function) { + return getReturnTypeToAnalyzeAsCallType(function, context); } return actual.getReturnType(context); } @@ -1037,8 +996,27 @@ public final class PyTypeChecker { } // TODO Make it a part of PyClassType interface + + /** + * Extract all type substitutions from a generic class instance. For instance, for the type of `x` + * + *
{@code
+   * class Base[T1, T2]: ...
+   * class Sub[T3](Base[int, T3]): ...
+   * x: Sub[str]
+   * }
+ *

+ * we should extract + * + *

{@code
+   * T1@Base -> int
+   * T2@Base -> Sub@T3
+   * Sub@T3 -> str
+   * }
+ */ private static @NotNull GenericSubstitutions collectTypeSubstitutions(@NotNull PyClassType classType, @NotNull TypeEvalContext context) { GenericSubstitutions result = new GenericSubstitutions(); + // Collect substitutions from ancestor classes. In the example above these are: T1@Base -> int and T2@Base -> Sub@T3 for (PyTypeProvider provider : PyTypeProvider.EP_NAME.getExtensionList()) { Map substitutionsFromClassDefinition = provider.getGenericSubstitutions(classType.getPyClass(), context); for (Map.Entry entry : substitutionsFromClassDefinition.entrySet()) { @@ -1054,11 +1032,30 @@ public final class PyTypeChecker { result.paramSpecs.put(specType, paramSpecType); } } + // Collect own type parameters. In the example above these are: Sub@T3 -> str PyCollectionType genericDefinitionType = as(provider.getGenericType(classType.getPyClass(), context), PyCollectionType.class); // TODO Re-use PyTypeParameterMapping, at the moment C[*Ts] <- C leads to *Ts being mapped to *tuple[], which breaks inference later on if (genericDefinitionType != null) { List definitionTypeParameters = genericDefinitionType.getElementTypes(); - if (!(classType instanceof PyCollectionType genericType)) { + + // Inside method bodies, where Self type can appear, we map class' own type parameters to themseleves, + // not to consider them unbound. For instance, here + // + // class B[T]:... + // class C[T2](B[T2]): + // attr: T + // def m(self): + // self.attr + // + // when collecting substitutions for `self` in `m`, we will map T@B -> T2@C, but T2@C -> T2@C. + // + // See Py3TypeTest.testApplyingSuperSubstitutionToBoundedGenericClass + // and Py3TypeTest#testApplyingSuperSubstitutionToGenericClass + if (classType instanceof PySelfType selfType) { + mapTypeParametersToSubstitutions(result, definitionTypeParameters, definitionTypeParameters, + Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY); + } + else if (!(classType instanceof PyCollectionType genericType)) { for (PyType typeParameter : definitionTypeParameters) { if (typeParameter instanceof PyTypeVarTupleType typeVarTupleType) { result.typeVarTuples.put(typeVarTupleType, null); @@ -1159,7 +1156,9 @@ public final class PyTypeChecker { } public static boolean isUnknown(@Nullable PyType type, boolean genericsAreUnknown, @NotNull TypeEvalContext context) { - if (type == null || (genericsAreUnknown && type instanceof PyTypeParameterType)) { + // TODO Don't consider other type parameters unknown (PY-85653) + // Since Self is always bound, don't consider it unknown, e.g. in Py3TypeCheckerInspectionTest.testSelfAssignedToOtherTypeBad + if (type == null || (genericsAreUnknown && type instanceof PyTypeParameterType && !(type instanceof PySelfType))) { return true; } if (type instanceof PyUnionType union) { @@ -1385,6 +1384,17 @@ public final class PyTypeChecker { if (qualifierType == null) { return selfType; } + // TODO change unification for calls on union types + // so that in the following + // + // class A: + // def a_method(self) -> Self: + // ... + // x: A | B + // + // A.a_method # type: Callable[[A], A] + // B wasn't considered as the receiver type in the first place, instead of filtering it out during substitution + // (see PyTypingTest.testMatchSelfUnionType) return PyTypeUtil.toStream(qualifierType) .map(qType -> { if (qType instanceof PyInstantiableType instantiableType) { @@ -1459,6 +1469,10 @@ public final class PyTypeChecker { parametersSubs instanceof PyParamSpecType paramSpec ? List.of(PyCallableParameterImpl.nonPsi(paramSpec)) : null, clone(callableType.getReturnType(context)) + , + callableType.getCallable(), + callableType.getModifier(), + callableType.getImplicitOffset() ); } diff --git a/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py index 97dfc5be4c61..ac289b59ab0c 100644 --- a/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py +++ b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py @@ -1,3 +1,3 @@ import c -c.A().foo() \ No newline at end of file +c.A().foo() \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/dictFromKeys.py b/python/testData/inspections/PyArgumentListInspection/dictFromKeys.py index fb2a7d6f1186..d670bf587933 100644 --- a/python/testData/inspections/PyArgumentListInspection/dictFromKeys.py +++ b/python/testData/inspections/PyArgumentListInspection/dictFromKeys.py @@ -1,2 +1,2 @@ -print(dict.fromkeys()) +print(dict.fromkeys()) print(dict.fromkeys(['foo', 'bar'])) diff --git a/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedArguments.py b/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedArguments.py index 7a3fbdeb4b6f..29c0e9398d9f 100644 --- a/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedArguments.py +++ b/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedArguments.py @@ -15,4 +15,4 @@ def f(): pass -f().foo(1, 2, 3) +f().foo(1, 2, 3) diff --git a/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedParameters.py b/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedParameters.py index 6b340be17531..44b00c6ec0e9 100644 --- a/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedParameters.py +++ b/python/testData/inspections/PyArgumentListInspection/multiResolveWhenAllResultsHaveUnmappedParameters.py @@ -15,4 +15,4 @@ def f(): pass -f().foo() +f().foo() diff --git a/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py b/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py index 4ba41c7f33ce..c9a744ba0d08 100644 --- a/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py +++ b/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py @@ -18,4 +18,4 @@ class A: return None -A().foo() \ No newline at end of file +A().foo() \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/slice.py b/python/testData/inspections/PyArgumentListInspection/slice.py index c582cd5b7c7b..bdd400d5af91 100644 --- a/python/testData/inspections/PyArgumentListInspection/slice.py +++ b/python/testData/inspections/PyArgumentListInspection/slice.py @@ -1,5 +1,5 @@ -print(slice()) +print(slice()) print(slice(1)) print(slice(1, 2)) print(slice(1, 2, 3)) -print(slice(1, 2, 3, 4)) +print(slice(1, 2, 3, 4)) diff --git a/python/testData/inspections/PyArgumentListInspection/unicodeConstructor.py b/python/testData/inspections/PyArgumentListInspection/unicodeConstructor.py index 76f3dff617c4..2ff77b51def4 100644 --- a/python/testData/inspections/PyArgumentListInspection/unicodeConstructor.py +++ b/python/testData/inspections/PyArgumentListInspection/unicodeConstructor.py @@ -2,4 +2,4 @@ print(unicode()) print(unicode('')) print(unicode('', 'utf-8')) print(unicode('', 'utf-8', 'ignore')) -print(unicode('', 'utf-8', 'ignore', foo)) +print(unicode('', 'utf-8', 'ignore', foo)) diff --git a/python/testData/inspections/PyArgumentListInspection/xRange.py b/python/testData/inspections/PyArgumentListInspection/xRange.py index 5b007f2e4e70..02a076652d7e 100644 --- a/python/testData/inspections/PyArgumentListInspection/xRange.py +++ b/python/testData/inspections/PyArgumentListInspection/xRange.py @@ -1,5 +1,5 @@ -print(xrange()) +print(xrange()) print(xrange(1)) print(xrange(1, 2)) print(xrange(1, 2, 3)) -print(xrange(1, 2, 3, 4)) +print(xrange(1, 2, 3, 4)) diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/methodSpecialAttributes.py b/python/testData/inspections/PyUnresolvedReferencesInspection/methodSpecialAttributes.py index f8706ab8199d..a5a6ca426a6c 100644 --- a/python/testData/inspections/PyUnresolvedReferencesInspection/methodSpecialAttributes.py +++ b/python/testData/inspections/PyUnresolvedReferencesInspection/methodSpecialAttributes.py @@ -9,7 +9,7 @@ class MyClass(object): # Unbound methods are still treated as __method in Python 2 MyClass.method.__func__ -MyClass.method.__defaults__ +MyClass.method.__defaults__ # Bound method with qualifier inst = MyClass() diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 9ba077cf573c..f27dabddef7a 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -3037,7 +3037,7 @@ public class PyTypeTest extends PyTestCase { public void testReplaceDefinitionInMethod() { doTest("Type[Derived]", """ - class Base: + class Base(object): def cls(self): return self.__class__ class Derived(Base): @@ -3046,7 +3046,7 @@ public class PyTypeTest extends PyTestCase { doTest("Type[Derived]", """ - class Base: + class Base(object): def cls(self): return self.__class__ class Derived(Base):