PY-84464 Support @property decorator when matching a protocol and its implementation

GitOrigin-RevId: fcbeeb57323336c7d921edad373ea595ef687d6b
This commit is contained in:
evgeny.bovykin
2025-10-21 09:52:41 +00:00
committed by intellij-monorepo-bot
parent 0bbcf6273e
commit 044c140516
17 changed files with 629 additions and 198 deletions
@@ -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);
}
@@ -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<PyTypedResolveResult> 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<? extends RatedResolveResult> 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();
}
/**
@@ -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
}
@@ -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;
}
@@ -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
}
}
@@ -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
@@ -301,6 +301,26 @@ public final class PyCustomType implements PyClassLikeType {
return result;
}
@Override
public @NotNull List<@NotNull PyTypeMember> getAllMembers(@NotNull PyResolveContext resolveContext) {
List<PyTypeMember> 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<PyTypeMember> members = type.findMember(name, resolveContext);
if (!members.isEmpty()) {
return members;
}
}
return List.of();
}
/**
* Predicate that filters completion using {@link #myFilter}
*/
@@ -77,15 +77,13 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
return null
}
override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List<PyTypedResolveResult>? {
override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List<PyTypeMember>? {
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
@@ -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<PyTypedResolveResult>? {
override fun getMemberTypes(type: PyType, name: String, location: PyExpression?, direction: AccessDirection, context: PyResolveContext): List<PyTypeMember>? {
if (type !is PyNamedTupleType) return null
type.fields[name]?.let {
return listOf(PyTypedResolveResult(null, it.type))
return listOf(PyTypeMember(null, it.type))
}
return null
}
@@ -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<PyTypedElement, List<PyTypedResolveResult>>
typealias ProtocolAndSubclassElements = Pair<PyTypeMember, List<PyTypeMember>>
fun inspectProtocolSubclass(protocol: PyClassType, subclass: PyClassType, context: TypeEvalContext): List<ProtocolAndSubclassElements> {
val resolveContext = PyResolveContext.defaultContext(context)
val result = mutableListOf<Pair<PyTypedElement, List<PyTypedResolveResult>>>()
val result = mutableListOf<Pair<PyTypeMember, List<PyTypeMember>>>()
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<PyFunction>()?.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<PyFunction>()?.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<PyTypedElement>()
.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<PyTypedElement>()
.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
}
@@ -124,22 +124,25 @@ class PyProtocolInspection : PyInspection() {
}
private fun checkMemberCompatibility(
protocolElement: PyTypedElement,
subclassElements: List<PyTypedResolveResult>,
expectedMember: PyTypeMember,
subclassMembers: List<PyTypeMember>,
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))
}
}
}
@@ -812,12 +812,11 @@ fun PyClassType.resolveImplicitlyInvokedMethods(
else resolveDunderCall(callSite, resolveContext)
}
fun PyClassType.getImplicitlyInvokedMethodTypes(
callSite: PyCallSiteExpression?,
fun PyClassType.getImplicitlyInvokedMethod(
resolveContext: PyResolveContext,
): List<PyTypedResolveResult> {
return if (isDefinition()) getConstructorTypes(callSite, resolveContext)
else getDunderCallType(callSite, resolveContext)
): List<PyTypeMember> {
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<PyTypedResolveResult> {
val initTypes = getMemberTypes(PyNames.INIT, callSite, AccessDirection.READ, resolveContext)
if (initTypes != null) {
return initTypes
private fun PyClassType.getConstructorTypes(resolveContext: PyResolveContext): List<PyTypeMember> {
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<PyTypedResolveResult> {
return getMemberTypes(PyNames.CALL, location, AccessDirection.READ, resolveContext) ?: emptyList()
private fun PyClassLikeType.getDunderCallType(resolveContext: PyResolveContext): List<PyTypeMember> {
return findMember(PyNames.CALL, resolveContext)
}
fun analyzeArguments(
@@ -152,7 +152,8 @@ public class PyClassPatternImpl extends PyElementImpl implements PyClassPattern,
@Nullable
static Ref<PyType> getMemberType(@NotNull PyType type, @NotNull String name, @NotNull TypeEvalContext context) {
final PyResolveContext resolveContext = PyResolveContext.defaultContext(context);
final List<PyTypedResolveResult> results = type.getMemberTypes(name, null, AccessDirection.READ, resolveContext);
return ContainerUtil.isEmpty(results) ? null : Ref.create(getFirstItem(results).getType());
List<PyTypeMember> members = type.findMember(name, resolveContext);
if (members.isEmpty()) return null;
return Ref.create(PyUnionType.union(ContainerUtil.map(members, PyTypeMember::getType)));
}
}
@@ -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<String> visited = new HashSet<>();
for (Map.Entry<String, Property> 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<PyTypeMember> 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<PyTypeMember> types = typeProvider.getMemberTypes(this, name, null, AccessDirection.READ, context);
if (types != null) {
return types;
}
}
List<? extends RatedResolveResult> 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<? super PsiElement> processor) {
final PsiScopeProcessor scopeProcessor = new PsiScopeProcessor() {
@Override
@@ -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.
*
* <p>
* For example {@code int} matches {@code object}, while {@code str} doesn't match {@code int}.
* Work for builtin types, classes, tuples etc.
*
* <p>
* 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<Boolean>} 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<Boolean>} 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.
*
* <p>
* Implementation details:
* <ul>
* <li>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*.
*
* <p>
* {@see PyTypeChecker#match(PyType, PyType, TypeEvalContext, Map)}
*/
private static @NotNull Optional<Boolean> matchObject(@NotNull PyClassType expected, @Nullable PyType actual) {
@@ -262,7 +262,7 @@ public final class PyTypeChecker {
/**
* Match {@code actual} versus {@link PyTypeVarType} expected.
*
* <p>
* 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<Boolean> match(@NotNull PyTupleType expected, @NotNull PyUnionType actual, @NotNull MatchContext context) {
private static @NotNull Optional<Boolean> 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<Boolean> match(@NotNull PyClassType expected, @NotNull PyClassType actual, @NotNull MatchContext matchContext) {
private static @NotNull Optional<Boolean> 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<PyTypedElement, List<PyTypedResolveResult>> pair : PyProtocolsKt.inspectProtocolSubclass(expected, actual, matchContext.context)) {
final List<PyType> subclassElementTypes = ContainerUtil.map(pair.getSecond(), member -> member.getType());
if (ContainerUtil.isEmpty(subclassElementTypes)) {
for (kotlin.Pair<PyTypeMember, List<PyTypeMember>> pair : PyProtocolsKt.inspectProtocolSubclass(expected, actual,
matchContext.context)) {
final PyTypeMember protocolMember = pair.getFirst();
final List<PyTypeMember> 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<String, PyTypedElement> 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<String, PyTypedElement> 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<PyTypedElement, List<PyTypedResolveResult>> pair : protocolElements) {
PyTypedElement protocolMember = pair.getFirst();
for (kotlin.Pair<PyTypeMember, List<PyTypeMember>> 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<Boolean> match(@NotNull PyTupleType expected, @NotNull PyTupleType actual, @NotNull MatchContext context) {
private static @NotNull Optional<Boolean> 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<PyType> actualTypeParams,
@NotNull TypeEvalContext context) {
@NotNull List<PyType> actualTypeParams,
@NotNull TypeEvalContext context) {
Generics typeParams = collectGenerics(genericType, context);
if (!typeParams.isEmpty()) {
List<PyType> expectedTypeParams = new ArrayList<>(new LinkedHashSet<>(typeParams.getAllTypeParameters()));
@@ -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] = <warning descr="Expected type 'MyCallable[[object], object]', got 'MyCallable[[int], str]' instead">MyCallable[[int], str]()</warning>
incompatible2: MyCallable[[int], str] = <warning descr="Expected type 'MyCallable[[int], str]', got 'MyCallable[[object], object]' instead">MyCallable[[object], object]()</warning>
""");
""");
}
// 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 = <warning descr="Expected type 'Template', got 'Concrete' instead">Concrete()</warning>
""");
}
// 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 = <warning descr="Expected type 'Template', got 'Concrete' instead">Concrete()</warning>
""");
}
// 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 = <warning descr="Expected type 'Template', got 'Concrete' instead">Concrete()</warning>
""");
}
// 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()
""");
}
}
@@ -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):
<warning descr="'val' is writable in protocol 'Template'">val</warning>: int = 0
""");
}
// PY-61857
public void testClassWithTypeParameterListNotReported() {
runWithLanguageLevel(LanguageLevel.PYTHON312, () -> {