From 044c14051680ac1f2d53c4b8b38ff4587c0bdb99 Mon Sep 17 00:00:00 2001 From: "evgeny.bovykin" Date: Wed, 20 Aug 2025 13:15:39 +0200 Subject: [PATCH] PY-84464 Support `@property` decorator when matching a protocol and its implementation GitOrigin-RevId: fcbeeb57323336c7d921edad373ea595ef687d6b --- .../python/psi/impl/PyTypeProvider.java | 12 +- .../jetbrains/python/psi/types/PyType.java | 40 +--- .../python/psi/types/PyTypeMember.kt | 38 ++++ .../python/psi/types/PyTypeProviderBase.java | 10 +- .../python/psi/types/PyTypedResolveResult.kt | 14 -- .../resources/messages/PyPsiBundle.properties | 1 + .../com/jetbrains/python/PyCustomType.java | 20 ++ .../stdlib/PyDataclassTypeProvider.kt | 25 ++- .../stdlib/PyNamedTupleTypeProvider.kt | 4 +- .../python/codeInsight/typing/PyProtocols.kt | 105 +++++----- .../inspections/PyProtocolInspection.kt | 21 +- .../python/psi/impl/PyCallExpressionHelper.kt | 27 ++- .../python/psi/impl/PyClassPatternImpl.java | 5 +- .../python/psi/types/PyClassTypeImpl.java | 116 ++++++++++- .../python/psi/types/PyTypeChecker.java | 124 +++++++----- .../Py3TypeCheckerInspectionTest.java | 189 +++++++++++++++++- .../inspections/PyProtocolInspectionTest.java | 76 +++++++ 17 files changed, 629 insertions(+), 198 deletions(-) create mode 100644 python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeMember.kt delete mode 100644 python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypedResolveResult.kt diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java b/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java index 2cbd3d9551b9..36cfea0fd578 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/impl/PyTypeProvider.java @@ -8,7 +8,7 @@ import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.types.PyCallableType; import com.jetbrains.python.psi.types.PyType; -import com.jetbrains.python.psi.types.PyTypedResolveResult; +import com.jetbrains.python.psi.types.PyTypeMember; import com.jetbrains.python.psi.types.TypeEvalContext; import org.jetbrains.annotations.ApiStatus; import org.jetbrains.annotations.NotNull; @@ -78,9 +78,9 @@ public interface PyTypeProvider { @NotNull TypeEvalContext context); @ApiStatus.Experimental - @Nullable List<@NotNull PyTypedResolveResult> getMemberTypes(@NotNull PyType type, - @NotNull String name, - @Nullable PyExpression location, - @NotNull AccessDirection direction, - @NotNull PyResolveContext context); + @Nullable List<@NotNull PyTypeMember> getMemberTypes(@NotNull PyType type, + @NotNull String name, + @Nullable PyExpression location, + @NotNull AccessDirection direction, + @NotNull PyResolveContext context); } diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyType.java b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyType.java index b2fa1897f6b7..f40288b8843d 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyType.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyType.java @@ -5,12 +5,9 @@ import com.intellij.openapi.util.Key; import com.intellij.openapi.util.NlsSafe; import com.intellij.psi.PsiElement; import com.intellij.util.ProcessingContext; -import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.psi.AccessDirection; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.PyQualifiedNameOwner; -import com.jetbrains.python.psi.PyTypedElement; -import com.jetbrains.python.psi.impl.PyTypeProvider; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.RatedResolveResult; import org.jetbrains.annotations.ApiStatus; @@ -51,34 +48,19 @@ public interface PyType { final @NotNull AccessDirection direction, final @NotNull PyResolveContext resolveContext); + @ApiStatus.Experimental - @Nullable - default List<@NotNull PyTypedResolveResult> getMemberTypes(@NotNull String name, - final @Nullable PyExpression location, - final @NotNull AccessDirection direction, - final @NotNull PyResolveContext context) { - for (PyTypeProvider typeProvider : PyTypeProvider.EP_NAME.getExtensionList()) { - List types = typeProvider.getMemberTypes(this, name, location, direction, context); - if (types != null) { - return types; - } - } + default @NotNull List<@NotNull PyTypeMember> getAllMembers(final @NotNull PyResolveContext resolveContext) { + return List.of(); + } - List results = resolveMember(name, location, direction, context); - if (results == null) { - return null; - } - - return ContainerUtil.map(results, result -> { - PsiElement element = result.getElement(); - if (element instanceof PyTypedElement typedElement) { - return new PyTypedResolveResult(typedElement, - context.getTypeEvalContext().getType(typedElement)); - } - else { - return new PyTypedResolveResult(element, null); - } - }); + /** + * Returns a list of members with a given name + * There can be several members with the same name (for example, methods with @overload) + */ + @ApiStatus.Experimental + default @NotNull List<@NotNull PyTypeMember> findMember(@NotNull String name, final @NotNull PyResolveContext resolveContext) { + return List.of(); } /** diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeMember.kt b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeMember.kt new file mode 100644 index 000000000000..7b3a3d0495e1 --- /dev/null +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeMember.kt @@ -0,0 +1,38 @@ +// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.jetbrains.python.psi.types + +import com.intellij.psi.PsiElement +import com.intellij.psi.PsiNamedElement +import com.jetbrains.python.psi.Property +import com.jetbrains.python.psi.resolve.RatedResolveResult +import org.jetbrains.annotations.ApiStatus + +/** + * Represents a high-level member of a class + * It can be a simple attribute, a method, a property, a dataclass attribute and so on + * + * For a property, one member represents its getter, setter and deleter + */ +@ApiStatus.Experimental +class PyTypeMember @JvmOverloads constructor( + val mainElement: PsiElement?, + val type: PyType?, + val isClassVar: Boolean = false, + val getter: PsiElement? = mainElement, + val setter: PsiElement? = mainElement, + val deleter: PsiElement? = mainElement, +) : RatedResolveResult(0, mainElement) { + + constructor(property: Property, type: PyType?) : this( + property.getter.value(), + type, + getter = property.getter.value(), + setter = property.setter.valueOrNull(), + deleter = property.deleter.valueOrNull(), + ) + + val isWritable: Boolean get() = setter != null + val isDeletable: Boolean get() = deleter != null + + val name: String? get() = if (mainElement is PsiNamedElement) mainElement.name else null +} \ No newline at end of file diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java index 2cf33250f629..bb809c0d4135 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java @@ -84,11 +84,11 @@ public class PyTypeProviderBase implements PyTypeProvider { } @Override - public @Nullable List<@NotNull PyTypedResolveResult> getMemberTypes(@NotNull PyType type, - @NotNull String name, - @Nullable PyExpression location, - @NotNull AccessDirection direction, - @NotNull PyResolveContext context) { + public @Nullable List<@NotNull PyTypeMember> getMemberTypes(@NotNull PyType type, + @NotNull String name, + @Nullable PyExpression location, + @NotNull AccessDirection direction, + @NotNull PyResolveContext context) { return null; } diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypedResolveResult.kt b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypedResolveResult.kt deleted file mode 100644 index 3b36c75ff8af..000000000000 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyTypedResolveResult.kt +++ /dev/null @@ -1,14 +0,0 @@ -// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. -package com.jetbrains.python.psi.types - -import com.intellij.psi.PsiElement -import com.intellij.psi.ResolveResult - -class PyTypedResolveResult(private val el: PsiElement?, val type: PyType?) : ResolveResult { - override fun getElement(): PsiElement? { - return el - } - override fun isValidResult(): Boolean { - return element != null - } -} \ No newline at end of file diff --git a/python/python-psi-impl/resources/messages/PyPsiBundle.properties b/python/python-psi-impl/resources/messages/PyPsiBundle.properties index f70b4ce83054..02a1332b96c0 100644 --- a/python/python-psi-impl/resources/messages/PyPsiBundle.properties +++ b/python/python-psi-impl/resources/messages/PyPsiBundle.properties @@ -1100,6 +1100,7 @@ INSP.protocol.only.runtime.checkable.protocols.can.be.used.with.instance.class.c INSP.protocol.newtype.cannot.be.used.with.protocol.classes=NewType cannot be used with protocol classes INSP.protocol.element.type.incompatible.with.protocol=Type of ''{0}'' is incompatible with ''{1}'' INSP.protocol.cannot.instantiate.protocol.class=Cannot instantiate protocol class ''{0}'' +INSP.protocol.element.type.not.writable=''{0}'' is writable in protocol ''{1}'' # PyShadowingBuiltinsInspection INSP.NAME.shadowing.builtins=Shadowing built-in names diff --git a/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java b/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java index da2b7cfb9f1f..a36dc705a2b4 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java +++ b/python/python-psi-impl/src/com/jetbrains/python/PyCustomType.java @@ -301,6 +301,26 @@ public final class PyCustomType implements PyClassLikeType { return result; } + @Override + public @NotNull List<@NotNull PyTypeMember> getAllMembers(@NotNull PyResolveContext resolveContext) { + List result = new ArrayList<>(); + for (PyClassLikeType type : myTypesToMimic) { + result.addAll(type.getAllMembers(resolveContext)); + } + return result; + } + + @Override + public @NotNull List<@NotNull PyTypeMember> findMember(@NotNull String name, @NotNull PyResolveContext resolveContext) { + for (PyClassLikeType type : myTypesToMimic) { + List members = type.findMember(name, resolveContext); + if (!members.isEmpty()) { + return members; + } + } + return List.of(); + } + /** * Predicate that filters completion using {@link #myFilter} */ diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt index 6339d5900ce0..910f7568dfe1 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt @@ -77,15 +77,13 @@ class PyDataclassTypeProvider : PyTypeProviderBase() { return null } - override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List? { + override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List? { if (type !is PyClassType) { return null } + val dataclassParameters = parseDataclassParameters(type.pyClass, context.typeEvalContext) ?: return null if (PyNames.HASH == name) { // See `unsafe_hash` section here https://docs.python.org/3/library/dataclasses.html - val dataclassParameters = parseDataclassParameters(type.pyClass, context.typeEvalContext) - if (dataclassParameters == null) return null - if (dataclassParameters.unsafeHash) { return null } @@ -102,7 +100,24 @@ class PyDataclassTypeProvider : PyTypeProviderBase() { if (resolvedMembers?.isNotEmpty() == true) { return null } - return listOf(PyTypedResolveResult(null, PyBuiltinCache.getInstance(type.pyClass).noneType)) + return listOf(PyTypeMember(null, PyBuiltinCache.getInstance(type.pyClass).noneType)) + } + else { + if (dataclassParameters.frozen) { + val resolvedMembers = type.resolveMember(name, location, direction, context, false) + if (resolvedMembers?.isNotEmpty() == true) { + return resolvedMembers.map { + val element = it.element + val type = if (element is PyTypedElement) { + context.typeEvalContext.getType(element) + } + else { + null + } + PyTypeMember(element, type, false, element, null, null) + } + } + } } return null diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleTypeProvider.kt index 77233cf0e512..9a1a1d93c0af 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleTypeProvider.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleTypeProvider.kt @@ -51,10 +51,10 @@ class PyNamedTupleTypeProvider : PyTypeProviderBase() { return if (type is PyNamedTupleType) Ref.create(type) else null } - override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List? { + override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List? { if (type !is PyNamedTupleType) return null type.fields[name]?.let { - return listOf(PyTypedResolveResult(null, it.type)) + return listOf(PyTypeMember(null, it.type)) } return null } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt index ba7e9296782f..b2043645ae48 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt @@ -6,7 +6,7 @@ import com.jetbrains.python.PyNames import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.PROTOCOL import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.PROTOCOL_EXT import com.jetbrains.python.psi.* -import com.jetbrains.python.psi.impl.getImplicitlyInvokedMethodTypes +import com.jetbrains.python.psi.impl.getImplicitlyInvokedMethod import com.jetbrains.python.psi.impl.resolveImplicitlyInvokedMethods import com.jetbrains.python.psi.resolve.PyResolveContext import com.jetbrains.python.psi.types.* @@ -23,68 +23,65 @@ fun matchingProtocolDefinitions(expected: PyType?, actual: PyType?, context: Typ isProtocol(expected, context) && isProtocol(actual, context) -typealias ProtocolAndSubclassElements = Pair> +typealias ProtocolAndSubclassElements = Pair> fun inspectProtocolSubclass(protocol: PyClassType, subclass: PyClassType, context: TypeEvalContext): List { val resolveContext = PyResolveContext.defaultContext(context) - val result = mutableListOf>>() + val result = mutableListOf>>() - protocol.toInstance().visitMembers( - { e -> - if (e is PyTypedElement) { - if (e is PyPossibleClassMember) { - val cls = e.containingClass - if (cls != null && !isProtocol(cls, context)) { - return@visitMembers true - } - } - if (e is PyTypeParameter) { - return@visitMembers true - } + val protocolMembers = protocol.toInstance().getAllMembers(resolveContext) + val superClassesMembers = protocol.toInstance().getSuperClassTypes(context) + .filter { isProtocol(it, context) } + .flatMap { it.toInstance().getAllMembers(resolveContext).asIterable() } + protocolMembers.addAll(superClassesMembers) - if (e.contextOfType()?.containingClass == protocol.pyClass) { - return@visitMembers true - } + for (protocolMember in protocolMembers) { + val protocolElement = protocolMember.mainElement ?: continue + if (protocolElement is PyPossibleClassMember) { + val cls = protocolElement.containingClass + if (cls != null && !isProtocol(cls, context)) { + continue + } + } + if (protocolElement is PyTypeParameter) { + continue + } - val name = e.name ?: return@visitMembers true - when (name) { - PyNames.SLOTS -> return@visitMembers true // __slots__ in a protocol definition are not considered to be a part of the protocol - PyNames.CLASS_GETITEM -> return@visitMembers true - PyNames.CALL -> { - val types = subclass.getImplicitlyInvokedMethodTypes(null, resolveContext) - if (types.isNotEmpty()) { - result.add(Pair(e, types)) + if (protocolElement.contextOfType()?.containingClass == protocol.pyClass) { + continue + } + + when (val name = protocolMember.name) { + null -> continue + PyNames.SLOTS -> continue // __slots__ in a protocol definition are not considered to be a part of the protocol + PyNames.CLASS_GETITEM -> continue + PyNames.CALL -> { + val invokedMethods = subclass.getImplicitlyInvokedMethod(resolveContext) + if (invokedMethods.isNotEmpty()) { + result.add(Pair(protocolMember, invokedMethods)) + } + else { + val fallbackTypes = subclass.resolveImplicitlyInvokedMethods(null, resolveContext) + .mapNotNull { it.element } + .filterIsInstance() + .mapNotNull { + val type = resolveContext.typeEvalContext.getType(it) + if (type != null) { + it to type + } + else { + null + } } - else { - val fallbackTypes = subclass.resolveImplicitlyInvokedMethods(null, resolveContext) - .mapNotNull { it.element } - .filterIsInstance() - .mapNotNull { - val type = resolveContext.typeEvalContext.getType(it) - if (type != null) { - it to type - } - else { - null - } - } - result.add(Pair(e, fallbackTypes.map { PyTypedResolveResult(it.first, it.second) })) - } - } - else -> { - val types = subclass.getMemberTypes(name, null, AccessDirection.READ, resolveContext) - if (types != null) { - result.add(Pair(e, types)) - } - } + result.add(Pair(protocolMember, fallbackTypes.map { PyTypeMember(it.first, it.second) })) } } - - true - }, - true, - context - ) + else -> { + val subclassMembers = subclass.findMember(name, resolveContext) + result.add(Pair(protocolMember, subclassMembers)) + } + } + } return result } diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyProtocolInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyProtocolInspection.kt index f78c3b5077a4..838f584b9648 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyProtocolInspection.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyProtocolInspection.kt @@ -124,22 +124,25 @@ class PyProtocolInspection : PyInspection() { } private fun checkMemberCompatibility( - protocolElement: PyTypedElement, - subclassElements: List, + expectedMember: PyTypeMember, + subclassMembers: List, type: PyClassType, protocol: PyClassType, ) { - val expectedMemberType = myTypeEvalContext.getType(protocolElement) - - subclassElements + subclassMembers .asSequence() - .filter { it.element?.containingFile == type.pyClass.containingFile } - .filterNot { PyTypeChecker.match(expectedMemberType, it.type, myTypeEvalContext) } + .filter { it.mainElement?.containingFile == type.pyClass.containingFile } .forEach { - val element = it.element + val element = it.mainElement val place = if (element is PsiNameIdentifierOwner) element.nameIdentifier else element ?: return@forEach val elementName = if (element is PsiNameIdentifierOwner) element.name else return@forEach - registerProblem(place, PyPsiBundle.message("INSP.protocol.element.type.incompatible.with.protocol", elementName, protocol.name)) + + if (!PyTypeChecker.match(expectedMember.type, it.type, myTypeEvalContext)) { + registerProblem(place, PyPsiBundle.message("INSP.protocol.element.type.incompatible.with.protocol", elementName, protocol.name)) + } + else if (expectedMember.isWritable && !it.isWritable || expectedMember.isDeletable && !it.isDeletable) { + registerProblem(place, PyPsiBundle.message("INSP.protocol.element.type.not.writable", elementName, protocol.name)) + } } } 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 55330332f918..3d41410221ce 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 @@ -812,12 +812,11 @@ fun PyClassType.resolveImplicitlyInvokedMethods( else resolveDunderCall(callSite, resolveContext) } -fun PyClassType.getImplicitlyInvokedMethodTypes( - callSite: PyCallSiteExpression?, +fun PyClassType.getImplicitlyInvokedMethod( resolveContext: PyResolveContext, -): List { - return if (isDefinition()) getConstructorTypes(callSite, resolveContext) - else getDunderCallType(callSite, resolveContext) +): List { + return if (isDefinition()) getConstructorTypes(resolveContext) + else getDunderCallType(resolveContext) } private fun PyClassType.changeToImplicitlyInvokedMethods( @@ -866,15 +865,15 @@ private fun PyClassType.resolveConstructors(callSite: PyCallSiteExpression?, res return initAndNew.preferInitOverNew().map { RatedResolveResult(PyReferenceImpl.getRate(it, context), it) } } -private fun PyClassType.getConstructorTypes(callSite: PyCallSiteExpression?, resolveContext: PyResolveContext): List { - val initTypes = getMemberTypes(PyNames.INIT, callSite, AccessDirection.READ, resolveContext) - if (initTypes != null) { - return initTypes +private fun PyClassType.getConstructorTypes(resolveContext: PyResolveContext): List { + val initFunc = findMember(PyNames.INIT, resolveContext) + if (initFunc.isNotEmpty()) { + return initFunc } - val newTypes = getMemberTypes(PyNames.NEW, callSite, AccessDirection.READ, resolveContext) - if (newTypes != null) { - return newTypes + val newFunc = findMember(PyNames.NEW, resolveContext) + if (newFunc.isNotEmpty()) { + return newFunc } return emptyList() @@ -923,8 +922,8 @@ private fun PyClassLikeType.resolveDunderCall(location: PyExpression?, resolveCo return resolveMember(PyNames.CALL, location, AccessDirection.READ, resolveContext) ?: emptyList() } -private fun PyClassLikeType.getDunderCallType(location: PyExpression?, resolveContext: PyResolveContext): List { - return getMemberTypes(PyNames.CALL, location, AccessDirection.READ, resolveContext) ?: emptyList() +private fun PyClassLikeType.getDunderCallType(resolveContext: PyResolveContext): List { + return findMember(PyNames.CALL, resolveContext) } fun analyzeArguments( diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java index 80bbde9803fe..77c091bc4643 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java @@ -152,7 +152,8 @@ public class PyClassPatternImpl extends PyElementImpl implements PyClassPattern, @Nullable static Ref getMemberType(@NotNull PyType type, @NotNull String name, @NotNull TypeEvalContext context) { final PyResolveContext resolveContext = PyResolveContext.defaultContext(context); - final List results = type.getMemberTypes(name, null, AccessDirection.READ, resolveContext); - return ContainerUtil.isEmpty(results) ? null : Ref.create(getFirstItem(results).getType()); + List members = type.findMember(name, resolveContext); + if (members.isEmpty()) return null; + return Ref.create(PyUnionType.union(ContainerUtil.map(members, PyTypeMember::getType))); } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java index b3874785d788..d4fbab28427f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java @@ -18,10 +18,7 @@ import com.jetbrains.python.codeInsight.PyCustomMemberUtils; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; import com.jetbrains.python.psi.*; -import com.jetbrains.python.psi.impl.PyBuiltinCache; -import com.jetbrains.python.psi.impl.PyCallExpressionHelper; -import com.jetbrains.python.psi.impl.PyResolveResultRater; -import com.jetbrains.python.psi.impl.ResolveResultList; +import com.jetbrains.python.psi.impl.*; import com.jetbrains.python.psi.impl.references.PyReferenceImpl; import com.jetbrains.python.psi.resolve.*; import com.jetbrains.python.pyi.PyiUtil; @@ -591,6 +588,117 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType { return result; } + @Override + public @NotNull List<@NotNull PyTypeMember> getAllMembers(@NotNull PyResolveContext resolveContext) { + List<@NotNull PyTypeMember> result = new ArrayList<>(); + Set visited = new HashSet<>(); + for (Map.Entry entry : myClass.getProperties().entrySet()) { + visited.add(entry.getKey()); + Property property = entry.getValue(); + PyType type = property.getType(null, resolveContext.getTypeEvalContext()); + result.add(new PyTypeMember(property, type)); + } + + visitMembers(element -> { + if (element instanceof PsiNamedElement namedElement) { + if (visited.add(namedElement.getName())) { + PyType type = null; + if (element instanceof PyTypedElement typedElement) { + type = resolveContext.getTypeEvalContext().getType(typedElement); + } + + result.add(new PyTypeMember(element, type)); + } + } + return true; + }, false, resolveContext.getTypeEvalContext()); + + processProvidedMembers( + member -> { + PyTypeMember typeMember = convertCustomMemberToTypeMember(member, resolveContext); + if (typeMember != null) { + result.add(typeMember); + } + return true; + }, + null, + resolveContext.getTypeEvalContext() + ); + + return result; + } + + @Override + public @NotNull List<@NotNull PyTypeMember> findMember(@NotNull String name, @NotNull PyResolveContext resolveContext) { + Property property = myClass.findProperty(name, true, resolveContext.getTypeEvalContext()); + if (property != null) { + PyType type = property.getType(null, resolveContext.getTypeEvalContext()); + return List.of(new PyTypeMember(property, type)); + } + List customMembers = new ArrayList<>(); + processProvidedMembers( + member -> { + if (member.getName().equals(name)) { + PyTypeMember typeMember = convertCustomMemberToTypeMember(member, resolveContext); + if (typeMember != null) { + customMembers.add(typeMember); + } + } + return true; + }, + null, + resolveContext.getTypeEvalContext() + ); + if (!customMembers.isEmpty()) { + return customMembers; + } + List<@NotNull PyTypeMember> types = getMemberTypes(name, resolveContext); + if (types != null && !types.isEmpty()) { + return types; + } + return List.of(); + } + + private @Nullable PyTypeMember convertCustomMemberToTypeMember(@NotNull PyCustomMember customMember, + @NotNull PyResolveContext resolveContext) { + PsiElement element = customMember.resolve(getPyClass(), resolveContext); + if (element != null) { + PyType type = null; + if (element instanceof PyTypedElement typedElement) { + type = resolveContext.getTypeEvalContext().getType(typedElement); + } + return new PyTypeMember(element, type, customMember.isClassVar()); + } + return null; + } + + @Nullable + private List<@NotNull PyTypeMember> getMemberTypes(@NotNull String name, + final @NotNull PyResolveContext context) { + for (PyTypeProvider typeProvider : PyTypeProvider.EP_NAME.getExtensionList()) { + List types = typeProvider.getMemberTypes(this, name, null, AccessDirection.READ, context); + if (types != null) { + return types; + } + } + + List results = resolveMember(name, null, AccessDirection.READ, context); + if (results == null) { + return null; + } + + return ContainerUtil.map(results, result -> { + PsiElement element = result.getElement(); + if (element instanceof PyTypedElement typedElement) { + return new PyTypeMember(typedElement, + context.getTypeEvalContext().getType(typedElement), typedElement, typedElement, typedElement); + } + else { + return new PyTypeMember(element, null, element, element, element); + } + }); + } + private void processMembers(@NotNull Processor processor) { final PsiScopeProcessor scopeProcessor = new PsiScopeProcessor() { @Override 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 a6d6ecfb3a58..d7026eaa7c89 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 @@ -4,6 +4,7 @@ package com.jetbrains.python.psi.types; import com.intellij.openapi.util.*; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; +import com.intellij.psi.PsiNamedElement; import com.intellij.util.ArrayUtil; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; @@ -51,19 +52,18 @@ public final class PyTypeChecker { /** * Checks whether a type {@code actual} can be placed where {@code expected} is expected. - * + *

* For example {@code int} matches {@code object}, while {@code str} doesn't match {@code int}. * Work for builtin types, classes, tuples etc. - * + *

* Whether it's unknown if {@code actual} match {@code expected} the method returns {@code true}. * - * @implNote This behavior may be changed in future by replacing {@code boolean} with {@code Optional} and updating the clients. - * - * @param expected expected type - * @param actual type to be matched against expected - * @param context type evaluation context + * @param expected expected type + * @param actual type to be matched against expected + * @param context type evaluation context * @param substitutions map of substitutions for {@code expected} type * @return {@code false} if {@code expected} and {@code actual} don't match, true otherwise + * @implNote This behavior may be changed in future by replacing {@code boolean} with {@code Optional} and updating the clients. */ public static boolean match(@Nullable PyType expected, @Nullable PyType actual, @@ -93,7 +93,7 @@ public final class PyTypeChecker { /** * Perform type matching. - * + *

* Implementation details: *

    *
  • The method mutates {@code context.substitutions} map adding new entries into it @@ -243,7 +243,7 @@ public final class PyTypeChecker { /** * Check whether {@code expected} is Python *object* or *type*. - * + *

    * {@see PyTypeChecker#match(PyType, PyType, TypeEvalContext, Map)} */ private static @NotNull Optional matchObject(@NotNull PyClassType expected, @Nullable PyType actual) { @@ -262,7 +262,7 @@ public final class PyTypeChecker { /** * Match {@code actual} versus {@link PyTypeVarType} expected. - * + *

    * The method mutates {@code context.substitutions} map adding new entries into it */ private static boolean match(@NotNull PyTypeVarType expected, @Nullable PyType actual, @NotNull MatchContext context) { @@ -459,7 +459,9 @@ public final class PyTypeChecker { return false; } - private static boolean match(@NotNull PyCallableParameterListType expectedParameters, @Nullable PyType actual, @NotNull MatchContext context) { + private static boolean match(@NotNull PyCallableParameterListType expectedParameters, + @Nullable PyType actual, + @NotNull MatchContext context) { if (actual == null) return true; if (!(actual instanceof PyCallableParameterListType actualParameters)) return false; return matchCallableParameters(expectedParameters.getParameters(), actualParameters.getParameters(), context); @@ -487,7 +489,9 @@ public final class PyTypeChecker { return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false)); } - private static @NotNull Optional match(@NotNull PyTupleType expected, @NotNull PyUnionType actual, @NotNull MatchContext context) { + private static @NotNull Optional match(@NotNull PyTupleType expected, + @NotNull PyUnionType actual, + @NotNull MatchContext context) { final int elementCount = expected.getElementCount(); if (!expected.isHomogeneous()) { @@ -514,7 +518,9 @@ public final class PyTypeChecker { return ContainerUtil.or(expected.getMembers(), type -> match(type, actual, context).orElse(true)); } - private static @NotNull Optional match(@NotNull PyClassType expected, @NotNull PyClassType actual, @NotNull MatchContext matchContext) { + private static @NotNull Optional match(@NotNull PyClassType expected, + @NotNull PyClassType actual, + @NotNull MatchContext matchContext) { if (expected.equals(actual)) { return Optional.of(true); } @@ -574,32 +580,38 @@ public final class PyTypeChecker { GenericSubstitutions substitutions = collectTypeSubstitutions(actual, matchContext.context); MatchContext protocolContext = new MatchContext(matchContext.context, new GenericSubstitutions(), matchContext.reversedSubstitutions); - for (kotlin.Pair> pair : PyProtocolsKt.inspectProtocolSubclass(expected, actual, matchContext.context)) { - final List subclassElementTypes = ContainerUtil.map(pair.getSecond(), member -> member.getType()); - if (ContainerUtil.isEmpty(subclassElementTypes)) { + for (kotlin.Pair> pair : PyProtocolsKt.inspectProtocolSubclass(expected, actual, + matchContext.context)) { + final PyTypeMember protocolMember = pair.getFirst(); + final List subclassElementMembers = pair.getSecond(); + if (ContainerUtil.isEmpty(subclassElementMembers)) { return false; } - final PyType protocolElementType = dropSelfIfNeeded(expected, matchContext.context.getType(pair.getFirst()), matchContext.context); - final boolean elementResult = StreamEx - .of(subclassElementTypes) - .map(type -> dropSelfIfNeeded(actual, type, matchContext.context)) - .map(type -> substitute(type, substitutions, matchContext.context)) - .anyMatch( - subclassElementType -> { - boolean matched = match(protocolElementType, subclassElementType, protocolContext).orElse(true); - if (!matched) return false; - if (!(protocolElementType instanceof PyCallableType callableProtocolElement) || - !(subclassElementType instanceof PyCallableType callableSubclassElement)) return matched; - var protocolReturnType = callableProtocolElement.getReturnType(protocolContext.context); - if (protocolReturnType instanceof PySelfType) { - var subclassReturnType = callableSubclassElement.getReturnType(protocolContext.context); - if (subclassReturnType instanceof PySelfType) return true; - return match(actual, subclassReturnType, matchContext).orElse(true); - } - return matched; - } - ); + final PyType protocolElementType = dropSelfIfNeeded(expected, pair.getFirst().getType(), matchContext.context); + final boolean elementResult = ContainerUtil.exists(subclassElementMembers, subclassElementMember -> { + if (protocolMember.isWritable() && !subclassElementMember.isWritable()) { + return false; + } + if (protocolMember.isDeletable() && !subclassElementMember.isDeletable()) { + return false; + } + PyType subclassElementType = dropSelfIfNeeded(actual, subclassElementMember.getType(), matchContext.context); + subclassElementType = substitute(subclassElementType, substitutions, matchContext.context); + boolean matched = match(protocolElementType, subclassElementType, protocolContext).orElse(true); + if (!matched) return false; + if (!(protocolElementType instanceof PyCallableType callableProtocolElement) || + !(subclassElementType instanceof PyCallableType callableSubclassElement)) { + return matched; + } + var protocolReturnType = callableProtocolElement.getReturnType(protocolContext.context); + if (protocolReturnType instanceof PySelfType) { + var subclassReturnType = callableSubclassElement.getReturnType(protocolContext.context); + if (subclassReturnType instanceof PySelfType) return true; + return match(actual, subclassReturnType, matchContext).orElse(true); + } + return matched; + }); if (!elementResult) { return false; @@ -610,7 +622,8 @@ public final class PyTypeChecker { if (expected instanceof PyCollectionType) { PyCollectionType genericSuperClass = findGenericDefinitionType(expected.getPyClass(), matchContext.context); if (genericSuperClass != null) { - PyCollectionType concreteSuperClass = (PyCollectionType)substitute(genericSuperClass, protocolContext.mySubstitutions, protocolContext.context); + PyCollectionType concreteSuperClass = + (PyCollectionType)substitute(genericSuperClass, protocolContext.mySubstitutions, protocolContext.context); assert concreteSuperClass != null; return matchGenericClassesParameterWise((PyCollectionType)expected, concreteSuperClass, matchContext); } @@ -622,25 +635,29 @@ public final class PyTypeChecker { private static boolean match(PyClassType expectedProtocol, PyModuleType actualModule, MatchContext matchContext) { PyFile module = actualModule.getModule(); - Map moduleElements = StreamEx.of(ContainerUtil.concat(module.getTopLevelAttributes(), module.getTopLevelFunctions())) - .filter(e -> { - var name = ((PyQualifiedNameOwner)e).getName(); - return name != null && !PyNamesKt.isPrivate(name) && !PyNamesKt.isProtected(name); - }) - .toMap(PyTypedElement::getName, v -> v); + Map moduleElements = + StreamEx.of(ContainerUtil.concat(module.getTopLevelAttributes(), module.getTopLevelFunctions())) + .filter(e -> { + var name = ((PyQualifiedNameOwner)e).getName(); + return name != null && !PyNamesKt.isPrivate(name) && !PyNamesKt.isProtected(name); + }) + .toMap(PyTypedElement::getName, v -> v); var protocolElements = PyProtocolsKt.inspectProtocolSubclass(expectedProtocol, expectedProtocol, matchContext.context); if (protocolElements.size() != moduleElements.size()) return false; GenericSubstitutions substitutions = collectTypeSubstitutions(expectedProtocol, matchContext.context); - for (kotlin.Pair> pair : protocolElements) { - PyTypedElement protocolMember = pair.getFirst(); + for (kotlin.Pair> pair : protocolElements) { + PsiElement pm = pair.getFirst().getMainElement(); + if (!(pm instanceof PsiNamedElement protocolMember)) { + continue; + } String name = protocolMember.getName(); PyTypedElement moduleElement = moduleElements.get(name); if (moduleElement != null) { PyType expectedProtocolMemberType = - substitute(dropSelfIfNeeded(expectedProtocol, matchContext.context.getType(pair.getFirst()), matchContext.context), substitutions, + substitute(dropSelfIfNeeded(expectedProtocol, pair.getFirst().getType(), matchContext.context), substitutions, matchContext.context); PyType actualModuleElementType = matchContext.context.getType(moduleElement); if (!match(expectedProtocolMemberType, actualModuleElementType, matchContext.context)) { @@ -665,7 +682,9 @@ public final class PyTypeChecker { } // https://typing.python.org/en/latest/spec/tuples.html#type-compatibility-rules - private static @NotNull Optional match(@NotNull PyTupleType expected, @NotNull PyTupleType actual, @NotNull MatchContext context) { + private static @NotNull Optional match(@NotNull PyTupleType expected, + @NotNull PyTupleType actual, + @NotNull MatchContext context) { if (actual.isHomogeneous()) { // The type tuple[Any, ...] is consistent with any tuple final PyType elementType = actual.getIteratedItemType(); @@ -939,7 +958,8 @@ public final class PyTypeChecker { assert entry.getValue() instanceof PyPositionalVariadicType; result.typeVarTuples.put(typeVarTuple, (PyPositionalVariadicType)entry.getValue()); } - else if (entry.getKey() instanceof PyParamSpecType specType && entry.getValue() instanceof PyCallableParameterVariadicType paramSpecType) { + else if (entry.getKey() instanceof PyParamSpecType specType && + entry.getValue() instanceof PyCallableParameterVariadicType paramSpecType) { result.paramSpecs.put(specType, paramSpecType); } } @@ -1104,8 +1124,8 @@ public final class PyTypeChecker { else { existingSubstitutions.paramSpecs.put(paramSpecType, new PyCallableParameterListTypeImpl( List.of(PyCallableParameterImpl.positionalNonPsi("args", null), - PyCallableParameterImpl.keywordNonPsi("kwargs", null))) - ); + PyCallableParameterImpl.keywordNonPsi("kwargs", null))) + ); } } } @@ -1700,8 +1720,8 @@ public final class PyTypeChecker { */ @ApiStatus.Internal public static @Nullable PyType parameterizeType(@NotNull PyType genericType, - @NotNull List actualTypeParams, - @NotNull TypeEvalContext context) { + @NotNull List actualTypeParams, + @NotNull TypeEvalContext context) { Generics typeParams = collectGenerics(genericType, context); if (!typeParams.isEmpty()) { List expectedTypeParams = new ArrayList<>(new LinkedHashSet<>(typeParams.getAllTypeParameters())); diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index b24259041d9b..db5e9f6c0a21 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -112,7 +112,7 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { public void testFunctionReturnTypePy3() { doTest(); } - + public void testFunctionYieldTypePy3() { doTest(); } @@ -2658,7 +2658,7 @@ def foo(param: str | int) -> TypeGuard[str]: compatible: MyCallable[[int], object] = MyCallable[[object], str]() incompatible1: MyCallable[[object], object] = MyCallable[[int], str]() incompatible2: MyCallable[[int], str] = MyCallable[[object], object]() - """); + """); } // PY-77541 @@ -3125,4 +3125,189 @@ def foo(param: str | int) -> TypeGuard[str]: var: Template = Concrete("value", 42) """); } + + // PY-76822 + public void testProtocolWithPropertyAndConcreteWithAttribute() { + doTestByText(""" + from typing import Protocol + + class Template(Protocol): + @property + def val1(self) -> int: + ... + + + class Concrete: + val1: int = 0 + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testProtocolWithPropertyAndConcreteWithProperty() { + doTestByText(""" + from typing import Protocol + + class Template(Protocol): + @property + def val1(self) -> int: + ... + + + class Concrete: + @property + def val1(self) -> int: + ... + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testProtocolWithPropertySetterAndConcreteWithPropertyDeleter() { + doTestByText(""" + from typing import Protocol + + class Template(Protocol): + @property + def val1(self) -> int: + ... + + @val1.setter + def val1(self, val: int) -> None: + ... + + + class Concrete: + @property + def val1(self) -> int: + ... + + @val1.deleter + def val1(self, val: int) -> None: + ... + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testProtocolWithPropertySetterAndFrozenDataclass() { + doTestByText(""" + from typing import Protocol + from dataclasses import dataclass + + class Template(Protocol): + @property + def val(self) -> int: + ... + + @val.setter + def val(self, val: int) -> None: + ... + + + @dataclass(frozen=True) + class Concrete: + val: int = 0 + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testProtocolWithPropertyDeleterAndFrozenDataclass() { + doTestByText(""" + from typing import Protocol + from dataclasses import dataclass + + class Template(Protocol): + @property + def val(self) -> int: + ... + + @val.deleter + def val(self, val: int) -> None: + ... + + + @dataclass(frozen=True) + class Concrete: + val: int = 0 + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testOverloadedMethodInConcreteClass() { + doTestByText(""" + from typing import Protocol, ClassVar, overload + + class Template(Protocol): + def f(self, x: int) -> int: ... + + + class Concrete: + @overload + def f(self, x: str) -> int: ... + + @overload + def f(self, x: int) -> int: ... + + def f(self, x) -> int: + return 1 + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testExplicitAnyInConcreteType() { + doTestByText(""" + from typing import Protocol, Any + + class Template(Protocol): + val: int + + + class Concrete: + val: Any + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testExplicitAnyInProtocol() { + doTestByText(""" + from typing import Protocol, Any + + class Template(Protocol): + val: Any + + + class Concrete: + val: int + + var: Template = Concrete() + """); + } + + // PY-76822 + public void testExplicitAnyInBothProtocolAndConcreteType() { + doTestByText(""" + from typing import Protocol, Any + + class Template(Protocol): + val: Any + + + class Concrete: + val: Any + + var: Template = Concrete() + """); + } } diff --git a/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java index f8e68f2710d0..c133a967b038 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java @@ -37,6 +37,82 @@ public class PyProtocolInspectionTest extends PyInspectionTestCase { doTest(); } + // PY-76822 + public void testProtocolWithPropertyAndConcreteWithAttribute() { + doTestByText(""" + from typing import Protocol + + class Template(Protocol): + @property + def val1(self) -> int: + ... + + + class Concrete(Template): + val1: int = 0 + """); + } + + // PY-76822 + public void testProtocolWithPropertyAndConcreteWithProperty() { + doTestByText(""" + from typing import Protocol + + class Template(Protocol): + @property + def val1(self) -> int: + ... + + + class Concrete(Template): + @property + def val1(self) -> int: + ... + """); + } + + // PY-76822 + public void testProtocolWithMutablePropertyAndClassAttribute() { + doTestByText(""" + from typing import Protocol, Sequence + + class Template(Protocol): + @property + def val(self) -> Sequence[float]: + ... + + @val.setter + def val(self, val: Sequence[float]) -> None: + ... + + + class Concrete(Template): + val: Sequence[float] = [0] + """); + } + + // PY-76822 + public void testProtocolAndFrozenDataclass() { + doTestByText(""" + from typing import Protocol + from dataclasses import dataclass + + class Template(Protocol): + @property + def val(self) -> int: + ... + + @val.setter + def val(self, val: int) -> None: + ... + + + @dataclass(frozen=True) + class Concrete(Template): + val: int = 0 + """); + } + // PY-61857 public void testClassWithTypeParameterListNotReported() { runWithLanguageLevel(LanguageLevel.PYTHON312, () -> {