Refactor logic of reporting errors in contracts

Rename 'doCheckContract' into 'parseContractAndReportErrors',
emphasizing that it works with definitely call to 'contract' from
stdlib.
In particular, it means that it should either return non-null
value or report some errors.

Note that this commit doesn't change behavior of this code (modulo cases
when something in 'parseContract' throws exception), but just refactors
the code to be more clear and easy to reason about.
This commit is contained in:
Dmitry Savvinov
2018-10-23 18:10:33 +03:00
parent 3dd4e9f04f
commit 339c55505a
2 changed files with 45 additions and 27 deletions
@@ -7,7 +7,6 @@ package org.jetbrains.kotlin.contracts.parsing
import org.jetbrains.kotlin.config.LanguageFeature import org.jetbrains.kotlin.config.LanguageFeature
import org.jetbrains.kotlin.config.LanguageVersionSettings import org.jetbrains.kotlin.config.LanguageVersionSettings
import org.jetbrains.kotlin.config.LanguageVersionSettingsImpl
import org.jetbrains.kotlin.diagnostics.Diagnostic import org.jetbrains.kotlin.diagnostics.Diagnostic
import org.jetbrains.kotlin.diagnostics.Errors import org.jetbrains.kotlin.diagnostics.Errors
import org.jetbrains.kotlin.psi.KtCallExpression import org.jetbrains.kotlin.psi.KtCallExpression
@@ -19,49 +18,52 @@ interface ContractParsingDiagnosticsCollector {
fun unsupportedFeature(languageVersionSettings: LanguageVersionSettings) fun unsupportedFeature(languageVersionSettings: LanguageVersionSettings)
fun contractNotAllowed(message: String) fun contractNotAllowed(message: String)
fun badDescription(message: String, reportOn: KtElement) fun badDescription(message: String, reportOn: KtElement)
fun addFallbackErrorIfNecessary()
fun flushDiagnostics(parsingFailed: Boolean) fun flushDiagnostics()
fun hasErrors(): Boolean fun hasErrors(): Boolean
object EMPTY : ContractParsingDiagnosticsCollector { object EMPTY : ContractParsingDiagnosticsCollector {
override fun contractNotAllowed(message: String) {} override fun contractNotAllowed(message: String) {}
override fun badDescription(message: String, reportOn: KtElement) {} override fun badDescription(message: String, reportOn: KtElement) {}
override fun unsupportedFeature(languageVersionSettings: LanguageVersionSettings) { } override fun unsupportedFeature(languageVersionSettings: LanguageVersionSettings) {}
override fun addFallbackErrorIfNecessary() { }
override fun flushDiagnostics(parsingFailed: Boolean) {} override fun flushDiagnostics() {}
override fun hasErrors(): Boolean = false override fun hasErrors(): Boolean = false
} }
} }
class TraceBasedCollector(private val bindingTrace: BindingTrace, mainCall: KtExpression) : ContractParsingDiagnosticsCollector { class TraceBasedCollector(private val bindingTrace: BindingTrace, mainCall: KtExpression) : ContractParsingDiagnosticsCollector {
private val diagnostics: MutableList<Diagnostic> = mutableListOf() private val reportedErrors: MutableList<Diagnostic> = mutableListOf()
private val mainCallReportTarget = (mainCall as? KtCallExpression)?.calleeExpression ?: mainCall private val mainCallReportTarget = (mainCall as? KtCallExpression)?.calleeExpression ?: mainCall
override fun contractNotAllowed(message: String) { override fun contractNotAllowed(message: String) {
diagnostics += Errors.CONTRACT_NOT_ALLOWED.on(mainCallReportTarget, message) reportedErrors += Errors.CONTRACT_NOT_ALLOWED.on(mainCallReportTarget, message)
} }
override fun badDescription(message: String, reportOn: KtElement) { override fun badDescription(message: String, reportOn: KtElement) {
diagnostics += Errors.ERROR_IN_CONTRACT_DESCRIPTION.on(reportOn, message) reportedErrors += Errors.ERROR_IN_CONTRACT_DESCRIPTION.on(reportOn, message)
} }
override fun unsupportedFeature(languageVersionSettings: LanguageVersionSettings) { override fun unsupportedFeature(languageVersionSettings: LanguageVersionSettings) {
diagnostics += Errors.UNSUPPORTED_FEATURE.on( reportedErrors += Errors.UNSUPPORTED_FEATURE.on(
mainCallReportTarget, mainCallReportTarget,
LanguageFeature.AllowContractsForCustomFunctions to languageVersionSettings LanguageFeature.AllowContractsForCustomFunctions to languageVersionSettings
) )
} }
override fun flushDiagnostics(parsingFailed: Boolean) { override fun addFallbackErrorIfNecessary() {
if (parsingFailed && diagnostics.isEmpty()) { if (reportedErrors.isEmpty())
diagnostics += Errors.ERROR_IN_CONTRACT_DESCRIPTION.on(mainCallReportTarget, "Error in contract description") reportedErrors += Errors.ERROR_IN_CONTRACT_DESCRIPTION.on(mainCallReportTarget, "Error in contract description")
} }
diagnostics.forEach { bindingTrace.report(it) } override fun flushDiagnostics() {
reportedErrors.forEach { bindingTrace.report(it) }
} }
override fun hasErrors(): Boolean = diagnostics.isNotEmpty() override fun hasErrors(): Boolean = reportedErrors.isNotEmpty()
} }
@@ -48,16 +48,13 @@ class ContractParsingServices(val languageVersionSettings: LanguageVersionSettin
// is a *necessary* (but not sufficient, actually) condition for presence of 'LazyContractProvider' // is a *necessary* (but not sufficient, actually) condition for presence of 'LazyContractProvider'
if (!expression.isContractDescriptionCallPsiCheck()) return if (!expression.isContractDescriptionCallPsiCheck()) return
val callContext = ContractCallContext(expression, isFirstStatement, scope, trace.bindingContext) val callContext = ContractCallContext(expression, isFirstStatement, scope, trace)
val contractProviderIfAny = (scope.ownerDescriptor as? FunctionDescriptor)?.getUserData(ContractProviderKey) val contractProviderIfAny = (scope.ownerDescriptor as? FunctionDescriptor)?.getUserData(ContractProviderKey)
var resultingContractDescription: ContractDescription? = null var resultingContractDescription: ContractDescription? = null
try { try {
if (!callContext.isContractDescriptionCallPreciseCheck()) return if (!callContext.isContractDescriptionCallPreciseCheck()) return
resultingContractDescription = parseContractAndReportErrors(callContext)
val collector = TraceBasedCollector(trace, expression)
resultingContractDescription = doCheckContract(collector, callContext)
collector.flushDiagnostics(parsingFailed = resultingContractDescription == null)
} finally { } finally {
contractProviderIfAny?.setContractDescription(resultingContractDescription) contractProviderIfAny?.setContractDescription(resultingContractDescription)
} }
@@ -66,14 +63,32 @@ class ContractParsingServices(val languageVersionSettings: LanguageVersionSettin
private fun ContractCallContext.isContractDescriptionCallPreciseCheck(): Boolean = private fun ContractCallContext.isContractDescriptionCallPreciseCheck(): Boolean =
contractCallExpression.isContractDescriptionCallPreciseCheck(bindingContext) contractCallExpression.isContractDescriptionCallPreciseCheck(bindingContext)
private fun doCheckContract(collector: ContractParsingDiagnosticsCollector, callContext: ContractCallContext): ContractDescription? { /**
checkFeatureEnabled(collector) * This function deals with some call that is guaranteed to resolve to 'contract' from stdlib, so,
checkContractAllowedHere(collector, callContext) * ideally, it should satisfy following condition: null returned <=> at least one error was reported
*/
private fun parseContractAndReportErrors(callContext: ContractCallContext): ContractDescription? {
val collector = TraceBasedCollector(callContext.trace, callContext.contractCallExpression)
return if (!collector.hasErrors()) try {
PsiContractParserDispatcher(collector, callContext).parseContract() checkFeatureEnabled(collector)
else checkContractAllowedHere(collector, callContext)
null
// Small optimization: do not even try to parse contract if we already have errors
if (collector.hasErrors()) return null
val parsedContract = PsiContractParserDispatcher(collector, callContext).parseContract()
// Make sure that at least generic error will be reported if we couldn't parse contract
// (null returned => at least one error was reported)
if (parsedContract == null) collector.addFallbackErrorIfNecessary()
// Make sure that we don't return non-null value if there were some errors
// (null returned <= at least one error was reported)
return parsedContract?.takeUnless { collector.hasErrors() }
} finally {
collector.flushDiagnostics()
}
} }
private fun checkFeatureEnabled(collector: ContractParsingDiagnosticsCollector) { private fun checkFeatureEnabled(collector: ContractParsingDiagnosticsCollector) {
@@ -117,8 +132,9 @@ class ContractCallContext(
val contractCallExpression: KtExpression, val contractCallExpression: KtExpression,
val isFirstStatement: Boolean, val isFirstStatement: Boolean,
val scope: LexicalScope, val scope: LexicalScope,
val bindingContext: BindingContext val trace: BindingTrace
) { ) {
val ownerDescriptor: DeclarationDescriptor = scope.ownerDescriptor val ownerDescriptor: DeclarationDescriptor = scope.ownerDescriptor
val functionDescriptor: FunctionDescriptor = ownerDescriptor as FunctionDescriptor val functionDescriptor: FunctionDescriptor = ownerDescriptor as FunctionDescriptor
val bindingContext: BindingContext = trace.bindingContext
} }