Create SimpleFunctionDescriptorImpl under nonCancelableSection

SimpleFunctionDescriptorImpl initialization consists of two phases:
ctor + initialize.
When SimpleFunctionDescriptorImpl is created wrapped descriptor
(e.g. ValueParameterDescriptorImpl) is leaked through bindingTrace
with not fully initialized `containingDeclaration` (that is
SimpleFunctionDescriptorImpl).
If PCE happens after this unsafe publication prior to `initialize` then
it will be case with NPE on fully initialized instance reading.

#KT-56364 Fixed
This commit is contained in:
Vladimir Dolzhenko
2023-02-02 17:15:12 +01:00
committed by Space Team
parent 9be4aa2e02
commit a049fda75b
2 changed files with 210 additions and 175 deletions
@@ -17,6 +17,8 @@
package org.jetbrains.kotlin.resolve package org.jetbrains.kotlin.resolve
import com.google.common.collect.HashMultimap import com.google.common.collect.HashMultimap
import com.intellij.openapi.diagnostic.ControlFlowException
import com.intellij.openapi.progress.ProgressManager
import com.intellij.openapi.util.ThrowableComputable import com.intellij.openapi.util.ThrowableComputable
import com.intellij.psi.PsiElement import com.intellij.psi.PsiElement
import com.intellij.util.AstLoadingFilter import com.intellij.util.AstLoadingFilter
@@ -71,6 +73,7 @@ import org.jetbrains.kotlin.types.isError
import org.jetbrains.kotlin.types.typeUtil.replaceAnnotations import org.jetbrains.kotlin.types.typeUtil.replaceAnnotations
import java.util.* import java.util.*
class FunctionDescriptorResolver( class FunctionDescriptorResolver(
private val typeResolver: TypeResolver, private val typeResolver: TypeResolver,
private val descriptorResolver: DescriptorResolver, private val descriptorResolver: DescriptorResolver,
@@ -129,6 +132,7 @@ class FunctionDescriptorResolver(
CallableMemberDescriptor.Kind.DECLARATION, CallableMemberDescriptor.Kind.DECLARATION,
function.toSourceElement() function.toSourceElement()
) )
return computeInNonCancelableSection {
initializeFunctionDescriptorAndExplicitReturnType( initializeFunctionDescriptorAndExplicitReturnType(
containingDescriptor, containingDescriptor,
scope, scope,
@@ -141,7 +145,8 @@ class FunctionDescriptorResolver(
) )
initializeFunctionReturnTypeBasedOnFunctionBody(scope, function, functionDescriptor, trace, dataFlowInfo, inferenceSession) initializeFunctionReturnTypeBasedOnFunctionBody(scope, function, functionDescriptor, trace, dataFlowInfo, inferenceSession)
BindingContextUtils.recordFunctionDeclarationToDescriptor(trace, function, functionDescriptor) BindingContextUtils.recordFunctionDeclarationToDescriptor(trace, function, functionDescriptor)
return functionDescriptor functionDescriptor
}
} }
private fun initializeFunctionReturnTypeBasedOnFunctionBody( private fun initializeFunctionReturnTypeBasedOnFunctionBody(
@@ -182,6 +187,7 @@ class FunctionDescriptorResolver(
dataFlowInfo: DataFlowInfo, dataFlowInfo: DataFlowInfo,
inferenceSession: InferenceSession? inferenceSession: InferenceSession?
) { ) {
try {
val headerScope = LexicalWritableScope( val headerScope = LexicalWritableScope(
scope, functionDescriptor, true, scope, functionDescriptor, true,
TraceBasedLocalRedeclarationChecker(trace, overloadChecker), LexicalScopeKind.FUNCTION_HEADER TraceBasedLocalRedeclarationChecker(trace, overloadChecker), LexicalScopeKind.FUNCTION_HEADER
@@ -294,6 +300,12 @@ class FunctionDescriptorResolver(
for (valueParameterDescriptor in valueParameterDescriptors) { for (valueParameterDescriptor in valueParameterDescriptors) {
ForceResolveUtil.forceResolveAllContents(valueParameterDescriptor.type.annotations) ForceResolveUtil.forceResolveAllContents(valueParameterDescriptor.type.annotations)
} }
} catch (e: Exception) {
if (e is ControlFlowException) {
throw IllegalStateException("Method should be run under nonCancelableSection", e)
}
throw e
}
} }
private fun getContractProvider( private fun getContractProvider(
@@ -444,8 +456,6 @@ class FunctionDescriptorResolver(
constructorDescriptor.isActual = modifierList?.hasActualModifier() == true || constructorDescriptor.isActual = modifierList?.hasActualModifier() == true ||
// We don't require 'actual' for constructors of actual annotations // We don't require 'actual' for constructors of actual annotations
classDescriptor.kind == ClassKind.ANNOTATION_CLASS && classDescriptor.isActual classDescriptor.kind == ClassKind.ANNOTATION_CLASS && classDescriptor.isActual
if (declarationToTrace is PsiElement)
trace.record(BindingContext.CONSTRUCTOR, declarationToTrace, constructorDescriptor)
val parameterScope = LexicalWritableScope( val parameterScope = LexicalWritableScope(
scope, scope,
constructorDescriptor, constructorDescriptor,
@@ -453,20 +463,27 @@ class FunctionDescriptorResolver(
TraceBasedLocalRedeclarationChecker(trace, overloadChecker), TraceBasedLocalRedeclarationChecker(trace, overloadChecker),
LexicalScopeKind.CONSTRUCTOR_HEADER LexicalScopeKind.CONSTRUCTOR_HEADER
) )
return computeInNonCancelableSection {
if (declarationToTrace is PsiElement)
trace.record(BindingContext.CONSTRUCTOR, declarationToTrace, constructorDescriptor)
val constructor = constructorDescriptor.initialize( val constructor = constructorDescriptor.initialize(
resolveValueParameters( resolveValueParameters(
constructorDescriptor, parameterScope, valueParameters, trace, null, inferenceSession constructorDescriptor, parameterScope, valueParameters, trace, null, inferenceSession
), ),
resolveVisibilityFromModifiers( resolveVisibilityFromModifiers(
modifierList, modifierList,
DescriptorUtils.getDefaultConstructorVisibility(classDescriptor, languageVersionSettings.supportsFeature(LanguageFeature.AllowSealedInheritorsInDifferentFilesOfSamePackage)) DescriptorUtils.getDefaultConstructorVisibility(
classDescriptor,
languageVersionSettings.supportsFeature(LanguageFeature.AllowSealedInheritorsInDifferentFilesOfSamePackage)
)
) )
) )
constructor.returnType = classDescriptor.defaultType constructor.returnType = classDescriptor.defaultType
if (DescriptorUtils.isAnnotationClass(classDescriptor)) { if (DescriptorUtils.isAnnotationClass(classDescriptor)) {
CompileTimeConstantUtils.checkConstructorParametersType(valueParameters, trace) CompileTimeConstantUtils.checkConstructorParametersType(valueParameters, trace)
} }
return constructor constructor
}
} }
private fun resolveValueParameters( private fun resolveValueParameters(
@@ -477,6 +494,7 @@ class FunctionDescriptorResolver(
expectedParameterTypes: List<KotlinType>?, expectedParameterTypes: List<KotlinType>?,
inferenceSession: InferenceSession? inferenceSession: InferenceSession?
): List<ValueParameterDescriptor> { ): List<ValueParameterDescriptor> {
try {
val result = ArrayList<ValueParameterDescriptor>() val result = ArrayList<ValueParameterDescriptor>()
for (i in valueParameters.indices) { for (i in valueParameters.indices) {
@@ -527,7 +545,16 @@ class FunctionDescriptorResolver(
result.add(valueParameterDescriptor) result.add(valueParameterDescriptor)
} }
return result return result
} catch (e: Exception) {
if (e is ControlFlowException) {
throw IllegalStateException("Method should be run under nonCancelableSection", e)
}
throw e
}
} }
private data class ContextReceiverTypeWithLabel(val type: KotlinType, val label: Name?) private data class ContextReceiverTypeWithLabel(val type: KotlinType, val label: Name?)
} }
private fun <T> computeInNonCancelableSection(action: () -> T): T =
ProgressManager.getInstance().computeInNonCancelableSection<T, Exception>(action)
@@ -6,6 +6,7 @@
package org.jetbrains.kotlin.types.expressions package org.jetbrains.kotlin.types.expressions
import com.google.common.collect.Lists import com.google.common.collect.Lists
import com.intellij.openapi.progress.ProgressManager
import com.intellij.psi.PsiElement import com.intellij.psi.PsiElement
import org.jetbrains.kotlin.builtins.* import org.jetbrains.kotlin.builtins.*
import org.jetbrains.kotlin.config.LanguageFeature import org.jetbrains.kotlin.config.LanguageFeature
@@ -209,9 +210,15 @@ internal class FunctionsTypingVisitor(facade: ExpressionTypingInternals) : Expre
context: ExpressionTypingContext context: ExpressionTypingContext
): AnonymousFunctionDescriptor { ): AnonymousFunctionDescriptor {
val functionLiteral = expression.functionLiteral val functionLiteral = expression.functionLiteral
val annotations = components.annotationResolver.resolveAnnotationsWithArguments(
context.scope,
expression.getAnnotationEntries(),
context.trace
)
return ProgressManager.getInstance().computeInNonCancelableSection<AnonymousFunctionDescriptor, Exception> {
val functionDescriptor = AnonymousFunctionDescriptor( val functionDescriptor = AnonymousFunctionDescriptor(
context.scope.ownerDescriptor, context.scope.ownerDescriptor,
components.annotationResolver.resolveAnnotationsWithArguments(context.scope, expression.getAnnotationEntries(), context.trace), annotations,
CallableMemberDescriptor.Kind.DECLARATION, functionLiteral.toSourceElement(), CallableMemberDescriptor.Kind.DECLARATION, functionLiteral.toSourceElement(),
context.expectedType.isSuspendFunctionType() context.expectedType.isSuspendFunctionType()
).let { ).let {
@@ -225,7 +232,8 @@ internal class FunctionsTypingVisitor(facade: ExpressionTypingInternals) : Expre
ForceResolveUtil.forceResolveAllContents(parameterDescriptor.annotations) ForceResolveUtil.forceResolveAllContents(parameterDescriptor.annotations)
} }
BindingContextUtils.recordFunctionDeclarationToDescriptor(context.trace, functionLiteral, functionDescriptor) BindingContextUtils.recordFunctionDeclarationToDescriptor(context.trace, functionLiteral, functionDescriptor)
return functionDescriptor functionDescriptor
}
} }
private fun KotlinType.isBuiltinFunctionalType() = private fun KotlinType.isBuiltinFunctionalType() =