From cdad986c9ac76e5a43bf068561bb59891a6202a4 Mon Sep 17 00:00:00 2001 From: Morgan Bartholomew Date: Sun, 15 Feb 2026 04:38:04 +1000 Subject: [PATCH] refactor python: j2k `PyTypeChecker` GitOrigin-RevId: 5bcb5920395cdd03f14b52b530b9cdf547d6b212 --- .../typing/PyTypingTypeProvider.kt | 9 +- .../inspections/PyTypeHintsInspection.kt | 12 +- .../python/psi/impl/PyCallExpressionHelper.kt | 2 +- .../psi/types/PyExpectedTypeJudgement.kt | 2 +- .../python/psi/types/PyTypeChecker.kt | 2631 +++++++++-------- 5 files changed, 1341 insertions(+), 1315 deletions(-) diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt index efe75a83e6e5..d2befa59b12f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt @@ -117,6 +117,7 @@ import com.jetbrains.python.psi.types.PySelfType import com.jetbrains.python.psi.types.PyTupleType import com.jetbrains.python.psi.types.PyType import com.jetbrains.python.psi.types.PyTypeChecker +import com.jetbrains.python.psi.types.PyTypeChecker.collectGenerics import com.jetbrains.python.psi.types.PyTypeParameterMapping import com.jetbrains.python.psi.types.PyTypeParameterType import com.jetbrains.python.psi.types.PyTypeParser @@ -1076,7 +1077,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext() { .map { getType(it!!, context) } .map { Ref.deref(it) } .flatMap { - val typeParams = PyTypeChecker.collectGenerics(it, context.typeContext) + val typeParams = it.collectGenerics(context.typeContext) StreamEx.of(typeParams.typeVars).append(typeParams.typeVarTuples) .append(StreamEx.of(typeParams.paramSpecs)) } @@ -2386,9 +2387,9 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext() { } .append(PyTypingTypeProvider().getReturnType(function, context)) .map { Ref.deref(it) } - .map { PyTypeChecker.collectGenerics(it, context) } + .map { it.collectGenerics(context) } .flatMap { - StreamEx.of(it!!.typeVars) + StreamEx.of(it.typeVars) .append(it.paramSpecs) .append(it.typeVarTuples) } @@ -2507,7 +2508,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext() { if (typeHint is PyReferenceExpression) { if (assignedType !is PyTypeParameterType) { val typeAliasTypeParams = - PyTypeChecker.collectGenerics(assignedType, context.typeContext).allTypeParameters + assignedType.collectGenerics(context.typeContext).allTypeParameters if (!typeAliasTypeParams.isEmpty()) { return Ref( PyTypeChecker.parameterizeType( 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 40d8780b84f6..f536bed0c75b 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 @@ -97,6 +97,8 @@ import com.jetbrains.python.psi.types.PySelfType import com.jetbrains.python.psi.types.PyTupleType import com.jetbrains.python.psi.types.PyType import com.jetbrains.python.psi.types.PyTypeChecker +import com.jetbrains.python.psi.types.PyTypeChecker.collectGenerics +import com.jetbrains.python.psi.types.PyTypeChecker.hasGenerics import com.jetbrains.python.psi.types.PyTypeParameterMapping import com.jetbrains.python.psi.types.PyTypeParameterType import com.jetbrains.python.psi.types.PyTypeVarTupleType @@ -607,7 +609,7 @@ class PyTypeHintsInspection : PyInspection() { if (it != null) { val type = PyTypingTypeProvider.getType(it, myTypeEvalContext)?.get() - if (PyTypeChecker.hasGenerics(type, myTypeEvalContext)) { + if (type.hasGenerics(myTypeEvalContext)) { registerProblem(it, PyPsiBundle.message("INSP.type.hints.typevar.constraints.cannot.be.parametrized.by.type.variables")) } } @@ -1230,12 +1232,12 @@ class PyTypeHintsInspection : PyInspection() { private fun collectTypeParametersFromTypeAlias(assignedValue: PyExpression, assignedValueType: PyType, isExplicitTypeAlias: Boolean): PyTypeChecker.Generics { if (isExplicitTypeAlias || !(assignedValue is PyReferenceExpression && assignedValueType is PyClassType)) { - return PyTypeChecker.collectGenerics(assignedValueType, myTypeEvalContext) + return assignedValueType.collectGenerics(myTypeEvalContext) } else { val genericDefinitionType = PyTypeChecker.findGenericDefinitionType(assignedValueType.pyClass, myTypeEvalContext) ?: return PyTypeChecker.Generics() - return PyTypeChecker.collectGenerics(genericDefinitionType, myTypeEvalContext) + return genericDefinitionType.collectGenerics(myTypeEvalContext) } } @@ -1323,7 +1325,7 @@ class PyTypeHintsInspection : PyInspection() { PyPsiBundle.message("INSP.type.hints.default.type.var.cannot.follow.type.var.tuple"), ProblemHighlightType.GENERIC_ERROR) } - val genericTypesInDefaultExpr = PyTypeChecker.collectGenerics(Ref.deref(defaultType), myTypeEvalContext) + val genericTypesInDefaultExpr = Ref.deref(defaultType).collectGenerics(myTypeEvalContext) val defaultOutOfScope = genericTypesInDefaultExpr.allTypeParameters .firstOrNull { typeVar -> typeVar.declarationElement != null && typeVar.declarationElement !in typeParamDeclarations } @@ -1554,7 +1556,7 @@ class PyTypeHintsInspection : PyInspection() { val selfAnnotationValue = selfParameter.annotation?.value ?: return val selfAnnotationType = PyTypingTypeProvider.getType(selfAnnotationValue, myTypeEvalContext) ?: return - val generics = PyTypeChecker.collectGenerics(selfAnnotationType.get(), myTypeEvalContext) + val generics = selfAnnotationType.get().collectGenerics(myTypeEvalContext) if (generics.typeVars.any { it.scopeOwner === containingClass }) { registerProblem(selfAnnotationValue, PyPsiBundle.message("INSP.type.hints.cannot.use.class.scope.type.variables.in.annotation.for.self.parameter.of__init__")) 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 fa68b7bfa7b0..d6bd440fb899 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 @@ -1237,7 +1237,7 @@ private fun PyFunction.matchesByArgumentTypes(callSite: PyCallSiteExpression, co val receiver = callSite.getReceiver(this) val mappedExplicitParameters = fullMapping.mappedParameters - val allMappedParameters = LinkedHashMap() + val allMappedParameters = LinkedHashMap() val firstImplicit = fullMapping.implicitParameters.firstOrNull() if (receiver != null && firstImplicit != null) { allMappedParameters.put(receiver, firstImplicit) diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt index f10c7f094f03..7802039fb9d2 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt @@ -302,7 +302,7 @@ object PyExpectedTypeJudgement { mappedParameters: Map, ctx: TypeEvalContext, ): PyType? { - if (!hasGenerics(paramType, ctx)) return paramType + if (!paramType.hasGenerics(ctx)) return paramType val receiver = callSite.getReceiver(null) val substitutions = unifyGenericCall(receiver, mappedParameters, ctx) // might cause recursion diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt index 8882501697cf..edfe0048c5b8 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt @@ -1,669 +1,663 @@ // Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. -package com.jetbrains.python.psi.types; +package com.jetbrains.python.psi.types -import com.intellij.openapi.util.Condition; -import com.intellij.openapi.util.Couple; -import com.intellij.openapi.util.Pair; -import com.intellij.openapi.util.RecursionManager; -import com.intellij.openapi.util.Ref; -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; -import com.jetbrains.python.PyNamesKt; -import com.jetbrains.python.PythonRuntimeService; -import com.jetbrains.python.ast.PyAstFunction; -import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; -import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; -import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; -import com.jetbrains.python.psi.AccessDirection; -import com.jetbrains.python.psi.LanguageLevel; -import com.jetbrains.python.psi.PyCallable; -import com.jetbrains.python.psi.PyClass; -import com.jetbrains.python.psi.PyExpression; -import com.jetbrains.python.psi.PyFile; -import com.jetbrains.python.psi.PyFunction; -import com.jetbrains.python.psi.PyListLiteralExpression; -import com.jetbrains.python.psi.PyParameter; -import com.jetbrains.python.psi.PyQualifiedNameOwner; -import com.jetbrains.python.psi.PySequenceExpression; -import com.jetbrains.python.psi.PyTupleExpression; -import com.jetbrains.python.psi.PyTypedElement; -import com.jetbrains.python.psi.PyUtil; -import com.jetbrains.python.psi.impl.ParamHelper; -import com.jetbrains.python.psi.impl.PyBuiltinCache; -import com.jetbrains.python.psi.impl.PyPsiUtils; -import com.jetbrains.python.psi.impl.PyTypeProvider; -import com.jetbrains.python.psi.resolve.PyResolveContext; -import com.jetbrains.python.psi.resolve.RatedResolveResult; -import com.jetbrains.python.psi.types.PyRecursiveTypeVisitor.PyTypeTraverser; -import com.jetbrains.python.psi.types.PyRecursiveTypeVisitor.Traversal; -import com.jetbrains.python.psi.types.PyTypeParameterMapping.Option; -import com.jetbrains.python.pyi.PyiFile; -import com.jetbrains.python.sdk.legacy.PythonSdkUtil; -import one.util.streamex.EntryStream; -import one.util.streamex.IntStreamEx; -import one.util.streamex.StreamEx; -import org.jetbrains.annotations.ApiStatus; -import org.jetbrains.annotations.NotNull; -import org.jetbrains.annotations.Nullable; - -import java.util.ArrayList; -import java.util.Collection; -import java.util.Collections; -import java.util.HashSet; -import java.util.LinkedHashMap; -import java.util.LinkedHashSet; -import java.util.List; -import java.util.Map; -import java.util.Objects; -import java.util.Optional; -import java.util.Set; - -import static com.jetbrains.python.PyNames.FUNCTION; -import static com.jetbrains.python.psi.PyUtil.as; -import static com.jetbrains.python.psi.PyUtil.getReturnTypeToAnalyzeAsCallType; -import static com.jetbrains.python.psi.PyUtil.isNewMethod; -import static com.jetbrains.python.psi.impl.PyCallExpressionHelper.getArgumentsMappedToKeywordContainer; -import static com.jetbrains.python.psi.impl.PyCallExpressionHelper.getArgumentsMappedToPositionalContainer; -import static com.jetbrains.python.psi.impl.PyCallExpressionHelper.getMappedKeywordContainer; -import static com.jetbrains.python.psi.impl.PyCallExpressionHelper.getMappedPositionalContainer; -import static com.jetbrains.python.psi.impl.PyCallExpressionHelper.getRegularMappedParameters; - -public final class PyTypeChecker { - private PyTypeChecker() { - } +import com.intellij.openapi.util.Computable +import com.intellij.openapi.util.RecursionManager +import com.intellij.openapi.util.Ref +import com.intellij.psi.PsiElement +import com.intellij.psi.PsiNameIdentifierOwner +import com.intellij.psi.PsiNamedElement +import com.intellij.util.ArrayUtil +import com.intellij.util.containers.ContainerUtil +import com.jetbrains.python.PyNames +import com.jetbrains.python.PythonRuntimeService +import com.jetbrains.python.ast.PyAstFunction +import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.codeInsight.typing.inspectProtocolSubclass +import com.jetbrains.python.codeInsight.typing.isProtocol +import com.jetbrains.python.isPrivate +import com.jetbrains.python.isProtected +import com.jetbrains.python.psi.AccessDirection +import com.jetbrains.python.psi.LanguageLevel +import com.jetbrains.python.psi.PyClass +import com.jetbrains.python.psi.PyExpression +import com.jetbrains.python.psi.PyFile +import com.jetbrains.python.psi.PyFunction +import com.jetbrains.python.psi.PyListLiteralExpression +import com.jetbrains.python.psi.PyQualifiedNameOwner +import com.jetbrains.python.psi.PySequenceExpression +import com.jetbrains.python.psi.PyTupleExpression +import com.jetbrains.python.psi.PyTypedElement +import com.jetbrains.python.psi.PyUtil +import com.jetbrains.python.psi.impl.ParamHelper +import com.jetbrains.python.psi.impl.PyBuiltinCache.Companion.getInstance +import com.jetbrains.python.psi.impl.PyPsiUtils +import com.jetbrains.python.psi.impl.PyTypeProvider +import com.jetbrains.python.psi.impl.getArgumentsMappedToKeywordContainer +import com.jetbrains.python.psi.impl.getArgumentsMappedToPositionalContainer +import com.jetbrains.python.psi.impl.getMappedKeywordContainer +import com.jetbrains.python.psi.impl.getMappedPositionalContainer +import com.jetbrains.python.psi.impl.getRegularMappedParameters +import com.jetbrains.python.psi.resolve.PyResolveContext +import com.jetbrains.python.psi.types.PyCallableParameterMapping.mapCallableParameters +import com.jetbrains.python.psi.types.PyLiteralStringType.Companion.match +import com.jetbrains.python.psi.types.PyLiteralType.Companion.match +import com.jetbrains.python.psi.types.PyRecursiveTypeVisitor.PyTypeTraverser +import com.jetbrains.python.psi.types.PyTypeChecker.match +import com.jetbrains.python.pyi.PyiFile +import com.jetbrains.python.sdk.legacy.PythonSdkUtil +import org.jetbrains.annotations.ApiStatus +import java.util.Collections +import java.util.Optional +object PyTypeChecker { /** - * See {@link PyTypeChecker#match(PyType, PyType, TypeEvalContext, Map)} for description. + * See [match] for description. */ - public static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context) { - @NotNull GenericSubstitutions substitutions = new GenericSubstitutions(); - return match(expected, actual, new MatchContext(context, substitutions, false)).orElse(true); + @JvmStatic + fun match(expected: PyType?, actual: PyType?, context: TypeEvalContext): Boolean { + val substitutions = GenericSubstitutions() + return match(expected, actual, MatchContext(context, substitutions, false)).orElse(true)!! } /** - * 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}. + * Checks whether a type `actual` can be placed where `expected` is expected. + * + * + * For example `int` matches `object`, while `str` doesn't match `int`. * Work for builtin types, classes, tuples etc. - *

- * Whether it's unknown if {@code actual} match {@code expected} the method returns {@code true}. + * + * + * Whether it's unknown if `actual` match `expected` the method returns `true`. * * @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. + * @param substitutions map of substitutions for `expected` type + * @return `false` if `expected` and `actual` don't match, true otherwise + * @implNote This behavior may be changed in future by replacing `boolean` with `Optional` and updating the clients. */ - public static boolean match(@Nullable PyType expected, - @Nullable PyType actual, - @NotNull TypeEvalContext context, - @NotNull Map typeVars) { - var substitutions = new GenericSubstitutions(typeVars); - return match(expected, actual, new MatchContext(context, substitutions, false)).orElse(true); + fun match( + expected: PyType?, + actual: PyType?, + context: TypeEvalContext, + typeVars: Map, + ): Boolean { + val substitutions = GenericSubstitutions(typeVars) + return match(expected, actual, MatchContext(context, substitutions, false)).orElse(true)!! } - public static boolean match(@Nullable PyType expected, - @Nullable PyType actual, - @NotNull TypeEvalContext context, - @NotNull GenericSubstitutions substitutions) { - return match(expected, actual, new MatchContext(context, substitutions, false)) - .orElse(true); + @JvmStatic + fun match( + expected: PyType?, + actual: PyType?, + context: TypeEvalContext, + substitutions: GenericSubstitutions, + ): Boolean { + return match(expected, actual, MatchContext(context, substitutions, false)) + .orElse(true)!! } - private static @NotNull Optional match(@Nullable PyType expected, @Nullable PyType actual, @NotNull MatchContext context) { - final Optional result = RecursionManager.doPreventingRecursion( - Pair.create(expected, actual), - false, - () -> matchImpl(expected, actual, context) - ); + private fun match(expected: PyType?, actual: PyType?, context: MatchContext): Optional { + val result = RecursionManager.doPreventingRecursion(expected to actual, false) { + matchImpl(expected, actual, context) + } - return result == null ? Optional.of(true) : result; + return result ?: Optional.of(true) } /** * Perform type matching. - *

+ * + * * Implementation details: - *

    - *
  • The method mutates {@code context.substitutions} map adding new entries into it - *
  • The order of match subroutine calls is important - *
  • The method may recursively call itself - *
+ * + * * The method mutates `context.substitutions` map adding new entries into it + * * The order of match subroutine calls is important + * * The method may recursively call itself + * */ - private static @NotNull Optional matchImpl(@Nullable PyType expected, @Nullable PyType actual, @NotNull MatchContext context) { - if (Objects.equals(expected, actual)) { - return Optional.of(true); + private fun matchImpl(expected: PyType?, actual: PyType?, context: MatchContext): Optional { + if (expected == actual) { + return Optional.of(true) } - for (PyTypeCheckerExtension extension : PyTypeCheckerExtension.EP_NAME.getExtensionList()) { - final Optional result = extension.match(expected, actual, context.context, context.mySubstitutions); - if (result.isPresent()) { - return result; + for (extension in PyTypeCheckerExtension.EP_NAME.extensionList) { + val result = extension.match(expected, actual, context.context, context.mySubstitutions) + if (result.isPresent) { + return result } } - if (expected instanceof PyClassType) { - Optional match = matchObject((PyClassType)expected, actual); - if (match.isPresent()) { - return match; + if (expected is PyClassType) { + val match = matchObject(expected, actual) + if (match.isPresent) { + return match } } - if (actual instanceof PyNarrowedType actualNarrowedType && expected instanceof PyNarrowedType expectedNarrowedType) { - return match(expectedNarrowedType, actualNarrowedType, context); + if (actual is PyNarrowedType && expected is PyNarrowedType) { + return match(expected, actual, context) } - if (actual instanceof PyTypeVarTupleType typeVarTupleType && context.reversedSubstitutions) { - return Optional.of(match(typeVarTupleType, expected, context)); + if (actual is PyTypeVarTupleType && context.reversedSubstitutions) { + return Optional.of(match(actual, expected, context)) } - if (expected instanceof PyPositionalVariadicType variadic) { - return Optional.of(match(variadic, actual, context)); + if (expected is PyPositionalVariadicType) { + return Optional.of(match(expected, actual, context)) } - if (actual instanceof PyTypeVarType typeVarType && context.reversedSubstitutions) { - return Optional.of(match(typeVarType, expected, context)); + if (actual is PyTypeVarType && context.reversedSubstitutions) { + return Optional.of(match(actual, expected, context)) } - if (expected instanceof PyTypeVarType typeVarType) { - return Optional.of(match(typeVarType, actual, context)); + if (expected is PyTypeVarType) { + return Optional.of(match(expected, actual, context)) } - if (expected instanceof PySelfType selfType) { - return Optional.of(match(selfType, actual, context)); + if (expected is PySelfType) { + return Optional.of(match(expected, actual, context)) } - if (actual instanceof PySelfType selfType && context.reversedSubstitutions) { - return Optional.of(match(selfType, expected, context)); + if (actual is PySelfType && context.reversedSubstitutions) { + return Optional.of(match(actual, expected, context)) } - if (expected instanceof PyParamSpecType paramSpecType) { - return Optional.of(match(paramSpecType, actual, context)); + if (expected is PyParamSpecType) { + return Optional.of(match(expected, actual, context)) } - if (expected instanceof PyConcatenateType concatenateType) { - return Optional.of(match(concatenateType, actual, context)); + if (expected is PyConcatenateType) { + return Optional.of(match(expected, actual, context)) } if (expected == null || actual == null || isUnknown(actual, context.context)) { - return Optional.of(true); + return Optional.of(true) } - if (expected instanceof PyCallableParameterListType callableParameterListType) { - return Optional.of(match(callableParameterListType, actual, context)); + if (expected is PyCallableParameterListType) { + return Optional.of(match(expected, actual, context)) } - if (actual instanceof PyNeverType) { - return Optional.of(true); + if (actual is PyNeverType) { + return Optional.of(true) } - if (expected instanceof PyNeverType) { - return Optional.of(false); + if (expected is PyNeverType) { + return Optional.of(false) } - if (actual instanceof PyUnionType) { - return Optional.of(match(expected, (PyUnionType)actual, context)); + if (actual is PyUnionType) { + return Optional.of(match(expected, actual, context)) } - if (expected instanceof PyUnionType) { - return Optional.of(match((PyUnionType)expected, actual, context)); + if (expected is PyUnionType) { + return Optional.of(match(expected, actual, context)) } - if (actual instanceof PyUnsafeUnionType weakUnionType) { - return Optional.of(match(expected, weakUnionType, context)); + if (actual is PyUnsafeUnionType) { + return Optional.of(match(expected, actual, context)) } - if (expected instanceof PyUnsafeUnionType weakUnionType) { - return Optional.of(match(weakUnionType, actual, context)); + if (expected is PyUnsafeUnionType) { + return Optional.of(match(expected, actual, context)) } - if (actual instanceof PyIntersectionType intersectionType) { - return Optional.of(match(expected, intersectionType, context)); + if (actual is PyIntersectionType) { + return Optional.of(match(expected, actual, context)) } - if (expected instanceof PyIntersectionType intersectionType) { - return Optional.of(match(intersectionType, actual, context)); + if (expected is PyIntersectionType) { + return Optional.of(match(expected, actual, context)) } - if (expected instanceof PyClassType && actual instanceof PyClassType) { - Optional match = match((PyClassType)expected, (PyClassType)actual, context); - if (match.isPresent()) { - return match; + if (expected is PyClassType && actual is PyClassType) { + val match = match(expected, actual, context) + if (match.isPresent) { + return match } } - if (actual instanceof PyStructuralType && ((PyStructuralType)actual).isInferredFromUsages()) { - return Optional.of(true); + if (actual is PyStructuralType && actual.isInferredFromUsages) { + return Optional.of(true) } - if (expected instanceof PyStructuralType) { - return Optional.of(match((PyStructuralType)expected, actual, context.context)); + if (expected is PyStructuralType) { + return Optional.of(match(expected, actual, context.context)) } - if (actual instanceof PyStructuralType && expected instanceof PyClassType) { - final Set expectedAttributes = ((PyClassType)expected).getMemberNames(true, context.context); - return Optional.of(expectedAttributes.containsAll(((PyStructuralType)actual).getAttributeNames())); + if (actual is PyStructuralType && expected is PyClassType) { + val expectedAttributes = expected.getMemberNames(true, context.context) + return Optional.of(expectedAttributes.containsAll(actual.attributeNames)) } - if (actual instanceof PyCallableType actualCallable && expected instanceof PyCallableType expectedCallable) { - final Optional match = match(expectedCallable, actualCallable, context); - if (match.isPresent()) { - return match; + if (actual is PyCallableType && expected is PyCallableType) { + val match = match(expected, actual, context) + if (match.isPresent) { + return match } } - if (expected instanceof PyModuleType) { - return Optional.of(actual instanceof PyModuleType && ((PyModuleType)expected).getModule() == ((PyModuleType)actual).getModule()); + if (expected is PyModuleType) { + return Optional.of(actual is PyModuleType && expected.module === actual.module) } - if (expected instanceof PyClassType classType && actual instanceof PyModuleType moduleType) { - if (PyProtocolsKt.isProtocol(classType, context.context)) { - return Optional.of(match(classType, moduleType, context)); + if (expected is PyClassType && actual is PyModuleType) { + if (expected.isProtocol(context.context)) { + return Optional.of(match(expected, actual, context)) } - return match(expected, moduleType.getModuleClassType(), context); + return match(expected, actual.moduleClassType, context) } - return Optional.of(matchNumericTypes(expected, actual)); + return Optional.of(matchNumericTypes(expected, actual)) } - private static @NotNull Optional match(PyNarrowedType expectedNarrowedType, - PyNarrowedType actualNarrowedType, - @NotNull MatchContext context) { - if (expectedNarrowedType.getTypeIs() != actualNarrowedType.getTypeIs()) { - return Optional.of(false); + private fun match( + expectedNarrowedType: PyNarrowedType, + actualNarrowedType: PyNarrowedType, + context: MatchContext, + ): Optional { + if (expectedNarrowedType.typeIs != actualNarrowedType.typeIs) { + return Optional.of(false) } - if (expectedNarrowedType.getTypeIs()) { - return Optional.of(Objects.equals(expectedNarrowedType.getNarrowedType(), actualNarrowedType.getNarrowedType())); + if (expectedNarrowedType.typeIs) { + return Optional.of(expectedNarrowedType.narrowedType == actualNarrowedType.narrowedType) } - return match(expectedNarrowedType.getNarrowedType(), actualNarrowedType.getNarrowedType(), context); + return match(expectedNarrowedType.narrowedType, actualNarrowedType.narrowedType, context) } /** - * Check whether {@code expected} is Python *object* or *type*. - *

+ * Check whether `expected` is Python *object* or *type*. + * + * * {@see PyTypeChecker#match(PyType, PyType, TypeEvalContext, Map)} */ - private static @NotNull Optional matchObject(@NotNull PyClassType expected, @Nullable PyType actual) { - if (ArrayUtil.contains(expected.getName(), PyNames.OBJECT, PyNames.TYPE)) { - final PyBuiltinCache builtinCache = PyBuiltinCache.getInstance(expected.getPyClass()); - if (expected.equals(builtinCache.getObjectType())) { - return Optional.of(true); + private fun matchObject(expected: PyClassType, actual: PyType?): Optional { + if (ArrayUtil.contains(expected.name, PyNames.OBJECT, PyNames.TYPE)) { + val builtinCache = getInstance(expected.pyClass) + if (expected == builtinCache.objectType) { + return Optional.of(true) } - if (expected.equals(builtinCache.getTypeType()) && - actual instanceof PyInstantiableType && ((PyInstantiableType)actual).isDefinition()) { - return Optional.of(true); + if (expected == builtinCache.typeType && + actual is PyInstantiableType<*> && actual.isDefinition + ) { + return Optional.of(true) } } - return Optional.empty(); + return Optional.empty() } /** - * Match {@code actual} versus {@link PyTypeVarType} expected. - *

- * The method mutates {@code context.substitutions} map adding new entries into it + * Match `actual` versus [PyTypeVarType] expected. + * + * + * The method mutates `context.substitutions` map adding new entries into it */ - private static boolean match(@NotNull PyTypeVarType expected, @Nullable PyType actual, @NotNull MatchContext context) { - if (expected.isDefinition() && actual instanceof PyInstantiableType && !((PyInstantiableType)actual).isDefinition()) { - return false; + private fun match(expected: PyTypeVarType, actual: PyType?, context: MatchContext): Boolean { + if (expected.isDefinition && actual is PyInstantiableType<*> && !actual.isDefinition) { + return false } - Ref substitutedRef = context.mySubstitutions.typeVars.get(expected); - Ref defaultTypeRef = expected.getDefaultType(); + var substitutedRef = context.mySubstitutions.typeVars[expected] + val defaultTypeRef = expected.defaultType if (defaultTypeRef != null) { - PyType defaultType = defaultTypeRef.get(); + val defaultType: PyType? = defaultTypeRef.get() // Skip default substitution - if (defaultType != null && defaultType.equals(Ref.deref(substitutedRef))) { - substitutedRef = null; + if (defaultType != null && defaultType == Ref.deref(substitutedRef)) { + substitutedRef = null } } - PyType bound = expected.getBound(); - List<@Nullable PyType> constraints = expected.getConstraints(); + var bound = expected.bound + var constraints = expected.constraints // Promote int in Type[TypeVar('T', int)] to Type[int] before checking that bounds match - if (expected.isDefinition()) { - bound = toClass(bound); - constraints = ContainerUtil.map(constraints, PyTypeChecker::toClass); + if (expected.isDefinition) { + bound = toClass(bound) + constraints = constraints.map { toClass(it) } } // Remove value-specific components from the actual type to make it safe to propagate - PyType safeActual = constraints.isEmpty() && bound instanceof PyLiteralStringType ? actual : replaceLiteralStringWithStr(actual); + var safeActual = if (constraints.isEmpty() && bound is PyLiteralStringType) actual else replaceLiteralStringWithStr(actual) if (substitutedRef != null) { - final PyType substitution = substitutedRef.get(); - if (expected.equals(safeActual) || expected.equals(substitution)) { - return true; + val substitution = substitutedRef.get() + if (expected == safeActual || expected == substitution) { + return true } - MyMatchHelper matchHelper = new MyMatchHelper(expected, context); + val matchHelper = MyMatchHelper(expected, context) if (matchHelper.match(substitution, safeActual)) { - return true; + return true } - if (context.mySubstitutions.frozenTypeVars.contains(expected)) { - return false; + if (expected in context.mySubstitutions.getFrozenTypeVars(KeyImpl)) { + return false } if (!matchHelper.match(safeActual, substitution)) { - safeActual = PyUnionType.union(safeActual, substitution); + safeActual = PyUnionType.union(safeActual, substitution) } } - int matchedConstraintIndex = -1; + var matchedConstraintIndex = -1 if (constraints.isEmpty()) { - Optional match = match(bound, safeActual, context); - if (match.isPresent() && !match.get()) { - return false; + val match = match(bound, safeActual, context) + if (match.isPresent && !match.get()) { + return false } } else { - final PyType finalSafeActual = safeActual; - matchedConstraintIndex = ContainerUtil.indexOf(constraints, constraint -> match(constraint, finalSafeActual, context).orElse(true)); + val finalSafeActual = safeActual + matchedConstraintIndex = + constraints.indexOfFirst { constraint -> match(constraint, finalSafeActual, context).orElse(true)!! } if (matchedConstraintIndex == -1) { - return false; + return false } } if (safeActual != null) { - PyType type = constraints.isEmpty() ? safeActual : constraints.get(matchedConstraintIndex); - context.mySubstitutions.typeVars.put(expected, Ref.create(type)); + val type = if (constraints.isEmpty()) safeActual else constraints[matchedConstraintIndex] + context.mySubstitutions.putTypeVar(expected, Ref(type), KeyImpl) } else { - PyType effectiveBound = PyTypeUtil.getEffectiveBound(expected); + val effectiveBound = PyTypeUtil.getEffectiveBound(expected) if (effectiveBound != null) { - context.mySubstitutions.typeVars.put(expected, Ref.create(PyUnionType.createWeakType(effectiveBound))); + context.mySubstitutions.putTypeVar(expected, Ref(PyUnionType.createWeakType(effectiveBound)), KeyImpl) } } - return true; + return true } - private static final class MyMatchHelper { - private final @NotNull Object myKey; - private final @NotNull MatchContext myContext; - - MyMatchHelper(@NotNull Object key, @NotNull MatchContext context) { - myKey = key; - myContext = context; - } - - boolean match(@Nullable PyType expected, @Nullable PyType actual) { - Optional result = RecursionManager.doPreventingRecursion( - myKey, false, myContext.reversedSubstitutions - ? () -> PyTypeChecker.match(actual, expected, myContext) - : () -> PyTypeChecker.match(expected, actual, myContext) - ); - return result != null && result.orElse(false); + private fun match(expected: PySelfType, actual: PyType?, context: MatchContext): Boolean { + if (actual == null) return true + val substitution = context.mySubstitutions.qualifierType + if (substitution != null && substitution !is PySelfType) { + return match(substitution, actual, context).orElse(false)!! } + if (actual !is PySelfType) return false + return expected.isDefinition == actual.isDefinition && + match(expected.scopeClassType, actual.scopeClassType, context).orElse(false)!! } - 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 (!(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) { + private fun toClass(type: PyType?): PyType? { return PyTypeUtil.toStream(type) - .map(t -> t instanceof PyInstantiableType instantiableType ? instantiableType.toClass() : t) - .collect(PyTypeUtil.toUnion(type)); + .map { t: PyType? -> if (t is PyInstantiableType<*>) t.toClass() else t } + .collect(PyTypeUtil.toUnion(type)) } - private static boolean match(@NotNull PyPositionalVariadicType expected, @Nullable PyType actual, @NotNull MatchContext context) { + private fun match(expected: PyPositionalVariadicType, actual: PyType?, context: MatchContext): Boolean { if (actual == null) { - return true; + return true } - if (!(actual instanceof PyPositionalVariadicType actualVariadic)) { - return false; + if (actual !is PyPositionalVariadicType) { + return false } - if (expected instanceof PyUnpackedTupleType expectedUnpackedTupleType) { + if (expected is PyUnpackedTupleType) { // The actual type is just a TypeVarTuple - if (!(actualVariadic instanceof PyUnpackedTupleType actualUnpackedTupleType)) { - return false; + if (actual !is PyUnpackedTupleType) { + return false } - if (expectedUnpackedTupleType.isUnbound()) { - PyType repeatedExpectedType = expectedUnpackedTupleType.getElementTypes().get(0); - if (actualUnpackedTupleType.isUnbound()) { - return match(repeatedExpectedType, actualUnpackedTupleType.getElementTypes().get(0), context).orElse(false); + if (expected.isUnbound) { + val repeatedExpectedType = expected.elementTypes[0] + if (actual.isUnbound) { + return match(repeatedExpectedType, actual.elementTypes[0], context).orElse(false)!! } else { - return ContainerUtil.all(actualUnpackedTupleType.getElementTypes(), - singleActualType -> match(repeatedExpectedType, singleActualType, context).orElse(false)); + return actual.elementTypes + .all { singleActualType -> match(repeatedExpectedType, singleActualType, context).orElse(false)!! } } } else { - if (actualUnpackedTupleType.isUnbound()) { - PyType repeatedActualType = actualUnpackedTupleType.getElementTypes().get(0); - return ContainerUtil.all(expectedUnpackedTupleType.getElementTypes(), - singleExpectedType -> match(singleExpectedType, repeatedActualType, context).orElse(false)); + if (actual.isUnbound) { + val repeatedActualType: PyType? = actual.elementTypes[0] + return expected.elementTypes + .all { singleExpectedType -> match(singleExpectedType, repeatedActualType, context).orElse(false)!! } } else { - return matchTypeParameters(expectedUnpackedTupleType.getElementTypes(), actualUnpackedTupleType.getElementTypes(), context); + return matchTypeParameters(expected.elementTypes, actual.elementTypes, context) } } } - // The expected type is just a TypeVarTuple else { - PyPositionalVariadicType substitution = context.mySubstitutions.typeVarTuples.get(expected); - if (substitution != null && !substitution.equals(PyUnpackedTupleTypeImpl.UNSPECIFIED)) { - if (expected.equals(actual) || substitution.equals(expected)) { - return true; + val substitution = context.mySubstitutions.typeVarTuples[expected] + if (substitution != null && substitution != PyUnpackedTupleTypeImpl.UNSPECIFIED) { + if (expected == actual || substitution == expected) { + return true } - return context.reversedSubstitutions ? match(actualVariadic, substitution, context) : match(substitution, actualVariadic, context); + return if (context.reversedSubstitutions) match(actual, substitution, context) + else match( + substitution, + actual, + context + ) } - context.mySubstitutions.typeVarTuples.put((PyTypeVarTupleType)expected, actualVariadic); + context.mySubstitutions.putTypeVarTuple(expected as PyTypeVarTupleType, actual, KeyImpl) } - return true; + return true } - private static @Nullable PyType replaceLiteralStringWithStr(@Nullable PyType actual) { + private fun replaceLiteralStringWithStr(actual: PyType?): PyType? { // TODO replace with PyTypeVisitor API once it's ready - if (actual instanceof PyLiteralStringType literalStringType) { - return new PyClassTypeImpl(literalStringType.getPyClass(), false); + if (actual is PyLiteralStringType) { + return PyClassTypeImpl(actual.pyClass, false) } - if (actual instanceof PyUnionType unionType) { - return unionType.map(PyTypeChecker::replaceLiteralStringWithStr); + if (actual is PyUnionType) { + return actual.map { replaceLiteralStringWithStr(it) } } - if (actual instanceof PyNamedTupleType) { - return actual; + if (actual is PyNamedTupleType) { + return actual } - if (actual instanceof PyTupleType tupleType) { - return new PyTupleType(tupleType.getPyClass(), - ContainerUtil.map(tupleType.getElementTypes(), PyTypeChecker::replaceLiteralStringWithStr), - tupleType.isHomogeneous(), tupleType.isDefinition()); + if (actual is PyTupleType) { + return PyTupleType( + actual.pyClass, + actual.elementTypes.map { replaceLiteralStringWithStr(it) }, + actual.isHomogeneous, actual.isDefinition + ) } - if (actual instanceof PyCollectionType generic) { - return new PyCollectionTypeImpl(generic.getPyClass(), generic.isDefinition(), - ContainerUtil.map(generic.getElementTypes(), PyTypeChecker::replaceLiteralStringWithStr)); + if (actual is PyCollectionType) { + return PyCollectionTypeImpl( + actual.pyClass, actual.isDefinition, + actual.elementTypes.map { replaceLiteralStringWithStr(it) } + ) } - return actual; + return actual } - private static boolean match(@NotNull PyParamSpecType expected, @Nullable PyType actual, @NotNull MatchContext context) { - if (actual == null) return true; - if (!(actual instanceof PyCallableParameterVariadicType actualParameters)) return false; - context.mySubstitutions.paramSpecs.put(expected, actualParameters); - return true; + private fun match(expected: PyParamSpecType, actual: PyType?, context: MatchContext): Boolean { + if (actual == null) return true + if (actual !is PyCallableParameterVariadicType) return false + context.mySubstitutions.putParamSpec(expected, actual, KeyImpl) + return true } - private static boolean match(@NotNull PyConcatenateType expected, @Nullable PyType actual, @NotNull MatchContext context) { - if (actual == null) return true; - List expectedFirstTypes = expected.getFirstTypes(); - int expectedPrefixSize = expectedFirstTypes.size(); - PyParamSpecType expectedParamSpec = expected.getParamSpec(); - if (actual instanceof PyConcatenateType actualConcatenateType) { - if (expectedPrefixSize > actualConcatenateType.getFirstTypes().size()) { - return false; + private fun match(expected: PyConcatenateType, actual: PyType?, context: MatchContext): Boolean { + if (actual == null) return true + val expectedFirstTypes = expected.firstTypes + val expectedPrefixSize = expectedFirstTypes.size + val expectedParamSpec = expected.paramSpec + if (actual is PyConcatenateType) { + if (expectedPrefixSize > actual.firstTypes.size) { + return false } - List actualFirstTypes = actualConcatenateType.getFirstTypes().subList(0, expectedPrefixSize); + val actualFirstTypes = actual.firstTypes.subList(0, expectedPrefixSize) if (!match(expectedFirstTypes, actualFirstTypes, context)) { - return false; + return false } if (expectedParamSpec == null) { - return true; + return true } - if (actualFirstTypes.size() > expectedPrefixSize) { - return match(expectedParamSpec, - new PyConcatenateType(ContainerUtil.subList(expectedFirstTypes, actualFirstTypes.size()), - actualConcatenateType.getParamSpec()), context); + if (actualFirstTypes.size > expectedPrefixSize) { + return match( + expectedParamSpec, + PyConcatenateType( + expectedFirstTypes.subList(0, actualFirstTypes.size), + actual.paramSpec + ), context + ) } else { - return match(expectedParamSpec, actualConcatenateType.getParamSpec(), context); + return match(expectedParamSpec, actual.paramSpec, context) } } - else if (actual instanceof PyCallableParameterListType actualParameters) { - if (expectedPrefixSize > actualParameters.getParameters().size()) { - return false; + else if (actual is PyCallableParameterListType) { + if (expectedPrefixSize > actual.parameters.size) { + return false } - List actualFirstParamTypes = ContainerUtil.map(actualParameters.getParameters().subList(0, expectedPrefixSize), - it -> it.getType(context.context)); + val actualFirstParamTypes = + actual.parameters.subList(0, expectedPrefixSize).map { it.getType(context.context) } if (!match(expectedFirstTypes, actualFirstParamTypes, context)) { - return false; + return false } if (expectedParamSpec == null) { - return true; + return true } - return match(expectedParamSpec, - new PyCallableParameterListTypeImpl(ContainerUtil.subList(actualParameters.getParameters(), expectedPrefixSize)), - context); + return match( + expectedParamSpec, + PyCallableParameterListTypeImpl(ContainerUtil.subList(actual.parameters, expectedPrefixSize)), + context + ) } - return false; + return false } - 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 match(expectedParameters, actualParameters, context); + private fun match( + expectedParameters: PyCallableParameterListType, + actual: PyType?, + context: MatchContext, + ): Boolean { + if (actual == null) return true + if (actual !is PyCallableParameterListType) return false + return match(expectedParameters, actual, context) } - private static boolean match(@NotNull PyType expected, @NotNull PyUnionType actual, @NotNull MatchContext context) { - if (expected instanceof PyTupleType expectedTupleType) { + private fun match(expected: PyType, actual: PyUnionType, context: MatchContext): Boolean { + if (expected is PyTupleType) { // XXX A type-widening hack for cases like PyTypeTest.testDictFromTuple - Optional match = match(expectedTupleType, actual, context); - if (match.isPresent()) { - return match.get(); + val match = match(expected, actual, context) + if (match.isPresent) { + return match.get() } } - if (!PyUnionType.isStrictSemanticsEnabled()) {// checking strictly separately until PY-24834 gets implemented - if (ContainerUtil.exists(actual.getMembers(), x -> x instanceof PyLiteralStringType || x instanceof PyLiteralType)) { - return ContainerUtil.and(actual.getMembers(), type -> match(expected, type, context).orElse(false)); + if (!PyUnionType.isStrictSemanticsEnabled()) { // checking strictly separately until PY-24834 gets implemented + if (actual.members.any { it is PyLiteralStringType || it is PyLiteralType }) { + return actual.members.all { type -> match(expected, type, context).orElse(false)!! } } - return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false)); + return actual.members.any { type -> match(expected, type, context).orElse(false)!! } } - return ContainerUtil.and(actual.getMembers(), type -> match(expected, type, context).orElse(false)); + return actual.members.all { type -> match(expected, type, context).orElse(false)!! } } - private static @NotNull Optional match(@NotNull PyTupleType expected, - @NotNull PyUnionType actual, - @NotNull MatchContext context) { - final int elementCount = expected.getElementCount(); + private fun match( + expected: PyTupleType, + actual: PyUnionType, + context: MatchContext, + ): Optional { + val elementCount = expected.elementCount - if (!expected.isHomogeneous()) { - PyTupleType widenedActual = widenUnionOfTuplesToTupleOfUnions(actual, elementCount); + if (!expected.isHomogeneous) { + val widenedActual = widenUnionOfTuplesToTupleOfUnions(actual, elementCount) if (widenedActual != null) { - return match(expected, widenedActual, context); + return match(expected, widenedActual, context) } } - return Optional.empty(); + return Optional.empty() } - private static boolean match(@NotNull PyUnionType expected, @NotNull PyType actual, @NotNull MatchContext context) { - if (expected.getMembers().contains(actual)) { - return true; + private fun match(expected: PyUnionType, actual: PyType, context: MatchContext): Boolean { + if (actual in expected.members) { + return true } - return ContainerUtil.or(expected.getMembers(), type -> match(type, actual, context).orElse(true)); + return expected.members.any { type: PyType? -> match(type, actual, context).orElse(true)!! } } - private static boolean match(@NotNull PyType expected, @NotNull PyIntersectionType actual, @NotNull MatchContext context) { - return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false)); + private fun match(expected: PyType, actual: PyIntersectionType, context: MatchContext): Boolean { + return actual.members.any { type: PyType? -> match(expected, type, context).orElse(false)!! } } - private static boolean match(@NotNull PyIntersectionType expected, @NotNull PyType actual, @NotNull MatchContext context) { - return ContainerUtil.all(expected.getMembers(), type -> match(type, actual, context).orElse(true)); + private fun match(expected: PyIntersectionType, actual: PyType, context: MatchContext): Boolean { + return expected.members.all { type: PyType? -> match(type, actual, context).orElse(true)!! } } - private static boolean match(@NotNull PyType expected, @NotNull PyUnsafeUnionType actual, @NotNull MatchContext context) { - return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false)); + private fun match(expected: PyType, actual: PyUnsafeUnionType, context: MatchContext): Boolean { + return actual.members.any { type: PyType? -> match(expected, type, context).orElse(false)!! } } - private static boolean match(@NotNull PyUnsafeUnionType expected, @NotNull PyType actual, @NotNull MatchContext context) { - if (expected.getMembers().contains(actual)) { - return true; + private fun match(expected: PyUnsafeUnionType, actual: PyType, context: MatchContext): Boolean { + if (actual in expected.members) { + return true } - return ContainerUtil.or(expected.getMembers(), type -> match(type, actual, context).orElse(true)); + return expected.members.any { type: PyType? -> match(type, actual, context).orElse(true)!! } } - private static @NotNull Optional match(@NotNull PyClassType expected, - @NotNull PyClassType actual, - @NotNull MatchContext matchContext) { - if (expected.equals(actual)) { - return Optional.of(true); + private fun match( + expected: PyClassType, + actual: PyClassType, + matchContext: MatchContext, + ): Optional { + if (expected == actual) { + return Optional.of(true) } - final TypeEvalContext context = matchContext.context; + val context = matchContext.context - if (expected.isDefinition() ^ actual.isDefinition() && !PyProtocolsKt.isProtocol(expected, context)) { - if (!expected.isDefinition() && actual.isDefinition()) { - final PyClassLikeType metaClass = actual.getMetaClassType(context, true); - return Optional.of(metaClass != null && match((PyType)expected, metaClass.toInstance(), matchContext).orElse(true)); + if (expected.isDefinition xor actual.isDefinition && !expected.isProtocol(context)) { + if (!expected.isDefinition && actual.isDefinition) { + val metaClass = actual.getMetaClassType(context, true) + return Optional.of( + metaClass != null && match(expected as PyType, metaClass.toInstance(), matchContext).orElse( + true + )!! + ) } - return Optional.of(false); + return Optional.of(false) } - if (expected instanceof PyTupleType && actual instanceof PyTupleType) { - return match((PyTupleType)expected, (PyTupleType)actual, matchContext); + if (expected is PyTupleType && actual is PyTupleType) { + return match(expected, actual, matchContext) } - if (expected instanceof PyLiteralType) { - return Optional.of(actual instanceof PyLiteralType && PyLiteralType.Companion.match((PyLiteralType)expected, (PyLiteralType)actual)); + if (expected is PyLiteralType) { + return Optional.of(actual is PyLiteralType && match(expected, actual)) } - if (expected instanceof PyTypedDictType && !(actual instanceof PyTypedDictType)) { - return Optional.of(false); + if (expected is PyTypedDictType && actual !is PyTypedDictType) { + return Optional.of(false) } - if (actual instanceof PyTypedDictType typedDictType) { - final Boolean matchResult = PyTypedDictType.match(expected, typedDictType, context); - if (matchResult != null) return Optional.of(matchResult); + if (actual is PyTypedDictType) { + val matchResult = PyTypedDictType.match(expected, actual, context) + if (matchResult != null) return Optional.of(matchResult) } - if (expected instanceof PyLiteralStringType) { - return Optional.of(PyLiteralStringType.Companion.match((PyLiteralStringType)expected, actual)); + if (expected is PyLiteralStringType) { + return Optional.of(match(expected, actual)) } - final PyClass superClass = expected.getPyClass(); - final PyClass subClass = actual.getPyClass(); + val superClass = expected.pyClass + val subClass = actual.pyClass - if (!subClass.isSubclass(superClass, context) && PyProtocolsKt.isProtocol(expected, context)) { - return Optional.of(matchProtocols(expected, actual, matchContext)); + if (!subClass.isSubclass(superClass, context) && expected.isProtocol(context)) { + return Optional.of(matchProtocols(expected, actual, matchContext)) } - if (expected instanceof PyCollectionType) { - return Optional.of(match((PyCollectionType)expected, actual, matchContext)); + if (expected is PyCollectionType) { + return Optional.of(match(expected, actual, matchContext)) } if (matchClasses(superClass, subClass, context)) { - if (expected instanceof PyTypingNewType && !expected.equals(actual) && superClass.equals(subClass)) { - return Optional.of(actual.getAncestorTypes(context).contains(expected)); + if (expected is PyTypingNewType && (expected != actual) && superClass == subClass) { + return Optional.of(expected in actual.getAncestorTypes(context)) } - return Optional.of(true); + return Optional.of(true) } - return Optional.empty(); + return Optional.empty() } - private static boolean matchProtocols(@NotNull PyClassType expected, @NotNull PyClassType actual, @NotNull MatchContext matchContext) { - GenericSubstitutions expectedSubstitutions = collectTypeSubstitutions(expected, matchContext.context); - GenericSubstitutions actualSubstitutions = collectTypeSubstitutions(actual, matchContext.context); + private fun matchProtocols(expected: PyClassType, actual: PyClassType, matchContext: MatchContext): Boolean { + val expectedSubstitutions = collectTypeSubstitutions(expected, matchContext.context) + val actualSubstitutions = collectTypeSubstitutions(actual, matchContext.context) // See https://typing.python.org/en/latest/spec/generics.html#use-in-protocols @@ -671,374 +665,383 @@ public final class PyTypeChecker { // 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(); - final List subclassElementMembers = pair.getSecond(); + val protocolSubstitutions = GenericSubstitutions() + protocolSubstitutions.qualifierType = actual.toInstance() + val protocolContext = + MatchContext(matchContext.context, protocolSubstitutions, matchContext.reversedSubstitutions) + for (pair in inspectProtocolSubclass( + expected, actual, + matchContext.context + )) { + val protocolMember = pair.first + val subclassElementMembers = pair.second if (ContainerUtil.isEmpty(subclassElementMembers)) { - return false; + return false } - final PyType rawProtocolElementType = - dropSelfIfNeeded(expected, protocolMember.getType(), matchContext.context); + val rawProtocolElementType = + dropSelfIfNeeded(expected, protocolMember.type, matchContext.context) - final PyType protocolElementType = substitute(rawProtocolElementType, expectedSubstitutions, matchContext.context); - final boolean elementResult = ContainerUtil.exists(subclassElementMembers, subclassElementMember -> { - if (protocolMember.isWritable() && !subclassElementMember.isWritable()) { - return false; - } - if (protocolMember.isDeletable() && !subclassElementMember.isDeletable()) { - return false; - } - boolean isProtocolMemberClassVar = protocolMember.isClassVar(); - boolean isSubclassMemberClassVar = subclassElementMember.isClassVar(); - if (isSubclassMemberClassVar != isProtocolMemberClassVar) { - return false; - } + val protocolElementType = substitute(rawProtocolElementType, expectedSubstitutions, matchContext.context) + val elementResult: Boolean = + subclassElementMembers.any { subclassElementMember: PyTypeMember? -> + if (protocolMember.isWritable && !subclassElementMember!!.isWritable) { + return@any false + } + if (protocolMember.isDeletable && !subclassElementMember!!.isDeletable) { + return@any false + } + val isProtocolMemberClassVar = protocolMember.isClassVar + val isSubclassMemberClassVar = subclassElementMember!!.isClassVar + if (isSubclassMemberClassVar != isProtocolMemberClassVar) { + return@any false + } - PyType subclassElementType = dropSelfIfNeeded(actual, subclassElementMember.getType(), matchContext.context); - subclassElementType = substitute(subclassElementType, actualSubstitutions, matchContext.context); - - return match(protocolElementType, subclassElementType, protocolContext).orElse(true); - }); + var subclassElementType = dropSelfIfNeeded(actual, subclassElementMember.type, matchContext.context) + subclassElementType = substitute(subclassElementType, actualSubstitutions, matchContext.context) + match(protocolElementType, subclassElementType, protocolContext).orElse(true)!! + } if (!elementResult) { - return false; + return false } } - if (expected instanceof PyCollectionType genericExpected && hasGenerics(expected, matchContext.context)) { - PyCollectionType concreteExpected = - (PyCollectionType)substitute(expected, protocolContext.mySubstitutions, protocolContext.context); - assert concreteExpected != null; - // This match is supposed to succeed since all protocol members were compatible. + if (expected is PyCollectionType && expected.hasGenerics(matchContext.context)) { + val concreteExpected: PyCollectionType = checkNotNull( + substitute(expected, protocolContext.mySubstitutions, protocolContext.context) as PyCollectionType? + ) + // This match is supposed to succeed since all protocol members were compatible. // Effectively, we're just copying substitutions for the terminal type parameters e.g. - // extracting: T@func -> str + // extracting: T@func -> str // from: Iterable[T@func] <- Iterable[str] - return matchGenericClassesParameterWise(genericExpected, concreteExpected, matchContext); + return matchGenericClassesParameterWise(expected, concreteExpected, matchContext) } - return true; + return true } - private static boolean match(PyClassType expectedProtocol, PyModuleType actualModule, MatchContext matchContext) { - PyFile module = actualModule.getModule(); + private fun match(expectedProtocol: PyClassType, actualModule: PyModuleType, matchContext: MatchContext): Boolean { + val module = actualModule.module - 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) { - PsiElement pm = pair.getFirst().getElement(); - if (!(pm instanceof PsiNamedElement protocolMember)) { - continue; - } - String name = protocolMember.getName(); - PyTypedElement moduleElement = moduleElements.get(name); - if (moduleElement != null) { - PyType expectedProtocolMemberType = - substitute(dropSelfIfNeeded(expectedProtocol, pair.getFirst().getType(), matchContext.context), substitutions, - matchContext.context); - PyType actualModuleElementType = matchContext.context.getType(moduleElement); - if (!match(expectedProtocolMemberType, actualModuleElementType, matchContext.context)) { - return false; + val moduleElements = + (module.topLevelAttributes + module.topLevelFunctions) + .asSequence() + .filter { e: PsiNameIdentifierOwner? -> + val name = (e as PyQualifiedNameOwner).name + name != null && !isPrivate(name) && !isProtected(name) } - continue; + .associateBy { it.name } + + val protocolElements = + inspectProtocolSubclass(expectedProtocol, expectedProtocol, matchContext.context) + + if (protocolElements.size != moduleElements.size) return false + + val substitutions = collectTypeSubstitutions(expectedProtocol, matchContext.context) + for (pair in protocolElements) { + val pm = pair.first.element + if (pm !is PsiNamedElement) { + continue } - return false; + val name = pm.name + val moduleElement = moduleElements[name] + if (moduleElement != null) { + val expectedProtocolMemberType = + substitute( + dropSelfIfNeeded(expectedProtocol, pair.first.type, matchContext.context), substitutions, + matchContext.context + ) + val actualModuleElementType = matchContext.context.getType(moduleElement) + if (!match(expectedProtocolMemberType, actualModuleElementType, matchContext.context)) { + return false + } + continue + } + return false } - return true; + return true } - private static @Nullable PyType dropSelfIfNeeded(@NotNull PyClassType classType, - @Nullable PyType elementType, - @NotNull TypeEvalContext context) { - if (elementType instanceof PyCallableType callableType) { - if (PyUtil.isInitOrNewMethod(callableType.getCallable()) || !classType.isDefinition()) { - return callableType.dropSelf(context); + private fun dropSelfIfNeeded( + classType: PyClassType, + elementType: PyType?, + context: TypeEvalContext, + ): PyType? { + if (elementType is PyCallableType) { + if (PyUtil.isInitOrNewMethod(elementType.callable) || !classType.isDefinition) { + return elementType.dropSelf(context) } } - return elementType; + return elementType } // 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) { - if (actual.isHomogeneous()) { + private fun match( + expected: PyTupleType, + actual: PyTupleType, + context: MatchContext, + ): Optional { + if (actual.isHomogeneous) { // The type tuple[Any, ...] is consistent with any tuple - final PyType elementType = actual.getIteratedItemType(); + val elementType = actual.iteratedItemType if (elementType == null) { - return Optional.of(true); + return Optional.of(true) } } // TODO Delegate to UnpackedTupleType here, once it's introduced - if (!expected.isHomogeneous() && !actual.isHomogeneous()) { - Condition isUnbound = - type -> (type instanceof PyUnpackedTupleType unpacked && unpacked.isUnbound()) || type instanceof PyTypeVarTupleType; - boolean expectedContainsUnboundTypes = - ContainerUtil.exists(expected.getElementTypes(), isUnbound); - boolean actualContainsUnboundTypes = - ContainerUtil.exists(actual.getElementTypes(), isUnbound); + if (!expected.isHomogeneous && !actual.isHomogeneous) { + val isUnbound = + { type: PyType? -> (type is PyUnpackedTupleType && type.isUnbound) || type is PyTypeVarTupleType } + val expectedContainsUnboundTypes = + expected.elementTypes.any(isUnbound) + val actualContainsUnboundTypes = + actual.elementTypes.any(isUnbound) if (actualContainsUnboundTypes && !expectedContainsUnboundTypes) { - return Optional.of(false); + return Optional.of(false) } - return Optional.of(matchTypeParameters(expected.getElementTypes(), actual.getElementTypes(), context)); + return Optional.of(matchTypeParameters(expected.elementTypes, actual.elementTypes, context)) } - if (expected.isHomogeneous() && !actual.isHomogeneous()) { - final PyType expectedElementType = expected.getIteratedItemType(); + if (expected.isHomogeneous && !actual.isHomogeneous) { + val expectedElementType = expected.iteratedItemType return Optional.of( - PyTypeUtil.toStream(actual.getIteratedItemType()).allMatch(type -> match(expectedElementType, type, context).orElse(true)) - ); + PyTypeUtil.toStream(actual.iteratedItemType) + .allMatch { type: PyType? -> match(expectedElementType, type, context).orElse(true)!! } + ) } - if (!expected.isHomogeneous() && actual.isHomogeneous()) { - return Optional.of(false); + if (!expected.isHomogeneous && actual.isHomogeneous) { + return Optional.of(false) } - return match(expected.getIteratedItemType(), actual.getIteratedItemType(), context); + return match(expected.iteratedItemType, actual.iteratedItemType, context) } - private static boolean match(@NotNull PyCollectionType expected, @NotNull PyClassType actual, @NotNull MatchContext context) { - if (actual instanceof PyTupleType) { - return match(expected, (PyTupleType)actual, context); + private fun match(expected: PyCollectionType, actual: PyClassType, context: MatchContext): Boolean { + if (actual is PyTupleType) { + return match(expected, actual, context) } - final PyClass superClass = expected.getPyClass(); - final PyClass subClass = actual.getPyClass(); + val superClass = expected.pyClass + val subClass = actual.pyClass - return matchClasses(superClass, subClass, context.context) && matchGenerics(expected, actual, context); + return matchClasses(superClass, subClass, context.context) && matchGenerics(expected, actual, context) } - private static boolean match(@NotNull PyCollectionType expected, @NotNull PyTupleType actual, @NotNull MatchContext context) { - if (!matchClasses(expected.getPyClass(), actual.getPyClass(), context.context)) { - return false; + private fun match(expected: PyCollectionType, actual: PyTupleType, context: MatchContext): Boolean { + if (!matchClasses(expected.pyClass, actual.pyClass, context.context)) { + return false } - final PyType superElementType = expected.getIteratedItemType(); - final PyType subElementType = actual.getIteratedItemType(); + val superElementType = expected.iteratedItemType + val subElementType = actual.iteratedItemType - return match(superElementType, subElementType, context).orElse(true); + return match(superElementType, subElementType, context).orElse(true)!! } - private static boolean match(@NotNull PyStructuralType expected, @NotNull PyType actual, @NotNull TypeEvalContext context) { - if (actual instanceof PyStructuralType) { - return match(expected, (PyStructuralType)actual); + private fun match(expected: PyStructuralType, actual: PyType, context: TypeEvalContext): Boolean { + if (actual is PyStructuralType) { + return match(expected, actual) } - if (actual instanceof PyClassType) { - return match(expected, (PyClassType)actual, context); + if (actual is PyClassType) { + return match(expected, actual, context) } - if (actual instanceof PyModuleType) { - final PyFile module = ((PyModuleType)actual).getModule(); - if (module.getLanguageLevel().isAtLeast(LanguageLevel.PYTHON37) && definesGetAttr(module, context)) { - return true; + if (actual is PyModuleType) { + val module = actual.module + if (module.languageLevel.isAtLeast(LanguageLevel.PYTHON37) && definesGetAttr(module, context)) { + return true } } - final PyResolveContext resolveContext = PyResolveContext.defaultContext(context); - return !ContainerUtil.exists(expected.getAttributeNames(), attribute -> ContainerUtil - .isEmpty(actual.resolveMember(attribute, null, AccessDirection.READ, resolveContext))); - } - - private static boolean match(@NotNull PyStructuralType expected, @NotNull PyStructuralType actual) { - if (expected.isInferredFromUsages()) { - return true; + val resolveContext = PyResolveContext.defaultContext(context) + return !expected.attributeNames.any { attribute: String? -> + ContainerUtil + .isEmpty(actual.resolveMember(attribute!!, null, AccessDirection.READ, resolveContext)) } - return expected.getAttributeNames().containsAll(actual.getAttributeNames()); } - private static boolean match(@NotNull PyStructuralType expected, @NotNull PyClassType actual, @NotNull TypeEvalContext context) { - if (overridesGetAttr(actual.getPyClass(), context)) { - return true; + private fun match(expected: PyStructuralType, actual: PyStructuralType): Boolean { + if (expected.isInferredFromUsages) { + return true } - final Set actualAttributes = actual.getMemberNames(true, context); - return actualAttributes.containsAll(expected.getAttributeNames()); + return expected.attributeNames.containsAll(actual.attributeNames) } - private static boolean match(@NotNull PyCallableParameterListType expectedParametersType, - @NotNull PyCallableParameterListType actualParametersType, - @NotNull MatchContext matchContext) { - TypeEvalContext context = matchContext.context; - List expectedParameters = expectedParametersType.getParameters(); - List actualParameters = actualParametersType.getParameters(); + private fun match(expected: PyStructuralType, actual: PyClassType, context: TypeEvalContext): Boolean { + if (overridesGetAttr(actual.pyClass, context)) { + return true + } + val actualAttributes = actual.getMemberNames(true, context) + return actualAttributes.containsAll(expected.attributeNames) + } - int startIndex = 0; + private fun match( + expectedParametersType: PyCallableParameterListType, + actualParametersType: PyCallableParameterListType, + matchContext: MatchContext, + ): Boolean { + val context = matchContext.context + val expectedParameters = expectedParametersType.parameters + val actualParameters = actualParametersType.parameters + + var startIndex = 0 if (!expectedParameters.isEmpty() && !actualParameters.isEmpty()) { - var firstExpectedParam = expectedParameters.getFirst(); - var firstActualParam = actualParameters.getFirst(); - if (firstExpectedParam.isSelf() && firstActualParam.isSelf()) { - if (!match(firstExpectedParam.getType(context), firstActualParam.getType(context), matchContext).orElse(true)) { - return false; + val firstExpectedParam: PyCallableParameter = expectedParameters.first() + val firstActualParam: PyCallableParameter = actualParameters.first() + if (firstExpectedParam.isSelf && firstActualParam.isSelf) { + if (!match(firstExpectedParam.getType(context), firstActualParam.getType(context), matchContext).orElse(true)!!) { + return false } - startIndex = 1; + startIndex = 1 } } - var mapping = PyCallableParameterMapping.mapCallableParameters(ContainerUtil.subList(expectedParameters, startIndex), - ContainerUtil.subList(actualParameters, startIndex), - context); + val mapping = mapCallableParameters( + ContainerUtil.subList(expectedParameters, startIndex), + ContainerUtil.subList(actualParameters, startIndex), + context + ) if (mapping == null) { - return false; + return false } // actual callable type could accept more general parameter type - for (Couple pair : mapping.getMappedTypes()) { - Optional matched = matchContext.reverseSubstitutions().reversedSubstitutions - ? match(pair.getSecond(), pair.getFirst(), matchContext.reverseSubstitutions()) - : match(pair.getFirst(), pair.getSecond(), matchContext.reverseSubstitutions()); - if (!matched.orElse(true)) { - return false; + for (pair in mapping.mappedTypes) { + val matched = if (matchContext.reverseSubstitutions().reversedSubstitutions) + match(pair.getSecond(), pair.getFirst(), matchContext.reverseSubstitutions()) + else + match(pair.getFirst(), pair.getSecond(), matchContext.reverseSubstitutions()) + if (!matched.orElse(true)!!) { + return false } } - return true; + return true } - private static @NotNull Optional match(@NotNull PyCallableType expected, - @NotNull PyCallableType actual, - @NotNull MatchContext matchContext) { - if (actual instanceof PyFunctionType && expected instanceof PyClassType && FUNCTION.equals(expected.getName()) - && expected.equals(PyBuiltinCache.getInstance(actual.getCallable()).getObjectType(FUNCTION))) { - return Optional.of(true); + private fun match( + expected: PyCallableType, + actual: PyCallableType, + matchContext: MatchContext, + ): Optional { + if (actual is PyFunctionType && expected is PyClassType && PyNames.FUNCTION == expected.name + && expected == getInstance(actual.callable).getObjectType(PyNames.FUNCTION) + ) { + return Optional.of(true) } - final TypeEvalContext context = matchContext.context; + val context = matchContext.context - if (expected instanceof PyClassLikeType && !isCallableProtocol((PyClassLikeType)expected, context)) { - return PyTypingTypeProvider.CALLABLE.equals(((PyClassLikeType)expected).getClassQName()) - ? Optional.of(actual.isCallable()) - : Optional.empty(); + if (expected is PyClassLikeType && !isCallableProtocol(expected, context)) { + return if (PyTypingTypeProvider.CALLABLE == expected.classQName) + Optional.of(actual.isCallable) + else + Optional.empty() } - if (expected.isCallable() && actual.isCallable()) { - final PyCallableParameterVariadicType expectedParametersType = expected.getParametersType(context); - final PyCallableParameterVariadicType actualParametersType = actual.getParametersType(context); + if (expected.isCallable && actual.isCallable) { + val expectedParametersType = expected.getParametersType(context) + val actualParametersType = actual.getParametersType(context) if (expectedParametersType != null && actualParametersType != null) { - if (!match(expectedParametersType, actualParametersType, matchContext).orElse(true)) { - return Optional.of(false); + if (!match(expectedParametersType, actualParametersType, matchContext).orElse(true)!!) { + return Optional.of(false) } } - if (!match(expected.getReturnType(context), getActualReturnType(actual, context), matchContext).orElse(true)) { - return Optional.of(false); + if (!match(expected.getReturnType(context), getActualReturnType(actual, context), matchContext).orElse(true)!!) { + return Optional.of(false) } - return Optional.of(true); + return Optional.of(true) } - return Optional.empty(); + return Optional.empty() } - private static @Nullable PyType getActualReturnType(@NotNull PyCallableType actual, @NotNull TypeEvalContext context) { - if (actual instanceof PyCallableTypeImpl) { - return actual.getReturnType(context); + private fun getActualReturnType(actual: PyCallableType, context: TypeEvalContext): PyType? { + if (actual is PyCallableTypeImpl) { + return actual.getReturnType(context) } - PyCallable callable = actual.getCallable(); - if (callable instanceof PyFunction) { - return getReturnTypeToAnalyzeAsCallType((PyFunction)callable, context); + val callable = actual.callable + if (callable is PyFunction) { + return PyUtil.getReturnTypeToAnalyzeAsCallType(callable, context) } - return actual.getReturnType(context); + return actual.getReturnType(context) } - private static boolean match(@NotNull List expected, @NotNull List actual, @NotNull MatchContext matchContext) { - if (expected.size() != actual.size()) return false; - for (int i = 0; i < expected.size(); ++i) { - if (!match(expected.get(i), actual.get(i), matchContext).orElse(true)) return false; + private fun match(expected: List, actual: List, matchContext: MatchContext): Boolean { + if (expected.size != actual.size) return false + for (i in expected.indices) { + if (!match(expected[i], actual[i], matchContext).orElse(true)!!) return false } - return true; + return true } - private static boolean isCallableProtocol(@NotNull PyClassLikeType expected, @NotNull TypeEvalContext context) { - return PyProtocolsKt.isProtocol(expected, context) && expected.getMemberNames(false, context).contains(PyNames.CALL); + private fun isCallableProtocol(expected: PyClassLikeType, context: TypeEvalContext): Boolean { + return expected.isProtocol(context) && PyNames.CALL in expected.getMemberNames(false, context) } - private static @Nullable PyTupleType widenUnionOfTuplesToTupleOfUnions(@NotNull PyUnionType unionType, int elementCount) { - boolean consistsOfSameSizeTuples = ContainerUtil.all(unionType.getMembers(), member -> member instanceof PyTupleType tupleType && - !tupleType.isHomogeneous() && - elementCount == tupleType.getElementCount()); + private fun widenUnionOfTuplesToTupleOfUnions(unionType: PyUnionType, elementCount: Int): PyTupleType? { + val consistsOfSameSizeTuples = + unionType.members.all { member -> member is PyTupleType && !member.isHomogeneous && elementCount == member.elementCount } if (!consistsOfSameSizeTuples) { - return null; + return null } - List newTupleElements = IntStreamEx.range(elementCount) - .mapToObj(index -> PyTypeUtil.toStream(unionType) - .select(PyTupleType.class) - .map(tupleType -> tupleType.getElementType(index)) - .collect(PyTypeUtil.toUnion())) - .toList(); - PyClass tupleClass = ((PyTupleType)ContainerUtil.getFirstItem(unionType.getMembers())).getPyClass(); - return new PyTupleType(tupleClass, newTupleElements, false); + val newTupleElements = List(elementCount) { index -> + unionType.members.asSequence() + .filterIsInstance() + .map { it.getElementType(index) } + .toList() + .let(PyUnionType::union) + } + val tupleClass = (unionType.members.firstOrNull() as PyTupleType).pyClass + return PyTupleType(tupleClass, newTupleElements, false) } - private static boolean matchGenerics(@NotNull PyCollectionType expected, @NotNull PyClassType actual, @NotNull MatchContext context) { - if (actual instanceof PyCollectionType && expected.getPyClass().equals(actual.getPyClass())) { - return matchGenericClassesParameterWise(expected, (PyCollectionType)actual, context); + private fun matchGenerics(expected: PyCollectionType, actual: PyClassType, context: MatchContext): Boolean { + if (actual is PyCollectionType && expected.pyClass == actual.pyClass) { + return matchGenericClassesParameterWise(expected, actual, context) } - PyCollectionType expectedGenericType = findGenericDefinitionType(expected.getPyClass(), context.context); + val expectedGenericType = findGenericDefinitionType(expected.pyClass, context.context) if (expectedGenericType != null) { - GenericSubstitutions actualSubstitutions = collectTypeSubstitutions(actual, context.context); - PyCollectionType concreteExpected = (PyCollectionType)substitute(expectedGenericType, actualSubstitutions, context.context); - assert concreteExpected != null; - return matchGenericClassesParameterWise(expected, concreteExpected, context); + val actualSubstitutions = collectTypeSubstitutions(actual, context.context) + val concreteExpected: PyCollectionType = + checkNotNull(substitute(expectedGenericType, actualSubstitutions, context.context) as PyCollectionType?) + return matchGenericClassesParameterWise(expected, concreteExpected, context) } - return true; + return true } // 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]
-   * }
- *

+ *

`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
-   * }
+ *
`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(); + private fun collectTypeSubstitutions(classType: PyClassType, context: TypeEvalContext): GenericSubstitutions { + val result = 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()) { - if (entry.getKey() instanceof PyTypeVarType typeVarType) { - result.typeVars.put(typeVarType, Ref.create(entry.getValue())); - } - else if (entry.getKey() instanceof PyTypeVarTupleType typeVarTuple) { - assert entry.getValue() instanceof PyPositionalVariadicType; - result.typeVarTuples.put(typeVarTuple, (PyPositionalVariadicType)entry.getValue()); - } - else if (entry.getKey() instanceof PyParamSpecType specType && - entry.getValue() instanceof PyCallableParameterVariadicType paramSpecType) { - result.paramSpecs.put(specType, paramSpecType); + for (provider in PyTypeProvider.EP_NAME.extensionList) { + val substitutionsFromClassDefinition = provider.getGenericSubstitutions(classType.pyClass, context) + for ((key, value) in substitutionsFromClassDefinition) { + when (key) { + is PyTypeVarType -> result.putTypeVar(key, Ref(value), KeyImpl) + is PyTypeVarTupleType -> result.putTypeVarTuple(key, value as PyPositionalVariadicType, KeyImpl) + is PyParamSpecType -> if (value is PyCallableParameterVariadicType) + result.putParamSpec(key, value, KeyImpl) } } // Collect own type parameters. In the example above these are: Sub@T3 -> str - PyCollectionType genericDefinitionType = as(provider.getGenericType(classType.getPyClass(), context), PyCollectionType.class); + val genericDefinitionType = provider.getGenericType(classType.pyClass, context) as? PyCollectionType // 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(); + val definitionTypeParameters = genericDefinitionType.elementTypes // Inside method bodies, where Self type can appear, we map class' own type parameters to themseleves, // not to consider them unbound. For instance, here @@ -1053,258 +1056,283 @@ public final class PyTypeChecker { // // 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); - } - else if (typeParameter instanceof PyParamSpecType paramSpecType) { - result.paramSpecs.put(paramSpecType, null); - } - else if (typeParameter instanceof PyTypeVarType typeVarType) { - result.typeVars.put(typeVarType, null); + when (classType) { + is PySelfType -> { + mapTypeParametersToSubstitutions( + result, definitionTypeParameters, definitionTypeParameters, + PyTypeParameterMapping.Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY + ) + } + !is PyCollectionType -> { + for (typeParameter in definitionTypeParameters) { + when (typeParameter) { + is PyTypeVarType -> result.putTypeVar(typeParameter, null, KeyImpl) + is PyTypeVarTupleType -> result.putTypeVarTuple(typeParameter, null, KeyImpl) + is PyParamSpecType -> result.putParamSpec(typeParameter, null, KeyImpl) + } } } - } - else { - mapTypeParametersToSubstitutions(result, definitionTypeParameters, genericType.getElementTypes(), - Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY); + else -> { + mapTypeParametersToSubstitutions( + result, definitionTypeParameters, classType.elementTypes, + PyTypeParameterMapping.Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY + ) + } } } - if (!result.typeVars.isEmpty() || !result.typeVarTuples.isEmpty() || !result.paramSpecs.isEmpty()) { - return result; + if (result.typeVars.isNotEmpty() || result.typeVarTuples.isNotEmpty() || result.paramSpecs.isNotEmpty()) { + return result } } - return result; + return result } + @JvmStatic @ApiStatus.Internal - public static @Nullable PyCollectionType findGenericDefinitionType(@NotNull PyClass pyClass, @NotNull TypeEvalContext context) { - for (PyTypeProvider provider : PyTypeProvider.EP_NAME.getExtensionList()) { - PyType definitionType = provider.getGenericType(pyClass, context); - if (definitionType instanceof PyCollectionType) { - return (PyCollectionType)definitionType; + fun findGenericDefinitionType(pyClass: PyClass, context: TypeEvalContext): PyCollectionType? { + for (provider in PyTypeProvider.EP_NAME.extensionList) { + val definitionType = provider.getGenericType(pyClass, context) + if (definitionType is PyCollectionType) { + return definitionType } } - return null; + return null } - private static boolean matchGenericClassesParameterWise(@NotNull PyCollectionType expected, - @NotNull PyCollectionType actual, - @NotNull MatchContext context) { - if (expected.equals(actual)) { - return true; + private fun matchGenericClassesParameterWise( + expected: PyCollectionType, + actual: PyCollectionType, + context: MatchContext, + ): Boolean { + if (expected == actual) { + return true } - if (!expected.getPyClass().equals(actual.getPyClass())) { - return false; + if (expected.pyClass != actual.pyClass) { + return false } - List expectedElementTypes = expected.getElementTypes(); - List actualElementTypes = actual.getElementTypes(); + val expectedElementTypes = expected.elementTypes + val actualElementTypes = actual.elementTypes if (context.reversedSubstitutions) { - return matchTypeParameters(actualElementTypes, expectedElementTypes, context.resetSubstitutions()); + return matchTypeParameters(actualElementTypes, expectedElementTypes, context.resetSubstitutions()) } else { - return matchTypeParameters(expectedElementTypes, actualElementTypes, context); + return matchTypeParameters(expectedElementTypes, actualElementTypes, context) } } - private static boolean matchTypeParameters(@NotNull List expectedTypeParameters, - @NotNull List actualTypeParameters, - @NotNull MatchContext context) { - PyTypeParameterMapping mapping = - PyTypeParameterMapping.mapByShape(expectedTypeParameters, actualTypeParameters, Option.USE_DEFAULTS); + private fun matchTypeParameters( + expectedTypeParameters: List, + actualTypeParameters: List, + context: MatchContext, + ): Boolean { + val mapping = + PyTypeParameterMapping.mapByShape(expectedTypeParameters, actualTypeParameters, PyTypeParameterMapping.Option.USE_DEFAULTS) if (mapping == null) { - return false; + return false } - for (Couple pair : mapping.getMappedTypes()) { - Optional matched = context.reversedSubstitutions - ? match(pair.getSecond(), pair.getFirst(), context) - : match(pair.getFirst(), pair.getSecond(), context); - if (!matched.orElse(true)) { - return false; + for (pair in mapping.mappedTypes) { + val matched = if (context.reversedSubstitutions) + match(pair.getSecond(), pair.getFirst(), context) + else + match(pair.getFirst(), pair.getSecond(), context) + if (!matched.orElse(true)!!) { + return false } } - return true; + return true } - private static boolean matchNumericTypes(PyType expected, PyType actual) { - if (expected instanceof PyClassType && actual instanceof PyClassType) { - final String superName = ((PyClassType)expected).getPyClass().getName(); - final String subName = ((PyClassType)actual).getPyClass().getName(); - final boolean subIsBool = "bool".equals(subName); - final boolean subIsInt = PyNames.TYPE_INT.equals(subName); - final boolean subIsLong = PyNames.TYPE_LONG.equals(subName); - final boolean subIsFloat = "float".equals(subName); - final boolean subIsComplex = "complex".equals(subName); + private fun matchNumericTypes(expected: PyType?, actual: PyType?): Boolean { + if (expected is PyClassType && actual is PyClassType) { + val superName = expected.pyClass.name + val subName = actual.pyClass.name + val subIsBool = "bool" == subName + val subIsInt = PyNames.TYPE_INT == subName + val subIsLong = PyNames.TYPE_LONG == subName + val subIsFloat = "float" == subName + val subIsComplex = "complex" == subName if (superName == null || subName == null || - superName.equals(subName) || - (PyNames.TYPE_INT.equals(superName) && subIsBool) || - ((PyNames.TYPE_LONG.equals(superName) || PyNames.ABC_INTEGRAL.equals(superName)) && (subIsBool || subIsInt)) || - (("float".equals(superName) || PyNames.ABC_REAL.equals(superName)) && (subIsBool || subIsInt || subIsLong)) || - (("complex".equals(superName) || PyNames.ABC_COMPLEX.equals(superName)) && (subIsBool || subIsInt || subIsLong || subIsFloat)) || - (PyNames.ABC_NUMBER.equals(superName) && (subIsBool || subIsInt || subIsLong || subIsFloat || subIsComplex))) { - return true; + superName == subName || + (PyNames.TYPE_INT == superName && subIsBool) || + ((PyNames.TYPE_LONG == superName || PyNames.ABC_INTEGRAL == superName) && (subIsBool || subIsInt)) || + (("float" == superName || PyNames.ABC_REAL == superName) && (subIsBool || subIsInt || subIsLong)) || + (("complex" == superName || PyNames.ABC_COMPLEX == superName) && (subIsBool || subIsInt || subIsLong || subIsFloat)) || + (PyNames.ABC_NUMBER == superName && (subIsBool || subIsInt || subIsLong || subIsFloat || subIsComplex)) + ) { + return true } } - return false; + return false } - public static boolean isUnknown(@Nullable PyType type, @NotNull TypeEvalContext context) { - return isUnknown(type, true, context); + @JvmStatic + fun isUnknown(type: PyType?, context: TypeEvalContext): Boolean { + return isUnknown(type, true, context) } - public static boolean isUnknown(@Nullable PyType type, boolean genericsAreUnknown, @NotNull TypeEvalContext context) { + @JvmStatic + fun isUnknown(type: PyType?, genericsAreUnknown: Boolean, context: TypeEvalContext): Boolean { // 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 == null || (genericsAreUnknown && type is PyTypeParameterType && (type !is PySelfType))) { + return true } - if (type instanceof PyUnionType union) { - if (!PyUnionType.isStrictSemanticsEnabled()) { - return ContainerUtil.exists(union.getMembers(), member -> isUnknown(member, genericsAreUnknown, context)); + return when (type) { + is PyUnionType -> { + if (!PyUnionType.isStrictSemanticsEnabled()) { + type.members.any { member: PyType? -> isUnknown(member, genericsAreUnknown, context) } + } + else type.members.all { member: PyType? -> isUnknown(member, genericsAreUnknown, context) } } - return ContainerUtil.all(union.getMembers(), member -> isUnknown(member, genericsAreUnknown, context)); + is PyUnsafeUnionType -> { + type.members.any { member: PyType? -> isUnknown(member, genericsAreUnknown, context) } + } + is PyIntersectionType -> { + type.members.any { member: PyType? -> isUnknown(member, genericsAreUnknown, context) } + } + else -> false } - if (type instanceof PyUnsafeUnionType weakUnion) { - return ContainerUtil.exists(weakUnion.getMembers(), member -> isUnknown(member, genericsAreUnknown, context)); - } - if (type instanceof PyIntersectionType intersectionType) { - return ContainerUtil.exists(intersectionType.getMembers(), member -> isUnknown(member, genericsAreUnknown, context)); - } - return false; } - public static @NotNull GenericSubstitutions getSubstitutionsWithUnresolvedReturnGenerics(@NotNull Collection parameters, - @Nullable PyType returnType, - @Nullable GenericSubstitutions substitutions, - @NotNull TypeEvalContext context) { - GenericSubstitutions existingSubstitutions = substitutions == null ? new GenericSubstitutions() : substitutions; - Generics typeParamsFromReturnType = collectGenerics(returnType, context); + @JvmStatic + fun getSubstitutionsWithUnresolvedReturnGenerics( + parameters: Collection, + returnType: PyType?, + substitutions: GenericSubstitutions?, + context: TypeEvalContext, + ): GenericSubstitutions { + val substitutions = substitutions ?: GenericSubstitutions() + val typeParamsFromReturnType = returnType.collectGenerics(context) // TODO Handle unmatched TypeVarTuples here as well if (typeParamsFromReturnType.typeVars.isEmpty() && typeParamsFromReturnType.paramSpecs.isEmpty()) { - return existingSubstitutions; + return substitutions } - Generics typeParamsFromParameterTypes = new Generics(); - for (PyCallableParameter parameter : parameters) { - collectGenerics(parameter.getArgumentType(context), context, typeParamsFromParameterTypes); + val typeParamsFromParameterTypes = GenericsImpl() + for (parameter in parameters) { + collectGenerics(parameter.getArgumentType(context), context, typeParamsFromParameterTypes) } - for (PyTypeVarType returnTypeParam : typeParamsFromReturnType.typeVars) { - boolean canGetBoundFromArguments = typeParamsFromParameterTypes.typeVars.contains(returnTypeParam) || - typeParamsFromParameterTypes.typeVars.contains(invert(returnTypeParam)); - boolean isAlreadyBound = existingSubstitutions.typeVars.containsKey(returnTypeParam) || - existingSubstitutions.typeVars.containsKey(invert(returnTypeParam)); + for (returnTypeParam in typeParamsFromReturnType.typeVars) { + val canGetBoundFromArguments = returnTypeParam in typeParamsFromParameterTypes.typeVars || + returnTypeParam.invert() in typeParamsFromParameterTypes.typeVars + val isAlreadyBound = returnTypeParam in substitutions.typeVars || + returnTypeParam.invert() in substitutions.typeVars if (canGetBoundFromArguments && !isAlreadyBound) { - existingSubstitutions.typeVars.put(returnTypeParam, (Ref)returnTypeParam.getDefaultType()); + @Suppress("UNCHECKED_CAST") + substitutions.putTypeVar(returnTypeParam, returnTypeParam.defaultType as Ref?, KeyImpl) } } - for (PyParamSpecType paramSpecType : typeParamsFromReturnType.paramSpecs) { - boolean canGetBoundFromArguments = typeParamsFromParameterTypes.paramSpecs.contains(paramSpecType); - boolean isAlreadyBound = existingSubstitutions.paramSpecs.containsKey(paramSpecType); + for (paramSpecType in typeParamsFromReturnType.paramSpecs) { + val canGetBoundFromArguments = paramSpecType in typeParamsFromParameterTypes.paramSpecs + val isAlreadyBound = paramSpecType in substitutions.paramSpecs if (canGetBoundFromArguments && !isAlreadyBound) { - if (paramSpecType.getDefaultType() != null) { - existingSubstitutions.paramSpecs.put(paramSpecType, Ref.deref(paramSpecType.getDefaultType())); + if (paramSpecType.defaultType != null) { + substitutions.putParamSpec(paramSpecType, Ref.deref(paramSpecType.defaultType), KeyImpl) } else { - existingSubstitutions.paramSpecs.put(paramSpecType, new PyCallableParameterListTypeImpl( - List.of(PyCallableParameterImpl.positionalNonPsi("args", null), - PyCallableParameterImpl.keywordNonPsi("kwargs", null))) - ); + substitutions.putParamSpec( + paramSpecType, + PyCallableParameterListTypeImpl( + listOf( + PyCallableParameterImpl.positionalNonPsi("args", null), + PyCallableParameterImpl.keywordNonPsi("kwargs", null), + ) + ), + KeyImpl, + ) } } } - return existingSubstitutions; + return substitutions } - private static @NotNull > T invert(@NotNull PyInstantiableType instantiable) { - return instantiable.isDefinition() ? instantiable.toInstance() : instantiable.toClass(); + private fun > PyInstantiableType.invert(): T { + return if (isDefinition) toInstance() else toClass() } - public static boolean hasGenerics(@Nullable PyType type, @NotNull TypeEvalContext context) { - return !collectGenerics(type, context).isEmpty(); + @JvmStatic + fun PyType?.hasGenerics(context: TypeEvalContext): Boolean { + return !this.collectGenerics(context).isEmpty } + @JvmStatic @ApiStatus.Internal - public static @NotNull Generics collectGenerics(@Nullable PyType type, @NotNull TypeEvalContext context) { - final var result = new Generics(); - collectGenerics(type, context, result); - return result; + fun PyType?.collectGenerics(context: TypeEvalContext): Generics { + val result = GenericsImpl() + collectGenerics(this, context, result) + return result } - private static void collectGenerics(@Nullable PyType type, - @NotNull TypeEvalContext context, - @NotNull Generics generics) { - PyRecursiveTypeVisitor.traverse(type, context, new PyTypeTraverser() { - @Override - public @NotNull Traversal visitPyTypeVarType(@NotNull PyTypeVarType typeVarType) { - generics.typeVars.add(typeVarType); - return super.visitPyTypeVarType(typeVarType); + private fun collectGenerics( + type: PyType?, + context: TypeEvalContext, + generics: GenericsImpl, + ) { + PyRecursiveTypeVisitor.traverse(type, context, object : PyTypeTraverser() { + override fun visitPyTypeVarType(typeVarType: PyTypeVarType): PyRecursiveTypeVisitor.Traversal { + generics.typeVars.add(typeVarType) + return super.visitPyTypeVarType(typeVarType) } - @Override - public @NotNull Traversal visitPyTypeVarTupleType(@NotNull PyTypeVarTupleType typeVarTupleType) { - generics.typeVarTuples.add(typeVarTupleType); - return super.visitPyTypeVarTupleType(typeVarTupleType); + override fun visitPyTypeVarTupleType(typeVarTupleType: PyTypeVarTupleType): PyRecursiveTypeVisitor.Traversal { + generics.typeVarTuples.add(typeVarTupleType) + return super.visitPyTypeVarTupleType(typeVarTupleType) } - @Override - public @NotNull Traversal visitPySelfType(@NotNull PySelfType selfType) { - generics.self = selfType; - return super.visitPySelfType(selfType); + override fun visitPySelfType(selfType: PySelfType): PyRecursiveTypeVisitor.Traversal { + generics.self = selfType + return super.visitPySelfType(selfType) } - @Override - public @NotNull Traversal visitPyParamSpecType(@NotNull PyParamSpecType paramSpecType) { - generics.paramSpecs.add(paramSpecType); - return super.visitPyParamSpecType(paramSpecType); + override fun visitPyParamSpecType(paramSpecType: PyParamSpecType): PyRecursiveTypeVisitor.Traversal { + generics.paramSpecs.add(paramSpecType) + return super.visitPyParamSpecType(paramSpecType) } - @Override - public @NotNull Traversal visitPyTypeParameterType(@NotNull PyTypeParameterType typeParameterType) { - generics.allTypeParameters.add(typeParameterType); - return Traversal.PRUNE; + override fun visitPyTypeParameterType(typeParameterType: PyTypeParameterType): PyRecursiveTypeVisitor.Traversal { + generics.allTypeParameters.add(typeParameterType) + return PyRecursiveTypeVisitor.Traversal.PRUNE } - }); + }) } - public static @Nullable PyType substitute(@Nullable PyType type, - @NotNull GenericSubstitutions substitutions, - @NotNull TypeEvalContext context) { - return PyCloningTypeVisitor.clone(type, new PyCloningTypeVisitor(context) { - private static @NotNull List<@Nullable PyType> flattenUnpackedTuple(@Nullable PyType type) { - if (type instanceof PyUnpackedTupleType unpackedTupleType && !unpackedTupleType.isUnbound()) { - return unpackedTupleType.getElementTypes(); + @JvmStatic + fun substitute( + type: PyType?, + substitutions: GenericSubstitutions, + context: TypeEvalContext, + ): PyType? { + return PyCloningTypeVisitor.clone(type, object : PyCloningTypeVisitor(context) { + fun flattenUnpackedTuple(type: PyType?): List { + return if (type is PyUnpackedTupleType && !type.isUnbound) { + type.elementTypes } - return Collections.singletonList(type); + else listOf(type) } - @Override - public PyType visitPyUnpackedTupleType(@NotNull PyUnpackedTupleType unpackedTupleType) { - return new PyUnpackedTupleTypeImpl(ContainerUtil.flatMap(unpackedTupleType.getElementTypes(), t -> flattenUnpackedTuple(clone(t))), - unpackedTupleType.isUnbound()); + override fun visitPyUnpackedTupleType(unpackedTupleType: PyUnpackedTupleType): PyType { + return PyUnpackedTupleTypeImpl( + unpackedTupleType.elementTypes.flatMap { + flattenUnpackedTuple(clone(it)) + }, + unpackedTupleType.isUnbound) } - @Override - public PyType visitPyTypeVarTupleType(@NotNull PyTypeVarTupleType typeVarTupleType) { - if (!substitutions.typeVarTuples.containsKey(typeVarTupleType)) { - return typeVarTupleType; + override fun visitPyTypeVarTupleType(typeVarTupleType: PyTypeVarTupleType): PyType? { + if (typeVarTupleType !in substitutions.typeVarTuples) { + return typeVarTupleType } - PyPositionalVariadicType substitution = substitutions.typeVarTuples.get(typeVarTupleType); - if (!typeVarTupleType.equals(substitution) && hasGenerics(substitution, context)) { - return clone(substitution); + val substitution = substitutions.typeVarTuples[typeVarTupleType] + if (typeVarTupleType != substitution && substitution.hasGenerics(context)) { + return clone(substitution) } // TODO This should happen in getSubstitutionsWithUnresolvedReturnGenerics, but it won't work for inherited constructors // as in testVariadicGenericCheckTypeAliasesRedundantParameter, investigate why // Replace unknown TypeVarTuples by *tuple[Any, ...] instead of plain Any - return substitution == null ? PyUnpackedTupleTypeImpl.UNSPECIFIED : substitution; + return substitution ?: PyUnpackedTupleTypeImpl.UNSPECIFIED } - @Override - public PyType visitPyTypeVarType(@NotNull PyTypeVarType typeVarType) { + override fun visitPyTypeVarType(typeVarType: PyTypeVarType): PyType? { // Both mappings of kind {T: T2} (and no mapping for T2) and {T: T} mean the substitution process should stop for T. // The first one occurs in the situations like // @@ -1317,79 +1345,64 @@ public final class PyTypeChecker { // def g() -> Callable[[T], T] // // where we want to prevent replacing T with Any, even though it cannot be inferred at a call site. - if (!substitutions.typeVars.containsKey(typeVarType) && !substitutions.typeVars.containsKey(invert(typeVarType))) { - PyTypeVarType substitution = StreamEx.of(substitutions.typeVars.keySet()) - .findFirst(typeVarType2 -> { - return typeVarType2.getDeclarationElement() != null - && (typeVarType.getScopeOwner() == null || (typeVarType2.getScopeOwner() == typeVarType.getScopeOwner())) - && typeVarType2.getDeclarationElement().equals(typeVarType.getDeclarationElement()); - }) - .orElse(null); + if (typeVarType !in substitutions.typeVars && typeVarType.invert() !in substitutions.typeVars) { + val substitution = substitutions.typeVars.keys.firstOrNull { typeVarType2 -> + typeVarType2.declarationElement != null && (typeVarType.scopeOwner == null || (typeVarType2.scopeOwner === typeVarType.scopeOwner)) + && typeVarType2.declarationElement == typeVarType.declarationElement + } if (substitution != null) { - return clone(substitution); + return clone(substitution) } - return typeVarType; + return typeVarType } - Ref substitutionRef = substitutions.typeVars.get(typeVarType); - PyType substitution = Ref.deref(substitutionRef); + val substitutionRef = substitutions.typeVars[typeVarType] + var substitution = Ref.deref(substitutionRef) if (substitutionRef == null) { - final PyInstantiableType invertedTypeVar = invert(typeVarType); - final PyInstantiableType invertedSubstitution = - as(Ref.deref(substitutions.typeVars.get(invertedTypeVar)), PyInstantiableType.class); + val invertedTypeVar: PyInstantiableType<*> = typeVarType.invert() + val invertedSubstitution = Ref.deref(substitutions.typeVars[invertedTypeVar]) as? PyInstantiableType<*> if (invertedSubstitution != null) { - substitution = invert(invertedSubstitution); + substitution = invertedSubstitution.invert() } } - if (substitution instanceof PyTypeVarType typeVarSubstitution) { + if (substitution is PyTypeVarType) { + val sameScopeSubstitution = substitutions.typeVars.keys.firstOrNull { typeVarType2 -> + (typeVarType2 != substitution) && typeVarType2.declarationElement != null && typeVarType2.declarationElement == substitution.declarationElement + } - PyTypeVarType sameScopeSubstitution = StreamEx.of(substitutions.typeVars.keySet()) - .findFirst(typeVarType2 -> { - return !typeVarType2.equals(typeVarSubstitution) - && typeVarType2.getDeclarationElement() != null - && typeVarType2.getDeclarationElement().equals(typeVarSubstitution.getDeclarationElement()); - }).orElse(null); - - if (sameScopeSubstitution != null && typeVarSubstitution.getDefaultType() != null) { - return clone(sameScopeSubstitution); + if (sameScopeSubstitution != null && substitution.defaultType != null) { + return clone(sameScopeSubstitution) } } // TODO remove !typeVar.equals(substitution) part, it's necessary due to the logic in unifyReceiverWithParamSpecs - if (!typeVarType.equals(substitution) && !(substitution instanceof PySelfType) && hasGenerics(substitution, context)) { - return clone(substitution); + if ((typeVarType != substitution) && (substitution !is PySelfType) && substitution.hasGenerics(context)) { + return clone(substitution) } - return substitution; + return substitution } - @Override - public PyType visitPyParamSpecType(@NotNull PyParamSpecType paramSpecType) { - if (!substitutions.paramSpecs.containsKey(paramSpecType)) { - PyParamSpecType sameScopeSubstitution = StreamEx.of(substitutions.paramSpecs.keySet()) - .findFirst(typeVarType -> { - return typeVarType.getDeclarationElement() != null - && typeVarType.getDeclarationElement().equals(paramSpecType.getDeclarationElement()); - }).orElse(null); + override fun visitPyParamSpecType(paramSpecType: PyParamSpecType): PyType? { + if (paramSpecType !in substitutions.paramSpecs) { + val sameScopeSubstitution = substitutions.paramSpecs.keys.firstOrNull { typeVarType -> + typeVarType.declarationElement != null && typeVarType.declarationElement == paramSpecType.declarationElement + } if (sameScopeSubstitution != null) { - return clone(sameScopeSubstitution); + return clone(sameScopeSubstitution) } - return paramSpecType; + return paramSpecType } - PyCallableParameterVariadicType substitution = substitutions.paramSpecs.get(paramSpecType); - if (substitution != null && !substitution.equals(paramSpecType) && hasGenerics(substitution, context)) { - return clone(substitution); + val substitution = substitutions.paramSpecs[paramSpecType] + if (substitution != null && (substitution != paramSpecType) && substitution.hasGenerics(context)) { + return clone(substitution) } // TODO For ParamSpecs, replace Any with (*args: Any, **kwargs: Any) as it's a logical "wildcard" for this kind of type parameter - return substitution; + return substitution } - @Override - public PyType visitPySelfType(@NotNull PySelfType selfType) { - var qualifierType = substitutions.qualifierType; - var selfScopeClassType = selfType.getScopeClassType(); - if (qualifierType == null) { - return selfType; - } + override fun visitPySelfType(selfType: PySelfType): PyType? { + val qualifierType = substitutions.qualifierType ?: return selfType + val selfScopeClassType = selfType.scopeClassType // TODO change unification for calls on union types // so that in the following // @@ -1402,364 +1415,382 @@ public final class PyTypeChecker { // 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) { - return selfScopeClassType.isDefinition() ? instantiableType.toClass() : instantiableType.toInstance(); + .map { qType: PyType? -> + if (qType is PyInstantiableType<*>) { + return@map if (selfScopeClassType.isDefinition) qType.toClass() else qType.toInstance() } - return qType; - }) - .filter(normalizedQType -> match(selfScopeClassType, normalizedQType, context)) - .collect(PyTypeUtil.toUnion(qualifierType)); + qType + } + .filter { normalizedQType: PyType? -> match(selfScopeClassType, normalizedQType, context) } + .collect(PyTypeUtil.toUnion(qualifierType)) } - @Override - public PyType visitPyGenericType(@NotNull PyCollectionType genericType) { - return new PyCollectionTypeImpl(genericType.getPyClass(), genericType.isDefinition(), - ContainerUtil.flatMap(genericType.getElementTypes(), t -> flattenUnpackedTuple(clone(t)))); + override fun visitPyGenericType(genericType: PyCollectionType): PyType { + return PyCollectionTypeImpl( + genericType.pyClass, genericType.isDefinition, + genericType.elementTypes.flatMap { + flattenUnpackedTuple(clone(it)) + } + ) } - @Override - public PyType visitPyTupleType(@NotNull PyTupleType tupleType) { - final PyClass tupleClass = tupleType.getPyClass(); - final List oldElementTypes = tupleType.isHomogeneous() - ? Collections.singletonList(tupleType.getIteratedItemType()) - : tupleType.getElementTypes(); - return new PyTupleType(tupleClass, - ContainerUtil.flatMap(oldElementTypes, elementType -> flattenUnpackedTuple(clone(elementType))), - tupleType.isHomogeneous()); + override fun visitPyTupleType(tupleType: PyTupleType): PyType { + val tupleClass = tupleType.pyClass + val oldElementTypes = if (tupleType.isHomogeneous) listOf(tupleType.iteratedItemType) + else + tupleType.elementTypes + return PyTupleType( + tupleClass, + oldElementTypes.flatMap { + flattenUnpackedTuple( + clone(it)) + }, + tupleType.isHomogeneous, + ) } - @Override - public PyType visitPyFunctionType(@NotNull PyFunctionType functionType) { + override fun visitPyFunctionType(functionType: PyFunctionType): PyType? { // TODO legacy behavior - if (hasGenerics(functionType, context)) { - return super.visitPyFunctionType(functionType); + if (functionType.hasGenerics(context)) { + return super.visitPyFunctionType(functionType) } - return functionType; + return functionType } - @Override - public PyType visitPyCallableType(@NotNull PyCallableType callableType) { - PyCallableParameterVariadicType substitutedParams = clone(callableType.getParametersType(context)); - return new PyCallableTypeImpl( + override fun visitPyCallableType(callableType: PyCallableType): PyType { + val substitutedParams = clone(callableType.getParametersType(context)) + return PyCallableTypeImpl( substitutedParams, clone(callableType.getReturnType(context)), - callableType.getCallable(), - callableType.getModifier(), - callableType.getImplicitOffset() - ); + callableType.callable, + callableType.modifier, + callableType.implicitOffset, + ) } - @Override - public PyType visitPyCallableParameterListType(@NotNull PyCallableParameterListType callableParameterListType) { - List parameters = callableParameterListType.getParameters(); + override fun visitPyCallableParameterListType(callableParameterListType: PyCallableParameterListType): PyType { + val parameters = callableParameterListType.parameters - boolean isSelfArgsKwargs = ParamHelper.isSelfArgsKwargsSignature(parameters); - boolean isArgsKwargs = ParamHelper.isArgsKwargsSignature(parameters); + val isSelfArgsKwargs = ParamHelper.isSelfArgsKwargsSignature(parameters) + val isArgsKwargs = ParamHelper.isArgsKwargsSignature(parameters) // substitutes (*args: P.args, **kwargs: P.kwargs) with actual signature if (isSelfArgsKwargs || isArgsKwargs) { - PyType kwargsType = parameters.getLast().getType(context); - PyCallableParameterVariadicType substituted = substitutions.getParamSpecs().get(kwargsType); + val kwargsType: PyType? = parameters.last().getType(context) + val substituted = substitutions.paramSpecs[kwargsType] - if (substituted instanceof PyCallableParameterListType listType) { - List newParams = listType.getParameters(); + if (substituted is PyCallableParameterListType) { + var newParams = substituted.parameters if (isSelfArgsKwargs) { - newParams = ContainerUtil.prepend(newParams, parameters.getFirst()); + newParams = ContainerUtil.prepend(newParams, parameters.first()) } - return new PyCallableParameterListTypeImpl(newParams); + return PyCallableParameterListTypeImpl(newParams) } } - List substitutedParams = StreamEx.of(parameters) - .mapToEntry(param -> param.getType(context)) - .flatMapKeyValue((param, paramType) -> { - PyParameter paramPsi = param.getParameter(); - return StreamEx.of(Collections.singletonList(param.getType(context))) - .flatCollection(t -> flattenUnpackedTuple(clone(t))) - .map(paramSubType -> paramPsi != null ? - PyCallableParameterImpl.psi(paramPsi, paramSubType) : - PyCallableParameterImpl.nonPsi(param.getName(), paramSubType, param.getDefaultValue())); - }) - .toList(); + val substitutedParams = parameters + .associateWith { it.getType(context) } + .flatMap { (param, _) -> + val paramPsi = param.parameter + flattenUnpackedTuple(clone(param.getType(context))) + .map { paramSubType -> + if (paramPsi != null) PyCallableParameterImpl.psi( + paramPsi, + paramSubType + ) + else PyCallableParameterImpl.nonPsi(param.name, paramSubType, param.defaultValue) + } + } - return new PyCallableParameterListTypeImpl(substitutedParams); + return PyCallableParameterListTypeImpl(substitutedParams) } - @Override - public PyCallableParameterVariadicType visitPyConcatenateType(@NotNull PyConcatenateType concatenateType) { - List firstParamTypeSubs = ContainerUtil.flatMap(concatenateType.getFirstTypes(), t -> flattenUnpackedTuple(clone(t))); - PyCallableParameterVariadicType paramSpecSubs = clone(concatenateType.getParamSpec()); - if (paramSpecSubs instanceof PyCallableParameterListType callableParams) { - return new PyCallableParameterListTypeImpl( - StreamEx.of(firstParamTypeSubs) - .map(PyCallableParameterImpl::nonPsi) - .append(callableParams.getParameters()) - .toList() - ); + override fun visitPyConcatenateType(concatenateType: PyConcatenateType): PyCallableParameterVariadicType? { + val firstParamTypeSubs = concatenateType.firstTypes.flatMap { + flattenUnpackedTuple(clone(it)) } - else if (paramSpecSubs instanceof PyParamSpecType paramSpecType) { - return new PyConcatenateType(firstParamTypeSubs, paramSpecType); + val paramSpecSubs = clone(concatenateType.paramSpec) + + return when (paramSpecSubs) { + is PyCallableParameterListType -> { + PyCallableParameterListTypeImpl( + firstParamTypeSubs.map { PyCallableParameterImpl.nonPsi(it) } + paramSpecSubs.parameters + ) + } + is PyParamSpecType -> PyConcatenateType(firstParamTypeSubs, paramSpecSubs) + is PyConcatenateType -> { + PyConcatenateType( + firstParamTypeSubs + paramSpecSubs.firstTypes, + paramSpecSubs.paramSpec + ) + } + else -> null } - else if (paramSpecSubs instanceof PyConcatenateType concatenateType2) { - return new PyConcatenateType( - ContainerUtil.concat(firstParamTypeSubs, concatenateType2.getFirstTypes()), - concatenateType2.getParamSpec() - ); - } - return null; } - }); + }) } - public static @Nullable GenericSubstitutions unifyGenericCall(@Nullable PyExpression receiver, - @NotNull Map arguments, - @NotNull TypeEvalContext context) { + @JvmStatic + fun unifyGenericCall( + receiver: PyExpression?, + arguments: Map, + context: TypeEvalContext, + ): GenericSubstitutions? { // UnifyReceiver retains the information about type parameters of ancestor generic classes, // which might be necessary if a method being called is defined in one them, not in the actual // class of the receiver. In theory, it could be replaced by matching the type of self parameter // with the receiver uniformly with the rest of the arguments. - final var substitutions = unifyReceiver(receiver, context); - for (Map.Entry entry : getRegularMappedParameters(arguments).entrySet()) { - final PyCallableParameter paramWrapper = entry.getValue(); - final PyType expectedType = paramWrapper.getArgumentType(context); - final PyType promotedToLiteral = PyLiteralType.Companion.promoteToLiteral(entry.getKey(), expectedType, context, substitutions); - var actualType = promotedToLiteral != null ? promotedToLiteral : context.getType(entry.getKey()); + val substitutions = unifyReceiver(receiver, context) + for ((key, paramWrapper) in arguments.getRegularMappedParameters().entries) { + val expectedType = paramWrapper.getArgumentType(context) + val promotedToLiteral = PyLiteralType.promoteToLiteral(key, expectedType, context, substitutions) + val actualType = promotedToLiteral ?: context.getType(key) // Matching with the type of "self" is necessary in particular for choosing the most specific overloads, e.g. // LiteralString-specific methods of str, or for instantiating the type parameters of the containing class // when it's not possible to infer them by other means, e.g. as in the following overload of dict[_KT, _VT].__init__: // def __init__(self: dict[str, _VT], **kwargs: _VT) -> None: ... - boolean matchedByTypes = matchParameterArgumentTypes(paramWrapper, expectedType, actualType, substitutions, context); + val matchedByTypes = matchParameterArgumentTypes(paramWrapper, expectedType, actualType, substitutions, context) if (!matchedByTypes) { - return null; + return null } } - if (!matchContainer(getMappedPositionalContainer(arguments), getArgumentsMappedToPositionalContainer(arguments), - substitutions, context)) { - return null; + if (!matchContainer( + arguments.getMappedPositionalContainer(), arguments.getArgumentsMappedToPositionalContainer(), + substitutions, context + ) + ) { + return null } - if (!matchContainer(getMappedKeywordContainer(arguments), getArgumentsMappedToKeywordContainer(arguments), - substitutions, context)) { - return null; + if (!matchContainer( + arguments.getMappedKeywordContainer(), arguments.getArgumentsMappedToKeywordContainer(), + substitutions, context + ) + ) { + return null } - return substitutions; + return substitutions } - public static @Nullable GenericSubstitutions unifyGenericCallOnArgumentTypes(@Nullable PyType receiverType, - @NotNull Map, PyCallableParameter> arguments, - @NotNull TypeEvalContext context) { - final var substitutions = unifyReceiver(receiverType, context); - for (Map.Entry, PyCallableParameter> entry : getRegularMappedParameters(arguments).entrySet()) { - final PyCallableParameter paramWrapper = entry.getValue(); - final PyType expectedType = paramWrapper.getArgumentType(context); - var actualType = Ref.deref(entry.getKey()); - boolean matchedByTypes = matchParameterArgumentTypes(paramWrapper, expectedType, actualType, substitutions, context); + @JvmStatic + fun unifyGenericCallOnArgumentTypes( + receiverType: PyType?, + arguments: Map?, PyCallableParameter>, + context: TypeEvalContext, + ): GenericSubstitutions? { + val substitutions = unifyReceiver(receiverType, context) + for ((key, paramWrapper) in arguments.getRegularMappedParameters()) { + val expectedType = paramWrapper.getArgumentType(context) + val actualType = Ref.deref(key) + val matchedByTypes = matchParameterArgumentTypes(paramWrapper, expectedType, actualType, substitutions, context) if (!matchedByTypes) { - return null; + return null } } - if (!matchContainerByType(getMappedPositionalContainer(arguments), - ContainerUtil.map(getArgumentsMappedToPositionalContainer(arguments), Ref::deref), + if (!matchContainerByType(arguments.getMappedPositionalContainer(), + arguments.getArgumentsMappedToPositionalContainer().map(Ref<*>::deref), substitutions, context)) { - return null; + return null } - if (!matchContainerByType(getMappedKeywordContainer(arguments), - ContainerUtil.map(getArgumentsMappedToKeywordContainer(arguments), Ref::deref), + if (!matchContainerByType(arguments.getMappedKeywordContainer(), + arguments.getArgumentsMappedToKeywordContainer().map(Ref<*>::deref), substitutions, context)) { - return null; + return null } - return substitutions; + return substitutions } - private static @Nullable PyType processSelfParameter(@NotNull PyCallableParameter paramWrapper, - @Nullable PyType expectedType, - @Nullable PyType actualType, - @NotNull GenericSubstitutions substitutions, - @NotNull TypeEvalContext context) { + private fun processSelfParameter( + paramWrapper: PyCallableParameter, + expectedType: PyType?, + actualType: PyType?, + substitutions: GenericSubstitutions, + context: TypeEvalContext, + ): PyType? { // TODO find out a better way to pass the corresponding function inside - final PyParameter param = paramWrapper.getParameter(); - final PyFunction function = as(ScopeUtil.getScopeOwner(param), PyFunction.class); - assert function != null; - if (function.getModifier() == PyAstFunction.Modifier.CLASSMETHOD) { + var actualType = actualType + val param = paramWrapper.parameter + val function: PyFunction = ScopeUtil.getScopeOwner(param) as PyFunction + if (function.modifier == PyAstFunction.Modifier.CLASSMETHOD) { actualType = PyTypeUtil.toStream(actualType) - .select(PyClassLikeType.class) - .map(PyClassLikeType::toClass) - .select(PyType.class) - .foldLeft(PyUnionType::union) - .orElse(actualType); + .select(PyClassLikeType::class.java) + .map { obj: PyClassLikeType? -> obj!!.toClass() } + .select(PyType::class.java) + .foldLeft { type1: PyType?, type2: PyType? -> PyUnionType.union(type1, type2) } + .orElse(actualType) } else if (PyUtil.isInitMethod(function)) { actualType = PyTypeUtil.toStream(actualType) - .select(PyInstantiableType.class) - .map(PyInstantiableType::toInstance) - .select(PyType.class) - .foldLeft(PyUnionType::union) - .orElse(actualType); + .select(PyInstantiableType::class.java) + .map { obj -> obj!!.toInstance() } + .select(PyType::class.java) + .foldLeft { type1, type2 -> PyUnionType.union(type1, type2) } + .orElse(actualType) } if (PyUnionType.isStrictSemanticsEnabled()) { - PyClass pyClass = function.getContainingClass(); - assert pyClass != null; - PyClassLikeType classType = as(context.getType(pyClass), PyClassLikeType.class); - assert classType != null; - PyClassLikeType superType = - function.getModifier() == PyAstFunction.Modifier.CLASSMETHOD || isNewMethod(function) ? classType : classType.toInstance(); + val pyClass: PyClass = checkNotNull(function.containingClass) + val classType: PyClassLikeType = context.getType(pyClass) as PyClassLikeType + val superType: PyClassLikeType = + (if (function.modifier == PyAstFunction.Modifier.CLASSMETHOD || PyUtil.isNewMethod(function)) classType else classType.toInstance()) // In a union receiver type, leave only members that actually have this function // TODO how does it work with qualified calls, e.g. SomeClass.method(receiver, arg1, arg2) // TODO how does it work with @classmethods? actualType = PyTypeUtil.toStream(actualType) - .filter(type -> match(superType, type, context)) - .collect(PyTypeUtil.toUnion(actualType)); + .filter { type: PyType? -> match(superType, type, context) } + .collect(PyTypeUtil.toUnion(actualType)) } - PyClass containingClass = function.getContainingClass(); - assert containingClass != null; - PyType genericClass = findGenericDefinitionType(containingClass, context); - if (genericClass instanceof PyInstantiableType instantiableType && - (isNewMethod(function) || function.getModifier() == PyAstFunction.Modifier.CLASSMETHOD)) { - genericClass = instantiableType.toClass(); + val containingClass: PyClass = checkNotNull(function.containingClass) + var genericClass: PyType? = findGenericDefinitionType(containingClass, context) + if (genericClass is PyInstantiableType<*> && + (PyUtil.isNewMethod(function) || function.modifier == PyAstFunction.Modifier.CLASSMETHOD) + ) { + genericClass = genericClass.toClass() } if (genericClass != null && !match(genericClass, expectedType, context, substitutions)) { - return null; + return null } - return actualType; + return actualType } - private static boolean matchParameterArgumentTypes(@NotNull PyCallableParameter paramWrapper, - @Nullable PyType expectedType, - @Nullable PyType actualType, - @NotNull GenericSubstitutions substitutions, - @NotNull TypeEvalContext context) { - if (paramWrapper.isSelf()) { - if (expectedType instanceof PySelfType selfType) { - final PyCollectionType genericType = findGenericDefinitionType(selfType.getScopeClassType().getPyClass(), context); - if (genericType != null) { - expectedType = selfType.isDefinition() ? genericType.toClass() : genericType; - } - else { - expectedType = selfType.getScopeClassType(); + private fun matchParameterArgumentTypes( + paramWrapper: PyCallableParameter, + expectedType: PyType?, + actualType: PyType?, + substitutions: GenericSubstitutions, + context: TypeEvalContext, + ): Boolean { + var expectedType = expectedType + var actualType = actualType + if (paramWrapper.isSelf) { + if (expectedType is PySelfType) { + val genericType = findGenericDefinitionType(expectedType.scopeClassType.pyClass, context) + expectedType = if (genericType != null) { + if (expectedType.isDefinition) genericType.toClass() else genericType } + else expectedType.scopeClassType } - actualType = processSelfParameter(paramWrapper, expectedType, actualType, substitutions, context); - if (actualType == null) { - return false; - } + actualType = processSelfParameter(paramWrapper, expectedType, actualType, substitutions, context) ?: return false } - return match(expectedType, actualType, context, substitutions); + return match(expectedType, actualType, context, substitutions) } - private static boolean matchContainer(@Nullable PyCallableParameter container, @NotNull List arguments, - @NotNull GenericSubstitutions substitutions, @NotNull TypeEvalContext context) { + private fun matchContainer( + container: PyCallableParameter?, arguments: List, + substitutions: GenericSubstitutions, context: TypeEvalContext, + ): Boolean { if (container == null) { - return true; + return true } - final List actualArgumentTypes = ContainerUtil.map(arguments, context::getType); - return matchContainerByType(container, actualArgumentTypes, substitutions, context); + val actualArgumentTypes = arguments.map { context.getType(it) } + return matchContainerByType(container, actualArgumentTypes, substitutions, context) } - private static boolean matchContainerByType(@Nullable PyCallableParameter container, @NotNull List actualArgumentTypes, - @NotNull GenericSubstitutions substitutions, @NotNull TypeEvalContext context) { + private fun matchContainerByType( + container: PyCallableParameter?, actualArgumentTypes: List, + substitutions: GenericSubstitutions, context: TypeEvalContext, + ): Boolean { if (container == null) { - return true; + return true } - final PyType expectedArgumentType = container.getArgumentType(context); - if (container.isPositionalContainer() && expectedArgumentType instanceof PyPositionalVariadicType variadicType) { - return match(variadicType, PyUnpackedTupleTypeImpl.create(actualArgumentTypes), - new MatchContext(context, substitutions, false)); + val expectedArgumentType = container.getArgumentType(context) + if (container.isPositionalContainer && expectedArgumentType is PyPositionalVariadicType) { + return match( + expectedArgumentType, PyUnpackedTupleTypeImpl.create(actualArgumentTypes), + MatchContext(context, substitutions, false) + ) } - return match(expectedArgumentType, PyUnionType.union(actualArgumentTypes), context, substitutions); + return match(expectedArgumentType, PyUnionType.union(actualArgumentTypes), context, substitutions) } - public static @NotNull GenericSubstitutions unifyReceiver(@Nullable PyExpression receiver, @NotNull TypeEvalContext context) { + @JvmStatic + fun unifyReceiver(receiver: PyExpression?, context: TypeEvalContext): GenericSubstitutions { if (receiver != null) { - PyType receiverType = context.getType(receiver); - return unifyReceiver(receiverType, context); + val receiverType = context.getType(receiver) + return unifyReceiver(receiverType, context) } - return new GenericSubstitutions(); + return GenericSubstitutions() } - public static @NotNull GenericSubstitutions unifyReceiver(@Nullable PyType receiverType, @NotNull TypeEvalContext context) { + @JvmStatic + fun unifyReceiver(receiverType: PyType?, context: TypeEvalContext): GenericSubstitutions { // Collect generic params of object type - final var substitutions = new GenericSubstitutions(); + val substitutions = GenericSubstitutions() if (receiverType != null) { // TODO properly handle union types here - if (receiverType instanceof PyClassType) { - substitutions.qualifierType = ((PyClassType)receiverType).toInstance(); + if (receiverType is PyClassType) { + substitutions.qualifierType = receiverType.toInstance() } else { - substitutions.qualifierType = receiverType; + substitutions.qualifierType = receiverType } PyTypeUtil.toStream(receiverType) - .select(PyClassType.class) - .map(type -> collectTypeSubstitutions(type, context)) - .forEach(newSubstitutions -> { - for (Map.Entry> typeVarMapping : newSubstitutions.typeVars.entrySet()) { - substitutions.typeVars.putIfAbsent(typeVarMapping.getKey(), typeVarMapping.getValue()); + .select(PyClassType::class.java) + .map { type: PyClassType? -> collectTypeSubstitutions(type!!, context) } + .forEach { newSubstitutions -> + for (typeVarMapping in newSubstitutions.typeVars.entries) { + substitutions.putTypeVar(typeVarMapping.key, typeVarMapping.value, KeyImpl, true) } - for (Map.Entry typeVarMapping : newSubstitutions.typeVarTuples.entrySet()) { - substitutions.typeVarTuples.putIfAbsent(typeVarMapping.getKey(), typeVarMapping.getValue()); + for (typeVarMapping in newSubstitutions.typeVarTuples.entries) { + substitutions.putTypeVarTuple(typeVarMapping.key, typeVarMapping.value, KeyImpl, true) } - for (Map.Entry paramSpecMapping : newSubstitutions.paramSpecs.entrySet()) { - substitutions.paramSpecs.putIfAbsent(paramSpecMapping.getKey(), paramSpecMapping.getValue()); + for (paramSpecMapping in newSubstitutions.paramSpecs.entries) { + substitutions.putParamSpec(paramSpecMapping.key, paramSpecMapping.value, KeyImpl, true) } - }); + } } - substitutions.frozenTypeVars = new HashSet<>(substitutions.typeVars.keySet()); - return substitutions; + substitutions.setFrozenTypeVars(HashSet(substitutions.typeVars.keys), KeyImpl) + return substitutions } - private static boolean matchClasses(@Nullable PyClass superClass, @Nullable PyClass subClass, @NotNull TypeEvalContext context) { - if (superClass == null || - subClass == null || - subClass.isSubclass(superClass, context) || - PyABCUtil.isSubclass(subClass, superClass, context) || - isStrUnicodeMatch(subClass, superClass) || - isBytearrayBytesStringMatch(subClass, superClass) || - PyUtil.hasUnresolvedAncestors(subClass, context)) { - return true; + private fun matchClasses(superClass: PyClass?, subClass: PyClass?, context: TypeEvalContext): Boolean { + return if (superClass == null || subClass == null || + subClass.isSubclass(superClass, context) || + PyABCUtil.isSubclass(subClass, superClass, context) || + isStrUnicodeMatch(subClass, superClass) || + isBytearrayBytesStringMatch(subClass, superClass) || + PyUtil.hasUnresolvedAncestors(subClass, context) + ) { + true } else { - final String superName = superClass.getName(); - return superName != null && superName.equals(subClass.getName()); + val superName = superClass.name + superName != null && superName == subClass.name } } - private static boolean isStrUnicodeMatch(@NotNull PyClass subClass, @NotNull PyClass superClass) { + private fun isStrUnicodeMatch(subClass: PyClass, superClass: PyClass): Boolean { // TODO: Check for subclasses as well - return PyNames.TYPE_STR.equals(subClass.getName()) && PyNames.TYPE_UNICODE.equals(superClass.getName()); + return PyNames.TYPE_STR == subClass.name && PyNames.TYPE_UNICODE == superClass.name } - private static boolean isBytearrayBytesStringMatch(@NotNull PyClass subClass, @NotNull PyClass superClass) { - if (!PyNames.TYPE_BYTEARRAY.equals(subClass.getName())) return false; + private fun isBytearrayBytesStringMatch(subClass: PyClass, superClass: PyClass): Boolean { + if (PyNames.TYPE_BYTEARRAY != subClass.name) return false - final PsiFile subClassFile = subClass.getContainingFile(); + val subClassFile = subClass.containingFile - final boolean isPy2 = subClassFile instanceof PyiFile - ? PythonRuntimeService.getInstance().getLanguageLevelForSdk(PythonSdkUtil.findPythonSdk(subClassFile)).isPython2() - : LanguageLevel.forElement(subClass).isPython2(); + val isPy2 = if (subClassFile is PyiFile) + PythonRuntimeService.getInstance().getLanguageLevelForSdk(PythonSdkUtil.findPythonSdk(subClassFile)).isPython2 + else + LanguageLevel.forElement(subClass).isPython2 - final String superClassName = superClass.getName(); - return isPy2 && PyNames.TYPE_STR.equals(superClassName) || !isPy2 && PyNames.TYPE_BYTES.equals(superClassName); + val superClassName = superClass.name + return isPy2 && PyNames.TYPE_STR == superClassName || !isPy2 && PyNames.TYPE_BYTES == superClassName } - public static @Nullable Boolean isCallable(@Nullable PyType type) { - if (type == null) { - return null; - } - if (type instanceof PyUnionType) { - return isUnionCallable((PyUnionType)type); - } - if (type instanceof PyCallableType) { - return ((PyCallableType)type).isCallable(); - } - if (type instanceof PyStructuralType && ((PyStructuralType)type).isInferredFromUsages()) { - return true; - } - if (type instanceof PyTypeVarType typeVarType) { - if (typeVarType.isDefinition()) { - return true; + @JvmStatic + fun isCallable(type: PyType?): Boolean? { + return when (type) { + null -> null + is PyUnionType -> isUnionCallable(type) + is PyCallableType -> type.isCallable + is PyStructuralType if type.isInferredFromUsages -> true + is PyTypeVarType -> { + if (type.isDefinition) { + true + } + else isCallable(PyTypeUtil.getEffectiveBound(type)) } - return isCallable(PyTypeUtil.getEffectiveBound(typeVarType)); + else -> false } - return false; } /** @@ -1767,277 +1798,269 @@ public final class PyTypeChecker { * If at least one is unknown -- it is unknown. * It is false otherwise. */ - private static @Nullable Boolean isUnionCallable(final @NotNull PyUnionType type) { - for (final PyType member : type.getMembers()) { - final Boolean callable = isCallable(member); + private fun isUnionCallable(type: PyUnionType): Boolean? { + for (member in type.members) { + val callable = isCallable(member) if (callable == null) { - return null; + return null } if (callable) { - return true; + return true } } - return false; + return false } - public static boolean definesGetAttr(@NotNull PyFile file, @NotNull TypeEvalContext context) { - if (file instanceof PyTypedElement) { - final PyType type = context.getType((PyTypedElement)file); + @JvmStatic + fun definesGetAttr(file: PyFile, context: TypeEvalContext): Boolean { + if (file is PyTypedElement) { + val type = context.getType(file as PyTypedElement) if (type != null) { - return resolveTypeMember(type, PyNames.GETATTR, context) != null; + return resolveTypeMember(type, PyNames.GETATTR, context) != null } } - return false; + return false } - public static boolean overridesGetAttr(@NotNull PyClass cls, @NotNull TypeEvalContext context) { - final PyType type = context.getType(cls); + @JvmStatic + fun overridesGetAttr(cls: PyClass, context: TypeEvalContext): Boolean { + val type = context.getType(cls) if (type != null) { if (resolveTypeMember(type, PyNames.GETATTR, context) != null) { - return true; + return true } - final PsiElement method = resolveTypeMember(type, PyNames.GETATTRIBUTE, context); - if (method != null && !PyBuiltinCache.getInstance(cls).isBuiltin(method)) { - return true; + val method = resolveTypeMember(type, PyNames.GETATTRIBUTE, context) + if (method != null && !getInstance(cls).isBuiltin(method)) { + return true } } - return false; + return false } - private static @Nullable PsiElement resolveTypeMember(@NotNull PyType type, @NotNull String name, @NotNull TypeEvalContext context) { - final PyResolveContext resolveContext = PyResolveContext.defaultContext(context); - final List results = type.resolveMember(name, null, AccessDirection.READ, resolveContext); - return !ContainerUtil.isEmpty(results) ? results.get(0).getElement() : null; + private fun resolveTypeMember(type: PyType, name: String, context: TypeEvalContext): PsiElement? { + val resolveContext = PyResolveContext.defaultContext(context) + val results = type.resolveMember(name, null, AccessDirection.READ, resolveContext) + return results?.getOrNull(0)?.element } - public static @Nullable PyType getTargetTypeFromTupleAssignment(@NotNull PyExpression target, - @NotNull PySequenceExpression parentTupleOrList, - @NotNull PyTupleType assignedTupleType) { - final int count = assignedTupleType.getElementCount(); - final PyExpression[] elements = parentTupleOrList.getElements(); - if (elements.length == count || assignedTupleType.isHomogeneous()) { - final int index = ArrayUtil.indexOf(elements, target); + @JvmStatic + fun getTargetTypeFromTupleAssignment( + target: PyExpression, + parentTupleOrList: PySequenceExpression, + assignedTupleType: PyTupleType, + ): PyType? { + val count = assignedTupleType.elementCount + val elements = parentTupleOrList.elements + if (elements.size == count || assignedTupleType.isHomogeneous) { + val index = elements.indexOf(target) if (index >= 0) { - return assignedTupleType.getElementType(index); + return assignedTupleType.getElementType(index) } - for (int i = 0; i < count; i++) { - PyExpression element = PyPsiUtils.flattenParens(elements[i]); - if (element instanceof PyTupleExpression || element instanceof PyListLiteralExpression) { - final PyType elementType = assignedTupleType.getElementType(i); - if (elementType instanceof PyTupleType nestedAssignedTupleType) { - final PyType result = getTargetTypeFromTupleAssignment(target, (PySequenceExpression)element, nestedAssignedTupleType); + for (i in 0.. actualTypeParams, - @NotNull TypeEvalContext context) { - Generics typeParams = collectGenerics(genericType, context); - if (!typeParams.isEmpty()) { - List expectedTypeParams = new ArrayList<>(new LinkedHashSet<>(typeParams.getAllTypeParameters())); - var substitutions = mapTypeParametersToSubstitutions(expectedTypeParams, - actualTypeParams, - Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY, - Option.USE_DEFAULTS); - if (substitutions == null) return null; - return substitute(genericType, substitutions, context); + @JvmStatic + fun parameterizeType( + genericType: PyType, + actualTypeParams: List, + context: TypeEvalContext, + ): PyType? { + val typeParams = genericType.collectGenerics(context) + if (!typeParams.isEmpty) { + val expectedTypeParams = ArrayList(LinkedHashSet(typeParams.allTypeParameters)) + val substitutions = mapTypeParametersToSubstitutions( + expectedTypeParams, + actualTypeParams, + PyTypeParameterMapping.Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY, + PyTypeParameterMapping.Option.USE_DEFAULTS, + ) ?: return null + return substitute(genericType, substitutions, context) } - // An already parameterized type, don't override existing values for type parameters - else if (genericType instanceof PyCollectionType) { - return genericType; + else if (genericType is PyCollectionType) { + return genericType } - else if (genericType instanceof PyClassType) { - PyClass cls = ((PyClassType)genericType).getPyClass(); - return new PyCollectionTypeImpl(cls, false, actualTypeParams); + else if (genericType is PyClassType) { + val cls = genericType.pyClass + return PyCollectionTypeImpl(cls, false, actualTypeParams) } - return null; + return null } + @JvmStatic @ApiStatus.Internal - public static @Nullable GenericSubstitutions mapTypeParametersToSubstitutions(@NotNull List expectedTypes, - @NotNull List actualTypes, - Option @NotNull ... options) { - return mapTypeParametersToSubstitutions(new GenericSubstitutions(), expectedTypes, actualTypes, options); + fun mapTypeParametersToSubstitutions( + expectedTypes: List, + actualTypes: List, + vararg options: PyTypeParameterMapping.Option, + ): GenericSubstitutions? { + return mapTypeParametersToSubstitutions(GenericSubstitutions(), expectedTypes, actualTypes, *options) } - private static @Nullable GenericSubstitutions mapTypeParametersToSubstitutions(@NotNull GenericSubstitutions substitutions, - @NotNull List expectedTypes, - @NotNull List actualTypes, - Option @NotNull ... options) { - PyTypeParameterMapping mapping = PyTypeParameterMapping.mapByShape(expectedTypes, actualTypes, options); - if (mapping != null) { - for (Couple pair : mapping.getMappedTypes()) { - if (pair.getFirst() instanceof PyTypeVarType typeVar) { - substitutions.typeVars.put(typeVar, Ref.create(pair.getSecond())); - } - else if (pair.getFirst() instanceof PyTypeVarTupleType typeVarTuple) { - substitutions.typeVarTuples.put(typeVarTuple, as(pair.getSecond(), PyPositionalVariadicType.class)); - } - else if (pair.getFirst() instanceof PyParamSpecType paramSpec) { - substitutions.paramSpecs.put(paramSpec, as(pair.getSecond(), PyCallableParameterVariadicType.class)); - } + private fun mapTypeParametersToSubstitutions( + substitutions: GenericSubstitutions, + expectedTypes: List, + actualTypes: List, + vararg options: PyTypeParameterMapping.Option, + ): GenericSubstitutions? { + val mapping = PyTypeParameterMapping.mapByShape(expectedTypes, actualTypes, *options) ?: return null + for (pair in mapping.mappedTypes) { + val first = pair.first + val second = pair.second + when (first) { + is PyTypeVarType -> substitutions.putTypeVar(first, Ref(second), KeyImpl) + is PyTypeVarTupleType -> substitutions.putTypeVarTuple(first, second as? PyPositionalVariadicType, KeyImpl) + is PyParamSpecType -> substitutions.putParamSpec(first, second as? PyCallableParameterVariadicType, KeyImpl) } - return substitutions; } - return null; + return substitutions } + @JvmStatic @ApiStatus.Internal - public static @Nullable PyType convertToType(@Nullable PyType type, @NotNull PyClassType superType, @NotNull TypeEvalContext context) { - MatchContext matchContext = new MatchContext(context, new GenericSubstitutions(), false); - Optional matched = match(superType, type, matchContext); - if (matched.orElse(false)) { - return substitute(superType, matchContext.mySubstitutions, context); + fun convertToType(type: PyType?, superType: PyClassType, context: TypeEvalContext): PyType? { + val matchContext = MatchContext(context, GenericSubstitutions(), false) + val matched = match(superType, type, matchContext) + if (matched.orElse(false)!!) { + return substitute(superType, matchContext.mySubstitutions, context) } - return null; + return null } + private class MyMatchHelper(private val myKey: Any, private val myContext: MatchContext) { + fun match(expected: PyType?, actual: PyType?): Boolean { + val result = RecursionManager.doPreventingRecursion( + myKey, false, if (myContext.reversedSubstitutions) + Computable { match(actual, expected, myContext) } + else + Computable { match(expected, actual, myContext) } + ) + return result != null && result.orElse(false)!! + } + } + + fun Generics(): Generics = GenericsImpl() + @ApiStatus.Internal - public static class Generics { - private final @NotNull Set typeVars = new LinkedHashSet<>(); + sealed interface Generics { + val typeVars: Set + val typeVarTuples: Set + val allTypeParameters: List + val paramSpecs: Set + val self: PySelfType? + val isEmpty: Boolean + } - private final @NotNull Set typeVarTuples = new LinkedHashSet<>(); - - private final @NotNull List allTypeParameters = new ArrayList<>(); - - private final @NotNull Set paramSpecs = new LinkedHashSet<>(); - - private @Nullable PySelfType self; - - public @NotNull Set getTypeVars() { - return Collections.unmodifiableSet(typeVars); - } - - public @NotNull Set getTypeVarTuples() { - return Collections.unmodifiableSet(typeVarTuples); - } - - public @NotNull List getAllTypeParameters() { - return Collections.unmodifiableList(allTypeParameters); - } - - public @NotNull Set getParamSpecs() { - return Collections.unmodifiableSet(paramSpecs); - } - - public boolean isEmpty() { - return typeVars.isEmpty() && typeVarTuples.isEmpty() && paramSpecs.isEmpty() && self == null; - } - - @Override - public String toString() { - return "Generics{" + - "typeVars=" + typeVars + - ", typeVarTuples" + typeVarTuples + - ", paramSpecs=" + paramSpecs + - '}'; - } + private data class GenericsImpl( + override val typeVars: MutableSet = LinkedHashSet(), + override val typeVarTuples: MutableSet = LinkedHashSet(), + override val allTypeParameters: MutableList = ArrayList(), + override val paramSpecs: MutableSet = LinkedHashSet(), + override var self: PySelfType? = null, + ) : Generics { + override val isEmpty: Boolean get() = typeVars.isEmpty() && typeVarTuples.isEmpty() && paramSpecs.isEmpty() && self == null } @ApiStatus.Experimental - public static class GenericSubstitutions { - private final @NotNull Map> typeVars; + class GenericSubstitutions() { + private val myTypeVars: MutableMap?> = LinkedHashMap() + val typeVars: Map?> + get() = Collections.unmodifiableMap(myTypeVars) - private final @NotNull Map typeVarTuples; + private val myTypeVarTuples: MutableMap = LinkedHashMap() + val typeVarTuples: Map + get() = Collections.unmodifiableMap(myTypeVarTuples) - private final @NotNull Map paramSpecs; + private val myParamSpecs: MutableMap = LinkedHashMap() + val paramSpecs: Map + get() = Collections.unmodifiableMap(myParamSpecs) - private @Nullable PyType qualifierType; + var qualifierType: PyType? = null + private var frozenTypeVars: Set = emptySet() - private @NotNull Set frozenTypeVars = Collections.emptySet(); - - public GenericSubstitutions(@NotNull Map typeParameters) { - this( - EntryStream.of(typeParameters) - .selectKeys(PyTypeVarType.class) - .mapValues(Ref::create) - .toCustomMap(LinkedHashMap::new), - EntryStream.of(typeParameters) - .selectKeys(PyTypeVarTupleType.class) - .selectValues(PyPositionalVariadicType.class) - .toCustomMap(LinkedHashMap::new), - EntryStream.of(typeParameters) - .selectKeys(PyParamSpecType.class) - .selectValues(PyCallableParameterVariadicType.class) - .toCustomMap(LinkedHashMap::new), - null - ); + constructor(typeParameters: Map) : this() { + for ((key, value) in typeParameters) { + when (key) { + is PyTypeVarType -> myTypeVars[key] = Ref(value) + is PyTypeVarTupleType -> if (value is PyPositionalVariadicType) myTypeVarTuples[key] = value + is PyParamSpecType -> if (value is PyCallableParameterVariadicType) myParamSpecs[key] = value + } + } } - public GenericSubstitutions() { - this(new LinkedHashMap<>(), new LinkedHashMap<>(), new LinkedHashMap<>(), null); + @ApiStatus.Internal + fun getFrozenTypeVars(@Suppress("unused") key: Key): Set = frozenTypeVars + + @ApiStatus.Internal + fun setFrozenTypeVars(value: Set, @Suppress("unused") key: Key) { + frozenTypeVars = value } - private GenericSubstitutions(@NotNull Map> typeVars, - @NotNull Map typeVarTuples, - @NotNull Map paramSpecs, - @Nullable PyType qualifierType) { - this.typeVars = typeVars; - this.typeVarTuples = typeVarTuples; - this.paramSpecs = paramSpecs; - this.qualifierType = qualifierType; + @ApiStatus.Internal + fun putTypeVar(typeVar: PyTypeVarType, substitute: Ref?, @Suppress("unused") key: Key, ifAbsent: Boolean = false) { + if (ifAbsent) myTypeVars.putIfAbsent(typeVar, substitute) + else myTypeVars[typeVar] = substitute } - public @NotNull Map getParamSpecs() { - return Collections.unmodifiableMap(paramSpecs); + @ApiStatus.Internal + fun putTypeVarTuple( + typeVarTuple: PyTypeVarTupleType, + substitute: PyPositionalVariadicType?, + @Suppress("unused") key: Key, + ifAbsent: Boolean = false, + ) { + if (ifAbsent) myTypeVarTuples.putIfAbsent(typeVarTuple, substitute) + else myTypeVarTuples[typeVarTuple] = substitute } - public @NotNull Map> getTypeVars() { - return Collections.unmodifiableMap(typeVars); + @ApiStatus.Internal + fun putParamSpec( + paramSpec: PyParamSpecType, + substitute: PyCallableParameterVariadicType?, + @Suppress("unused") key: Key, + ifAbsent: Boolean = false, + ) { + if (ifAbsent) myParamSpecs.putIfAbsent(paramSpec, substitute) + else myParamSpecs[paramSpec] = substitute } - public @NotNull Map getTypeVarTuples() { - return Collections.unmodifiableMap(typeVarTuples); - } - - public @Nullable PyType getQualifierType() { - return qualifierType; - } - - @Override - public String toString() { - return "GenericSubstitutions{" + - "typeVars=" + typeVars + - ", typeVarTuples" + typeVarTuples + - ", paramSpecs=" + paramSpecs + - '}'; - } + override fun toString(): String = + "GenericSubstitutions(typeVars=$typeVars, typeVarTuples$typeVarTuples, paramSpecs=$paramSpecs)" } - private static class MatchContext { + sealed class Key + private object KeyImpl : Key() - private final @NotNull TypeEvalContext context; - - private final @NotNull GenericSubstitutions mySubstitutions; - - private final boolean reversedSubstitutions; - - MatchContext(@NotNull TypeEvalContext context, @NotNull GenericSubstitutions substitutions, boolean reversedSubstitutions) { - this.context = context; - mySubstitutions = substitutions; - this.reversedSubstitutions = reversedSubstitutions; + private class MatchContext( + val context: TypeEvalContext, + val mySubstitutions: GenericSubstitutions, + val reversedSubstitutions: Boolean, + ) { + fun reverseSubstitutions(): MatchContext { + return MatchContext(context, mySubstitutions, !reversedSubstitutions) } - public @NotNull MatchContext reverseSubstitutions() { - return new MatchContext(context, mySubstitutions, !reversedSubstitutions); - } - - public @NotNull MatchContext resetSubstitutions() { - return new MatchContext(context, mySubstitutions, false); + fun resetSubstitutions(): MatchContext { + return MatchContext(context, mySubstitutions, false) } } }