Generate assertions for expressions with enhanced nullability
If an expression with type annotated with @EnhancedNullability is used as a function expression body, or property initializer, or variable initializer, and corresponding type can not contain null, generate nullability assertions for this expression.
This commit is contained in:
@@ -51,6 +51,7 @@ import org.jetbrains.kotlin.codegen.state.KotlinTypeMapper;
|
|||||||
import org.jetbrains.kotlin.codegen.when.SwitchCodegen;
|
import org.jetbrains.kotlin.codegen.when.SwitchCodegen;
|
||||||
import org.jetbrains.kotlin.codegen.when.SwitchCodegenProvider;
|
import org.jetbrains.kotlin.codegen.when.SwitchCodegenProvider;
|
||||||
import org.jetbrains.kotlin.config.ApiVersion;
|
import org.jetbrains.kotlin.config.ApiVersion;
|
||||||
|
import org.jetbrains.kotlin.config.LanguageFeature;
|
||||||
import org.jetbrains.kotlin.descriptors.*;
|
import org.jetbrains.kotlin.descriptors.*;
|
||||||
import org.jetbrains.kotlin.descriptors.impl.AnonymousFunctionDescriptor;
|
import org.jetbrains.kotlin.descriptors.impl.AnonymousFunctionDescriptor;
|
||||||
import org.jetbrains.kotlin.descriptors.impl.LocalVariableDescriptor;
|
import org.jetbrains.kotlin.descriptors.impl.LocalVariableDescriptor;
|
||||||
@@ -307,7 +308,12 @@ public class ExpressionCodegen extends KtVisitor<StackValue, StackValue> impleme
|
|||||||
|
|
||||||
RuntimeAssertionInfo runtimeAssertionInfo = null;
|
RuntimeAssertionInfo runtimeAssertionInfo = null;
|
||||||
if (selector instanceof KtExpression) {
|
if (selector instanceof KtExpression) {
|
||||||
runtimeAssertionInfo = bindingContext.get(JvmBindingContextSlices.RUNTIME_ASSERTION_INFO, (KtExpression) selector);
|
KtExpression expression = (KtExpression) selector;
|
||||||
|
runtimeAssertionInfo = bindingContext.get(JvmBindingContextSlices.RUNTIME_ASSERTION_INFO, expression);
|
||||||
|
if (runtimeAssertionInfo == null &&
|
||||||
|
state.getLanguageVersionSettings().supportsFeature(LanguageFeature.StrictJavaNullabilityAssertions)) {
|
||||||
|
runtimeAssertionInfo = bindingContext.get(JvmBindingContextSlices.BODY_RUNTIME_ASSERTION_INFO, expression);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (BuiltinSpecialBridgesKt.isValueArgumentForCallToMethodWithTypeCheckBarrier(selector, bindingContext)) return stackValue;
|
if (BuiltinSpecialBridgesKt.isValueArgumentForCallToMethodWithTypeCheckBarrier(selector, bindingContext)) return stackValue;
|
||||||
|
|||||||
+28
@@ -32,6 +32,7 @@ import org.jetbrains.kotlin.codegen.state.GenerationState;
|
|||||||
import org.jetbrains.kotlin.codegen.state.TypeMapperUtilsKt;
|
import org.jetbrains.kotlin.codegen.state.TypeMapperUtilsKt;
|
||||||
import org.jetbrains.kotlin.codegen.when.SwitchCodegenProvider;
|
import org.jetbrains.kotlin.codegen.when.SwitchCodegenProvider;
|
||||||
import org.jetbrains.kotlin.codegen.when.WhenByEnumsMapping;
|
import org.jetbrains.kotlin.codegen.when.WhenByEnumsMapping;
|
||||||
|
import org.jetbrains.kotlin.config.LanguageVersionSettings;
|
||||||
import org.jetbrains.kotlin.coroutines.CoroutineUtilKt;
|
import org.jetbrains.kotlin.coroutines.CoroutineUtilKt;
|
||||||
import org.jetbrains.kotlin.descriptors.*;
|
import org.jetbrains.kotlin.descriptors.*;
|
||||||
import org.jetbrains.kotlin.descriptors.annotations.Annotations;
|
import org.jetbrains.kotlin.descriptors.annotations.Annotations;
|
||||||
@@ -56,6 +57,7 @@ import org.jetbrains.kotlin.resolve.constants.ConstantValue;
|
|||||||
import org.jetbrains.kotlin.resolve.constants.EnumValue;
|
import org.jetbrains.kotlin.resolve.constants.EnumValue;
|
||||||
import org.jetbrains.kotlin.resolve.constants.NullValue;
|
import org.jetbrains.kotlin.resolve.constants.NullValue;
|
||||||
import org.jetbrains.kotlin.resolve.descriptorUtil.DescriptorUtilsKt;
|
import org.jetbrains.kotlin.resolve.descriptorUtil.DescriptorUtilsKt;
|
||||||
|
import org.jetbrains.kotlin.resolve.jvm.RuntimeAssertionsOnDeclarationBodyChecker;
|
||||||
import org.jetbrains.kotlin.resolve.scopes.receivers.ReceiverValue;
|
import org.jetbrains.kotlin.resolve.scopes.receivers.ReceiverValue;
|
||||||
import org.jetbrains.kotlin.resolve.scopes.receivers.TransientReceiver;
|
import org.jetbrains.kotlin.resolve.scopes.receivers.TransientReceiver;
|
||||||
import org.jetbrains.kotlin.types.KotlinType;
|
import org.jetbrains.kotlin.types.KotlinType;
|
||||||
@@ -86,6 +88,8 @@ class CodegenAnnotatingVisitor extends KtVisitorVoid {
|
|||||||
private final JvmRuntimeTypes runtimeTypes;
|
private final JvmRuntimeTypes runtimeTypes;
|
||||||
private final TypeMappingConfiguration<Type> typeMappingConfiguration;
|
private final TypeMappingConfiguration<Type> typeMappingConfiguration;
|
||||||
private final SwitchCodegenProvider switchCodegenProvider;
|
private final SwitchCodegenProvider switchCodegenProvider;
|
||||||
|
private final LanguageVersionSettings languageVersionSettings;
|
||||||
|
private final ClassBuilderMode classBuilderMode;
|
||||||
|
|
||||||
public CodegenAnnotatingVisitor(@NotNull GenerationState state) {
|
public CodegenAnnotatingVisitor(@NotNull GenerationState state) {
|
||||||
this.bindingTrace = state.getBindingTrace();
|
this.bindingTrace = state.getBindingTrace();
|
||||||
@@ -94,6 +98,8 @@ class CodegenAnnotatingVisitor extends KtVisitorVoid {
|
|||||||
this.runtimeTypes = state.getJvmRuntimeTypes();
|
this.runtimeTypes = state.getJvmRuntimeTypes();
|
||||||
this.typeMappingConfiguration = state.getTypeMapper().getTypeMappingConfiguration();
|
this.typeMappingConfiguration = state.getTypeMapper().getTypeMappingConfiguration();
|
||||||
this.switchCodegenProvider = new SwitchCodegenProvider(state);
|
this.switchCodegenProvider = new SwitchCodegenProvider(state);
|
||||||
|
this.languageVersionSettings = state.getLanguageVersionSettings();
|
||||||
|
this.classBuilderMode = state.getClassBuilderMode();
|
||||||
}
|
}
|
||||||
|
|
||||||
@NotNull
|
@NotNull
|
||||||
@@ -411,6 +417,8 @@ class CodegenAnnotatingVisitor extends KtVisitorVoid {
|
|||||||
// working around a problem with shallow analysis
|
// working around a problem with shallow analysis
|
||||||
if (descriptor == null) return;
|
if (descriptor == null) return;
|
||||||
|
|
||||||
|
checkRuntimeAsserionsOnDeclarationBody(property, descriptor);
|
||||||
|
|
||||||
if (descriptor instanceof LocalVariableDescriptor) {
|
if (descriptor instanceof LocalVariableDescriptor) {
|
||||||
recordLocalVariablePropertyMetadata((LocalVariableDescriptor) descriptor);
|
recordLocalVariablePropertyMetadata((LocalVariableDescriptor) descriptor);
|
||||||
}
|
}
|
||||||
@@ -447,6 +455,14 @@ class CodegenAnnotatingVisitor extends KtVisitorVoid {
|
|||||||
nameStack.pop();
|
nameStack.pop();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private void checkRuntimeAsserionsOnDeclarationBody(@NotNull KtDeclaration declaration, DeclarationDescriptor descriptor) {
|
||||||
|
if (classBuilderMode.generateBodies) {
|
||||||
|
// NB This is required only for bodies generation.
|
||||||
|
// In light class generation can cause recursion in types resolution.
|
||||||
|
RuntimeAssertionsOnDeclarationBodyChecker.check(declaration, descriptor, bindingTrace, languageVersionSettings);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
@NotNull
|
@NotNull
|
||||||
private Type getMetadataOwner(@NotNull KtProperty property) {
|
private Type getMetadataOwner(@NotNull KtProperty property) {
|
||||||
for (int i = classStack.size() - 1; i >= 0; i--) {
|
for (int i = classStack.size() - 1; i >= 0; i--) {
|
||||||
@@ -469,12 +485,24 @@ class CodegenAnnotatingVisitor extends KtVisitorVoid {
|
|||||||
return Type.getObjectType(JvmFileClassUtil.getFileClassInternalName(property.getContainingKtFile()));
|
return Type.getObjectType(JvmFileClassUtil.getFileClassInternalName(property.getContainingKtFile()));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@Override
|
||||||
|
public void visitPropertyAccessor(@NotNull KtPropertyAccessor accessor) {
|
||||||
|
PropertyAccessorDescriptor accessorDescriptor = bindingContext.get(PROPERTY_ACCESSOR, accessor);
|
||||||
|
if (accessorDescriptor != null) {
|
||||||
|
checkRuntimeAsserionsOnDeclarationBody(accessor, accessorDescriptor);
|
||||||
|
}
|
||||||
|
|
||||||
|
super.visitPropertyAccessor(accessor);
|
||||||
|
}
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
public void visitNamedFunction(@NotNull KtNamedFunction function) {
|
public void visitNamedFunction(@NotNull KtNamedFunction function) {
|
||||||
FunctionDescriptor functionDescriptor = (FunctionDescriptor) bindingContext.get(DECLARATION_TO_DESCRIPTOR, function);
|
FunctionDescriptor functionDescriptor = (FunctionDescriptor) bindingContext.get(DECLARATION_TO_DESCRIPTOR, function);
|
||||||
// working around a problem with shallow analysis
|
// working around a problem with shallow analysis
|
||||||
if (functionDescriptor == null) return;
|
if (functionDescriptor == null) return;
|
||||||
|
|
||||||
|
checkRuntimeAsserionsOnDeclarationBody(function, functionDescriptor);
|
||||||
|
|
||||||
String nameForClassOrPackageMember = getNameForClassOrPackageMember(functionDescriptor);
|
String nameForClassOrPackageMember = getNameForClassOrPackageMember(functionDescriptor);
|
||||||
|
|
||||||
if (functionDescriptor instanceof SimpleFunctionDescriptor && functionDescriptor.isSuspend()) {
|
if (functionDescriptor instanceof SimpleFunctionDescriptor && functionDescriptor.isSuspend()) {
|
||||||
|
|||||||
@@ -32,6 +32,9 @@ object JvmBindingContextSlices {
|
|||||||
@JvmField
|
@JvmField
|
||||||
val RECEIVER_RUNTIME_ASSERTION_INFO: WritableSlice<ExpressionReceiver, RuntimeAssertionInfo> = BasicWritableSlice(RewritePolicy.DO_NOTHING)
|
val RECEIVER_RUNTIME_ASSERTION_INFO: WritableSlice<ExpressionReceiver, RuntimeAssertionInfo> = BasicWritableSlice(RewritePolicy.DO_NOTHING)
|
||||||
|
|
||||||
|
@JvmField
|
||||||
|
val BODY_RUNTIME_ASSERTION_INFO: WritableSlice<KtExpression, RuntimeAssertionInfo> = BasicWritableSlice(RewritePolicy.DO_NOTHING)
|
||||||
|
|
||||||
@JvmField
|
@JvmField
|
||||||
val LOAD_FROM_JAVA_SIGNATURE_ERRORS: WritableSlice<DeclarationDescriptor, List<String>> = Slices.createCollectiveSlice()
|
val LOAD_FROM_JAVA_SIGNATURE_ERRORS: WritableSlice<DeclarationDescriptor, List<String>> = Slices.createCollectiveSlice()
|
||||||
|
|
||||||
|
|||||||
@@ -18,9 +18,12 @@ package org.jetbrains.kotlin.resolve.jvm
|
|||||||
|
|
||||||
import com.intellij.openapi.util.text.StringUtil
|
import com.intellij.openapi.util.text.StringUtil
|
||||||
import com.intellij.psi.PsiElement
|
import com.intellij.psi.PsiElement
|
||||||
import org.jetbrains.kotlin.descriptors.ReceiverParameterDescriptor
|
import org.jetbrains.kotlin.config.LanguageFeature
|
||||||
|
import org.jetbrains.kotlin.config.LanguageVersionSettings
|
||||||
|
import org.jetbrains.kotlin.descriptors.*
|
||||||
import org.jetbrains.kotlin.load.java.typeEnhancement.hasEnhancedNullability
|
import org.jetbrains.kotlin.load.java.typeEnhancement.hasEnhancedNullability
|
||||||
import org.jetbrains.kotlin.psi.KtExpression
|
import org.jetbrains.kotlin.psi.*
|
||||||
|
import org.jetbrains.kotlin.resolve.BindingTrace
|
||||||
import org.jetbrains.kotlin.resolve.calls.callUtil.isSafeCall
|
import org.jetbrains.kotlin.resolve.calls.callUtil.isSafeCall
|
||||||
import org.jetbrains.kotlin.resolve.calls.checkers.AdditionalTypeChecker
|
import org.jetbrains.kotlin.resolve.calls.checkers.AdditionalTypeChecker
|
||||||
import org.jetbrains.kotlin.resolve.calls.checkers.CallChecker
|
import org.jetbrains.kotlin.resolve.calls.checkers.CallChecker
|
||||||
@@ -31,9 +34,9 @@ import org.jetbrains.kotlin.resolve.calls.smartcasts.DataFlowValue
|
|||||||
import org.jetbrains.kotlin.resolve.calls.smartcasts.DataFlowValueFactory
|
import org.jetbrains.kotlin.resolve.calls.smartcasts.DataFlowValueFactory
|
||||||
import org.jetbrains.kotlin.resolve.scopes.receivers.ExpressionReceiver
|
import org.jetbrains.kotlin.resolve.scopes.receivers.ExpressionReceiver
|
||||||
import org.jetbrains.kotlin.resolve.scopes.receivers.ReceiverValue
|
import org.jetbrains.kotlin.resolve.scopes.receivers.ReceiverValue
|
||||||
import org.jetbrains.kotlin.types.KotlinType
|
import org.jetbrains.kotlin.types.*
|
||||||
import org.jetbrains.kotlin.types.TypeUtils
|
import org.jetbrains.kotlin.types.checker.isClassType
|
||||||
import org.jetbrains.kotlin.types.isError
|
import org.jetbrains.kotlin.types.typeUtil.immediateSupertypes
|
||||||
import org.jetbrains.kotlin.utils.addToStdlib.safeAs
|
import org.jetbrains.kotlin.utils.addToStdlib.safeAs
|
||||||
|
|
||||||
class RuntimeAssertionInfo(val needNotNullAssertion: Boolean, val message: String) {
|
class RuntimeAssertionInfo(val needNotNullAssertion: Boolean, val message: String) {
|
||||||
@@ -83,6 +86,9 @@ class RuntimeAssertionInfo(val needNotNullAssertion: Boolean, val message: Strin
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private val KtExpression.textForRuntimeAssertionInfo
|
||||||
|
get() = StringUtil.trimMiddle(text, 50)
|
||||||
|
|
||||||
class RuntimeAssertionsDataFlowExtras(
|
class RuntimeAssertionsDataFlowExtras(
|
||||||
private val c: ResolutionContext<*>,
|
private val c: ResolutionContext<*>,
|
||||||
private val dataFlowValue: DataFlowValue,
|
private val dataFlowValue: DataFlowValue,
|
||||||
@@ -93,7 +99,7 @@ class RuntimeAssertionsDataFlowExtras(
|
|||||||
override val possibleTypes: Set<KotlinType>
|
override val possibleTypes: Set<KotlinType>
|
||||||
get() = c.dataFlowInfo.getCollectedTypes(dataFlowValue)
|
get() = c.dataFlowInfo.getCollectedTypes(dataFlowValue)
|
||||||
override val presentableText: String
|
override val presentableText: String
|
||||||
get() = StringUtil.trimMiddle(expression.text, 50)
|
get() = expression.textForRuntimeAssertionInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
object RuntimeAssertionsTypeChecker : AdditionalTypeChecker {
|
object RuntimeAssertionsTypeChecker : AdditionalTypeChecker {
|
||||||
@@ -139,3 +145,100 @@ object RuntimeAssertionsOnExtensionReceiverCallChecker : CallChecker {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
object RuntimeAssertionsOnDeclarationBodyChecker {
|
||||||
|
@JvmStatic
|
||||||
|
fun check(
|
||||||
|
declaration: KtDeclaration,
|
||||||
|
descriptor: DeclarationDescriptor,
|
||||||
|
bindingTrace: BindingTrace,
|
||||||
|
languageVersionSettings: LanguageVersionSettings
|
||||||
|
) {
|
||||||
|
if (!languageVersionSettings.supportsFeature(LanguageFeature.StrictJavaNullabilityAssertions)) return
|
||||||
|
|
||||||
|
when {
|
||||||
|
declaration is KtProperty && descriptor is VariableDescriptor ->
|
||||||
|
checkLocalVariable(declaration, descriptor, bindingTrace)
|
||||||
|
declaration is KtFunction && descriptor is FunctionDescriptor ->
|
||||||
|
checkFunction(declaration, descriptor, bindingTrace)
|
||||||
|
declaration is KtProperty && descriptor is PropertyDescriptor ->
|
||||||
|
checkProperty(declaration, descriptor, bindingTrace)
|
||||||
|
declaration is KtPropertyAccessor && descriptor is PropertyAccessorDescriptor ->
|
||||||
|
checkPropertyAccessor(declaration, descriptor, bindingTrace)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun checkLocalVariable(
|
||||||
|
declaration: KtProperty,
|
||||||
|
descriptor: VariableDescriptor,
|
||||||
|
bindingTrace: BindingTrace
|
||||||
|
) {
|
||||||
|
if (declaration.typeReference != null) return
|
||||||
|
|
||||||
|
checkNullabilityAssertion(declaration.initializer ?: return, descriptor.type, bindingTrace)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun checkFunction(
|
||||||
|
declaration: KtFunction,
|
||||||
|
descriptor: FunctionDescriptor,
|
||||||
|
bindingTrace: BindingTrace
|
||||||
|
) {
|
||||||
|
if (declaration.typeReference != null || declaration.hasBlockBody()) return
|
||||||
|
|
||||||
|
checkNullabilityAssertion(declaration.bodyExpression ?: return, descriptor.returnType ?: return,
|
||||||
|
bindingTrace)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun checkProperty(
|
||||||
|
declaration: KtProperty,
|
||||||
|
descriptor: PropertyDescriptor,
|
||||||
|
bindingTrace: BindingTrace
|
||||||
|
) {
|
||||||
|
if (declaration.typeReference != null) return
|
||||||
|
|
||||||
|
// TODO nullability assertion on delegate initialization expression, see KT-20823
|
||||||
|
if (declaration.hasDelegateExpression()) return
|
||||||
|
|
||||||
|
checkNullabilityAssertion(declaration.initializer ?: return, descriptor.type, bindingTrace)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun checkPropertyAccessor(
|
||||||
|
declaration: KtPropertyAccessor,
|
||||||
|
descriptor: PropertyAccessorDescriptor,
|
||||||
|
bindingTrace: BindingTrace
|
||||||
|
) {
|
||||||
|
if (declaration.property.typeReference != null || declaration.hasBlockBody()) return
|
||||||
|
|
||||||
|
checkNullabilityAssertion(declaration.bodyExpression ?: return, descriptor.correspondingProperty.type,
|
||||||
|
bindingTrace)
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
private fun checkNullabilityAssertion(
|
||||||
|
expression: KtExpression,
|
||||||
|
declarationType: KotlinType,
|
||||||
|
bindingTrace: BindingTrace
|
||||||
|
) {
|
||||||
|
if (declarationType.unwrap().canContainNull()) return
|
||||||
|
|
||||||
|
val expressionType = bindingTrace.getType(expression) ?: return
|
||||||
|
if (expressionType.isError) return
|
||||||
|
|
||||||
|
if (!expressionType.hasEnhancedNullability()) return
|
||||||
|
|
||||||
|
bindingTrace.record(
|
||||||
|
JvmBindingContextSlices.BODY_RUNTIME_ASSERTION_INFO,
|
||||||
|
expression,
|
||||||
|
RuntimeAssertionInfo(true, expression.textForRuntimeAssertionInfo)
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun UnwrappedType.canContainNull(): Boolean {
|
||||||
|
val upper = upperIfFlexible()
|
||||||
|
return when {
|
||||||
|
upper.isMarkedNullable -> true
|
||||||
|
upper.isClassType -> false
|
||||||
|
else -> upper.immediateSupertypes().all { it.unwrap().canContainNull() }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user