Merge pull request #39 from bnorm/generic-diagramming

Generic Parameter Diagramming
This commit is contained in:
Brian Norman
2021-04-04 18:54:31 -05:00
committed by GitHub
12 changed files with 515 additions and 259 deletions
@@ -16,157 +16,170 @@
package com.bnorm.power package com.bnorm.power
import com.bnorm.power.diagram.buildDiagramNesting
import com.bnorm.power.diagram.buildTree
import com.bnorm.power.diagram.info
import com.bnorm.power.diagram.irDiagramString
import com.bnorm.power.diagram.substring
import com.bnorm.power.internal.ReturnableBlockTransformer import com.bnorm.power.internal.ReturnableBlockTransformer
import java.io.File import org.jetbrains.kotlin.backend.common.IrElementTransformerVoidWithContext
import org.jetbrains.kotlin.backend.common.*
import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext
import org.jetbrains.kotlin.backend.common.ir.asSimpleLambda import org.jetbrains.kotlin.backend.common.ir.asSimpleLambda
import org.jetbrains.kotlin.backend.common.ir.inline import org.jetbrains.kotlin.backend.common.ir.inline
import org.jetbrains.kotlin.backend.common.lower.DeclarationIrBuilder import org.jetbrains.kotlin.backend.common.lower.DeclarationIrBuilder
import org.jetbrains.kotlin.backend.common.lower.at import org.jetbrains.kotlin.cli.common.messages.CompilerMessageLocation
import org.jetbrains.kotlin.cli.common.messages.* import org.jetbrains.kotlin.cli.common.messages.CompilerMessageSeverity
import org.jetbrains.kotlin.cli.common.messages.MessageCollector
import org.jetbrains.kotlin.descriptors.DescriptorVisibilities import org.jetbrains.kotlin.descriptors.DescriptorVisibilities
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.backend.js.utils.asString
import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope
import org.jetbrains.kotlin.ir.builders.declarations.buildFun import org.jetbrains.kotlin.ir.builders.declarations.buildFun
import org.jetbrains.kotlin.ir.builders.irBlockBody import org.jetbrains.kotlin.ir.builders.irBlockBody
import org.jetbrains.kotlin.ir.builders.irCall import org.jetbrains.kotlin.ir.builders.irCall
import org.jetbrains.kotlin.ir.builders.irCallOp import org.jetbrains.kotlin.ir.builders.irCallOp
import org.jetbrains.kotlin.ir.builders.irFalse
import org.jetbrains.kotlin.ir.builders.irReturn import org.jetbrains.kotlin.ir.builders.irReturn
import org.jetbrains.kotlin.ir.builders.irString import org.jetbrains.kotlin.ir.builders.irString
import org.jetbrains.kotlin.ir.builders.parent import org.jetbrains.kotlin.ir.builders.parent
import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin
import org.jetbrains.kotlin.ir.declarations.IrFile import org.jetbrains.kotlin.ir.declarations.IrFile
import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.declarations.path import org.jetbrains.kotlin.ir.declarations.path
import org.jetbrains.kotlin.ir.expressions.* import org.jetbrains.kotlin.ir.expressions.IrCall
import org.jetbrains.kotlin.ir.expressions.IrConst
import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin
import org.jetbrains.kotlin.ir.expressions.IrStringConcatenation
import org.jetbrains.kotlin.ir.expressions.impl.IrFunctionExpressionImpl import org.jetbrains.kotlin.ir.expressions.impl.IrFunctionExpressionImpl
import org.jetbrains.kotlin.ir.types.* import org.jetbrains.kotlin.ir.symbols.IrTypeParameterSymbol
import org.jetbrains.kotlin.ir.util.* import org.jetbrains.kotlin.ir.types.IrSimpleType
import org.jetbrains.kotlin.ir.visitors.IrElementVisitorVoid import org.jetbrains.kotlin.ir.types.IrType
import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid import org.jetbrains.kotlin.ir.types.IrTypeArgument
import org.jetbrains.kotlin.ir.visitors.acceptVoid import org.jetbrains.kotlin.ir.types.IrTypeProjection
import org.jetbrains.kotlin.ir.types.classifierOrNull
import org.jetbrains.kotlin.ir.types.getClass
import org.jetbrains.kotlin.ir.types.isBoolean
import org.jetbrains.kotlin.ir.types.isSubtypeOf
import org.jetbrains.kotlin.ir.util.deepCopyWithSymbols
import org.jetbrains.kotlin.ir.util.functions
import org.jetbrains.kotlin.ir.util.isFunctionOrKFunction
import org.jetbrains.kotlin.ir.util.kotlinFqName
import org.jetbrains.kotlin.name.FqName import org.jetbrains.kotlin.name.FqName
import org.jetbrains.kotlin.name.Name import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.util.OperatorNameConventions import org.jetbrains.kotlin.util.OperatorNameConventions
fun FileLoweringPass.runOnFileInOrder(irFile: IrFile) {
irFile.acceptVoid(object : IrElementVisitorVoid {
override fun visitElement(element: IrElement) {
element.acceptChildrenVoid(this)
}
override fun visitFile(declaration: IrFile) {
lower(declaration)
super.visitFile(declaration)
}
})
}
class PowerAssertCallTransformer( class PowerAssertCallTransformer(
private val file: IrFile,
private val fileSource: String,
private val context: IrPluginContext, private val context: IrPluginContext,
private val messageCollector: MessageCollector, private val messageCollector: MessageCollector,
private val functions: Set<FqName> private val functions: Set<FqName>
) : IrElementTransformerVoidWithContext(), FileLoweringPass { ) : IrElementTransformerVoidWithContext() {
private lateinit var file: IrFile
private lateinit var fileSource: String
override fun lower(irFile: IrFile) {
file = irFile
fileSource = File(irFile.path).readText()
.replace("\r\n", "\n") // https://youtrack.jetbrains.com/issue/KT-41888
irFile.transformChildrenVoid()
}
override fun visitCall(expression: IrCall): IrExpression { override fun visitCall(expression: IrCall): IrExpression {
val fqName = expression.symbol.owner.kotlinFqName val function = expression.symbol.owner
if (functions.none { fqName == it }) val fqName = function.kotlinFqName
if (function.valueParameters.isEmpty() || functions.none { fqName == it })
return super.visitCall(expression) return super.visitCall(expression)
// Find a valid delegate function or do not translate // Find a valid delegate function or do not translate
val delegate = findDelegate(fqName) ?: run { val delegate = findDelegate(function) ?: run {
val valueType = function.valueParameters[0].type.asString()
messageCollector.warn( messageCollector.warn(
expression, expression,
"Unable to find overload for function $fqName callable as $fqName(Boolean, String) or $fqName(Boolean, () -> String) for power-assert transformation" "Unable to find overload for function $fqName callable as $fqName($valueType, String) or $fqName($valueType, () -> String) for power-assert transformation"
) )
return super.visitCall(expression) return super.visitCall(expression)
} }
val function = expression.symbol.owner // TODO - support more arguments by currying?
val assertionArgument = expression.getValueArgument(0)!! val assertionArgument = expression.getValueArgument(0)!!
val messageArgument = if (function.valueParameters.size == 2) expression.getValueArgument(1) else null val messageArgument = if (function.valueParameters.size == 2) expression.getValueArgument(1) else null
// If the tree does not contain any children, the expression is not transformable // If the tree does not contain any children, the expression is not transformable
val tree = buildAssertTree(assertionArgument) val root = buildTree(assertionArgument) ?: run {
val root = tree.children.singleOrNull() ?: run {
messageCollector.info(expression, "Expression is constant and will not be power-assert transformed") messageCollector.info(expression, "Expression is constant and will not be power-assert transformed")
return super.visitCall(expression) return super.visitCall(expression)
} }
// println(root.dump())
val symbol = currentScope!!.scope.scopeOwnerSymbol val symbol = currentScope!!.scope.scopeOwnerSymbol
DeclarationIrBuilder(context, symbol).run { val builder = DeclarationIrBuilder(context, symbol, expression.startOffset, expression.endOffset)
at(expression) return builder.buildDiagramNesting(root) { argument, variables ->
val lambda = messageArgument?.asSimpleLambda()
val generator = object : PowerAssertGenerator() { val title = when {
override fun IrBuilderWithScope.buildAssertThrow(subStack: List<IrStackVariable>): IrExpression { messageArgument is IrConst<*> -> messageArgument
messageArgument is IrStringConcatenation -> messageArgument
val lambda = messageArgument?.asSimpleLambda() lambda != null -> lambda.deepCopyWithSymbols(parent).inline(parent)
val title = when { .transform(ReturnableBlockTransformer(context, symbol), null)
messageArgument is IrConst<*> -> messageArgument messageArgument != null -> {
messageArgument is IrStringConcatenation -> messageArgument val invoke =
lambda != null -> lambda.deepCopyWithSymbols(parent).inline(parent).transform(ReturnableBlockTransformer(context, symbol), null) messageArgument.type.getClass()!!.functions.single { it.name == OperatorNameConventions.INVOKE }
messageArgument != null -> { irCallOp(invoke.symbol, invoke.returnType, messageArgument)
val invoke = messageArgument.type.getClass()!!.functions.single { it.name == OperatorNameConventions.INVOKE }
irCallOp(invoke.symbol, invoke.returnType, messageArgument)
}
// TODO what should the default message be?
else -> irString("Assertion failed")
}
return delegate.buildCall(this, expression, buildMessage(file, fileSource, title.deepCopyWithSymbols(parent), expression, subStack))
} }
// TODO what should the default message be?
assertionArgument.type.isBoolean() -> irString("Assertion failed")
else -> null
} }
// println(expression.dump()) val prefix = title?.deepCopyWithSymbols(parent)
// println(tree.dump()) val diagram = irDiagramString(file, fileSource, prefix, expression, variables)
delegate.buildCall(this, expression, argument, diagram)
return generator.buildAssert(this, root)
// .also { println(it.dump()) }
} }
// .also { println(expression.dump()) }
// .also { println(it.dump()) }
// .also { println(expression.dumpKotlinLike()) }
// .also { println(it.dumpKotlinLike()) }
} }
private interface FunctionDelegate { private interface FunctionDelegate {
fun buildCall(builder: IrBuilderWithScope, original: IrCall, message: IrExpression): IrExpression fun buildCall(
builder: IrBuilderWithScope,
original: IrCall,
argument: IrExpression,
message: IrExpression
): IrExpression
} }
private fun findDelegate(fqName: FqName): FunctionDelegate? { private fun findDelegate(function: IrFunction): FunctionDelegate? {
return context.referenceFunctions(fqName) if (function.valueParameters.isEmpty()) return null
return context.referenceFunctions(function.kotlinFqName)
.mapNotNull { overload -> .mapNotNull { overload ->
// TODO allow other signatures than (Boolean, String) and (Boolean, () -> String) // TODO allow other signatures than (Boolean, String) and (Boolean, () -> String)
val parameters = overload.owner.valueParameters val parameters = overload.owner.valueParameters
if (parameters.size != 2) return@mapNotNull null if (parameters.size != 2) return@mapNotNull null
if (!parameters[0].type.isBoolean()) return@mapNotNull null if (!function.valueParameters[0].type.isAssignableTo(parameters[0].type)) return@mapNotNull null
val messageParameter = parameters.last()
return@mapNotNull when { return@mapNotNull when {
isStringSupertype(parameters[1].type) -> { isStringSupertype(messageParameter.type) -> {
object : FunctionDelegate { object : FunctionDelegate {
override fun buildCall(builder: IrBuilderWithScope, original: IrCall, message: IrExpression): IrExpression = with(builder) { override fun buildCall(
irCall(overload, type = overload.owner.returnType).apply { builder: IrBuilderWithScope,
original: IrCall,
argument: IrExpression,
message: IrExpression
): IrExpression = with(builder) {
irCall(overload, type = original.type).apply {
dispatchReceiver = original.dispatchReceiver?.deepCopyWithSymbols(parent) dispatchReceiver = original.dispatchReceiver?.deepCopyWithSymbols(parent)
extensionReceiver = original.extensionReceiver?.deepCopyWithSymbols(parent) extensionReceiver = original.extensionReceiver?.deepCopyWithSymbols(parent)
for (i in 0 until original.typeArgumentsCount) { for (i in 0 until original.typeArgumentsCount) {
putTypeArgument(i, original.getTypeArgument(i)) putTypeArgument(i, original.getTypeArgument(i))
} }
putValueArgument(0, irFalse()) putValueArgument(0, argument)
putValueArgument(1, message) putValueArgument(1, message)
} }
} }
} }
} }
isStringFunction(parameters[1].type) -> { isStringFunction(messageParameter.type) -> {
object : FunctionDelegate { object : FunctionDelegate {
override fun buildCall(builder: IrBuilderWithScope, original: IrCall, message: IrExpression): IrExpression = with(builder) { override fun buildCall(
builder: IrBuilderWithScope,
original: IrCall,
argument: IrExpression,
message: IrExpression
): IrExpression = with(builder) {
val scope = this val scope = this
val lambda = builder.context.irFactory.buildFun { val lambda = builder.context.irFactory.buildFun {
name = Name.special("<anonymous>") name = Name.special("<anonymous>")
@@ -180,14 +193,20 @@ class PowerAssertCallTransformer(
} }
parent = scope.parent parent = scope.parent
} }
val expression = IrFunctionExpressionImpl(original.startOffset, original.endOffset, parameters[1].type, lambda, IrStatementOrigin.LAMBDA) val expression = IrFunctionExpressionImpl(
irCall(overload, type = overload.owner.returnType).apply { original.startOffset,
original.endOffset,
messageParameter.type,
lambda,
IrStatementOrigin.LAMBDA
)
irCall(overload, type = original.type).apply {
dispatchReceiver = original.dispatchReceiver?.deepCopyWithSymbols(parent) dispatchReceiver = original.dispatchReceiver?.deepCopyWithSymbols(parent)
extensionReceiver = original.extensionReceiver?.deepCopyWithSymbols(parent) extensionReceiver = original.extensionReceiver?.deepCopyWithSymbols(parent)
for (i in 0 until original.typeArgumentsCount) { for (i in 0 until original.typeArgumentsCount) {
putTypeArgument(i, original.getTypeArgument(i)) putTypeArgument(i, original.getTypeArgument(i))
} }
putValueArgument(0, irFalse()) putValueArgument(0, argument)
putValueArgument(1, expression) putValueArgument(1, expression)
} }
} }
@@ -210,6 +229,12 @@ class PowerAssertCallTransformer(
private fun isStringSupertype(type: IrType): Boolean = private fun isStringSupertype(type: IrType): Boolean =
context.irBuiltIns.stringType.isSubtypeOf(type, context.irBuiltIns) context.irBuiltIns.stringType.isSubtypeOf(type, context.irBuiltIns)
private fun IrType.isAssignableTo(type: IrType): Boolean {
if (isSubtypeOf(type, context.irBuiltIns)) return true
val superTypes = (type.classifierOrNull as? IrTypeParameterSymbol)?.owner?.superTypes
return superTypes != null && superTypes.all { isSubtypeOf(it, context.irBuiltIns) }
}
private fun MessageCollector.info(expression: IrElement, message: String) { private fun MessageCollector.info(expression: IrElement, message: String) {
report(expression, CompilerMessageSeverity.INFO, message) report(expression, CompilerMessageSeverity.INFO, message)
} }
@@ -1,104 +0,0 @@
/*
* Copyright (C) 2020 Brian Norman
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.bnorm.power
import org.jetbrains.kotlin.backend.common.lower.irIfThen
import org.jetbrains.kotlin.backend.common.lower.irNot
import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope
import org.jetbrains.kotlin.ir.builders.IrStatementsBuilder
import org.jetbrains.kotlin.ir.builders.irBlock
import org.jetbrains.kotlin.ir.builders.irGet
import org.jetbrains.kotlin.ir.builders.irTemporary
import org.jetbrains.kotlin.ir.builders.parent
import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.IrWhen
import org.jetbrains.kotlin.ir.util.deepCopyWithSymbols
import org.jetbrains.kotlin.ir.visitors.IrElementTransformerVoid
abstract class PowerAssertGenerator {
abstract fun IrBuilderWithScope.buildAssertThrow(subStack: List<IrStackVariable>): IrExpression
fun buildAssert(
builder: IrBuilderWithScope,
root: Node
): IrExpression {
return builder.irBlock {
buildAssert(root, mutableListOf()) { subStack ->
buildAssertThrow(subStack)
}
}
}
private fun IrStatementsBuilder<*>.buildAssert(
node: Node,
stack: MutableList<IrStackVariable>,
thenPart: IrStatementsBuilder<*>.(stack: MutableList<IrStackVariable>) -> IrExpression
) {
fun IrStatementsBuilder<*>.nest(children: List<Node>, index: Int, stack: MutableList<IrStackVariable>) {
val child = children[index]
buildAssert(child, stack) { subStack ->
if (index + 1 == children.size) buildAssertThrow(subStack)
else irBlock { nest(children, index + 1, subStack) }
}
}
when (node) {
is ExpressionNode -> {
+irIfNotThan(stack, node, thenPart)
}
is AndNode -> {
for (child in node.children) {
buildAssert(child, stack, thenPart)
}
}
is OrNode -> {
nest(node.children, 0, stack)
}
}
}
private inline fun IrStatementsBuilder<*>.irIfNotThan(
stack: MutableList<IrStackVariable>,
node: ExpressionNode,
thenPart: IrStatementsBuilder<*>.(subStack: MutableList<IrStackVariable>) -> IrExpression
): IrWhen {
val expressions = node.getExpressionsCopy(this.parent)
val stackTransformer = StackBuilder(this, stack, expressions)
val transformed = expressions.first().transform(stackTransformer, null)
return irIfThen(irNot(transformed), thenPart(stack.toMutableList()))
}
class StackBuilder(
private val builder: IrStatementsBuilder<*>,
private val stack: MutableList<IrStackVariable>,
private val transform: List<IrExpression>
) : IrElementTransformerVoid() {
override fun visitExpression(expression: IrExpression): IrExpression {
return if (expression in transform) {
with(builder) {
val copy = expression.deepCopyWithSymbols(scope.getLocalDeclarationParent())
val variable = irTemporary(super.visitExpression(expression))
stack.add(IrStackVariable(variable, copy))
irGet(variable)
}
} else {
super.visitExpression(expression)
}
}
}
}
@@ -20,7 +20,9 @@ import org.jetbrains.kotlin.backend.common.extensions.IrGenerationExtension
import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext
import org.jetbrains.kotlin.cli.common.messages.MessageCollector import org.jetbrains.kotlin.cli.common.messages.MessageCollector
import org.jetbrains.kotlin.ir.declarations.IrModuleFragment import org.jetbrains.kotlin.ir.declarations.IrModuleFragment
import org.jetbrains.kotlin.ir.declarations.path
import org.jetbrains.kotlin.name.FqName import org.jetbrains.kotlin.name.FqName
import java.io.File
class PowerAssertIrGenerationExtension( class PowerAssertIrGenerationExtension(
private val messageCollector: MessageCollector, private val messageCollector: MessageCollector,
@@ -28,7 +30,11 @@ class PowerAssertIrGenerationExtension(
) : IrGenerationExtension { ) : IrGenerationExtension {
override fun generate(moduleFragment: IrModuleFragment, pluginContext: IrPluginContext) { override fun generate(moduleFragment: IrModuleFragment, pluginContext: IrPluginContext) {
for (file in moduleFragment.files) { for (file in moduleFragment.files) {
PowerAssertCallTransformer(pluginContext, messageCollector, functions).runOnFileInOrder(file) val fileSource = File(file.path).readText()
.replace("\r\n", "\n") // https://youtrack.jetbrains.com/issue/KT-41888
PowerAssertCallTransformer(file, fileSource, pluginContext, messageCollector, functions)
.visitFile(file)
} }
} }
} }
@@ -0,0 +1,142 @@
/*
* Copyright (C) 2020 Brian Norman
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.bnorm.power.diagram
import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope
import org.jetbrains.kotlin.ir.builders.irBlock
import org.jetbrains.kotlin.ir.builders.irFalse
import org.jetbrains.kotlin.ir.builders.irIfThenElse
import org.jetbrains.kotlin.ir.builders.irTrue
import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.util.deepCopyWithSymbols
fun IrBuilderWithScope.buildDiagramNesting(
root: Node,
call: IrBuilderWithScope.(IrExpression, List<IrTemporaryVariable>) -> IrExpression
): IrExpression {
return buildExpression(root, listOf()) { argument, subStack ->
call(argument, subStack)
}
}
private fun IrBuilderWithScope.buildExpression(
node: Node,
variables: List<IrTemporaryVariable>,
call: IrBuilderWithScope.(IrExpression, List<IrTemporaryVariable>) -> IrExpression
): IrExpression = when (node) {
is ExpressionNode -> add(node, variables, call)
is AndNode -> nest(node, 0, variables, call)
is OrNode -> nest(node, 0, variables, call)
else -> TODO("Unknown node type=$node")
}
/**
* ```
* val result = call(1 + 2 + 3)
* ```
* Transforms to
* ```
* val result = run {
* val tmp0 = 1 + 2
* val tmp1 = tmp0 + 3
* call(tmp1, <diagram>)
* }
* ```
*/
private fun IrBuilderWithScope.add(
node: ExpressionNode,
variables: List<IrTemporaryVariable>,
call: IrBuilderWithScope.(IrExpression, List<IrTemporaryVariable>) -> IrExpression
): IrExpression {
return irBlock {
val head = node.expressions.first().deepCopyWithSymbols(scope.getLocalDeclarationParent())
val expressions = (buildTree(head) as ExpressionNode).expressions
val transformer = IrTemporaryExtractionTransformer(this@irBlock, expressions.toSet())
val transformed = expressions.first().transform(transformer, null)
+call(transformed, variables + transformer.variables)
}
}
/**
* ```
* val result = call(1 == 1 && 2 == 2)
* ```
* Transforms to
* ```
* val result = run {
* val tmp0 = 1 == 1
* if (tmp0) {
* val tmp1 = 2 == 2
* call(tmp1, <diagram>)
* }
* else call(false, <diagram>)
* }
* ```
*/
private fun IrBuilderWithScope.nest(
node: AndNode,
index: Int,
variables: List<IrTemporaryVariable>,
call: IrBuilderWithScope.(IrExpression, List<IrTemporaryVariable>) -> IrExpression
): IrExpression {
val children = node.children
val child = children[index]
return buildExpression(child, variables) { argument, newVariables ->
if (index + 1 == children.size) call(argument, newVariables) // last expression, result is false
else irIfThenElse(
context.irBuiltIns.anyType,
argument,
nest(node, index + 1, newVariables, call), // more expressions, continue nesting
call(irFalse(), newVariables), // short-circuit result to false
)
}
}
/**
* ```
* val result = call(1 == 1 || 2 == 2)
* ```
* Transforms to
* ```
* val result = run {
* val tmp0 = 1 == 1
* if (tmp0) call(true, <diagram>)
* else {
* val tmp1 = 2 == 2
* call(tmp1, <diagram>)
* }
* }
* ```
*/
private fun IrBuilderWithScope.nest(
node: OrNode,
index: Int,
variables: List<IrTemporaryVariable>,
call: IrBuilderWithScope.(IrExpression, List<IrTemporaryVariable>) -> IrExpression
): IrExpression {
val children = node.children
val child = children[index]
return buildExpression(child, variables) { argument, newVariables ->
if (index + 1 == children.size) call(argument, newVariables) // last expression, result is false
else irIfThenElse(
context.irBuiltIns.anyType,
argument,
call(irTrue(), newVariables), // short-circuit result to true
nest(node, index + 1, newVariables, call), // more expressions, continue nesting
)
}
}
@@ -14,25 +14,31 @@
* limitations under the License. * limitations under the License.
*/ */
package com.bnorm.power package com.bnorm.power.diagram
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.declarations.IrDeclarationParent
import org.jetbrains.kotlin.ir.expressions.IrCall import org.jetbrains.kotlin.ir.expressions.IrCall
import org.jetbrains.kotlin.ir.expressions.IrConst import org.jetbrains.kotlin.ir.expressions.IrConst
import org.jetbrains.kotlin.ir.expressions.IrContainerExpression import org.jetbrains.kotlin.ir.expressions.IrContainerExpression
import org.jetbrains.kotlin.ir.expressions.IrExpression import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin
import org.jetbrains.kotlin.ir.expressions.IrWhen import org.jetbrains.kotlin.ir.expressions.IrWhen
import org.jetbrains.kotlin.ir.util.deepCopyWithSymbols import org.jetbrains.kotlin.ir.util.dumpKotlinLike
import org.jetbrains.kotlin.ir.visitors.IrElementVisitor import org.jetbrains.kotlin.ir.visitors.IrElementVisitor
sealed class Node { abstract class Node {
abstract val parent: Node? private val _children = mutableListOf<Node>()
val mutableChildren: MutableList<Node> = mutableListOf() val children: List<Node> get() = _children
val children: List<Node> get() = mutableChildren
protected fun dump(builder: StringBuilder, indent: Int) { fun addChild(node: Node) {
_children.add(node)
}
fun dump(): String = buildString {
dump(this, 0)
}
private fun dump(builder: StringBuilder, indent: Int) {
builder.append(" ".repeat(indent)).append(this).appendLine() builder.append(" ".repeat(indent)).append(this).appendLine()
for (child in children) { for (child in children) {
child.dump(builder, indent + 1) child.dump(builder, indent + 1)
@@ -40,56 +46,30 @@ sealed class Node {
} }
} }
class AndNode(override val parent: Node) : Node() { class AndNode : Node() {
init {
parent.mutableChildren.add(this)
}
override fun toString() = "AndNode" override fun toString() = "AndNode"
} }
class OrNode(override val parent: Node) : Node() { class OrNode : Node() {
init {
parent.mutableChildren.add(this)
}
override fun toString() = "OrNode" override fun toString() = "OrNode"
} }
class ExpressionNode( class ExpressionNode : Node() {
override val parent: Node private val _expressions = mutableListOf<IrExpression>()
) : Node() { val expressions: List<IrExpression> = _expressions
init {
parent.mutableChildren.add(this)
}
private val _expressions: MutableList<IrExpression> = mutableListOf()
fun add(expression: IrExpression) { fun add(expression: IrExpression) {
_expressions.add(expression) _expressions.add(expression)
} }
fun getExpressionsCopy(initialParent: IrDeclarationParent?): List<IrExpression> { override fun toString() = "ExpressionNode(${_expressions.map { it.dumpKotlinLike() }})"
// Return a copy of all the expression by creating a deep copy of the head
// expression and running back through the assertion tree builder
val headCopy = _expressions.first().deepCopyWithSymbols(initialParent)
return (buildAssertTree(headCopy).children.single() as ExpressionNode)._expressions
}
override fun toString() = "ExpressionNode($_expressions)"
} }
class RootNode : Node() { fun buildTree(expression: IrExpression): Node? {
override val parent: Node? = null class RootNode : Node() {
override fun toString() = "RootNode"
fun dump(): String = buildString {
dump(this, 0)
} }
override fun toString() = "RootNode"
}
fun buildAssertTree(expression: IrExpression): RootNode {
val tree = RootNode() val tree = RootNode()
expression.accept(object : IrElementVisitor<Unit, Node> { expression.accept(object : IrElementVisitor<Unit, Node> {
val INCREMENT_DECREMENT_OPERATORS = setOf( val INCREMENT_DECREMENT_OPERATORS = setOf(
@@ -104,14 +84,14 @@ fun buildAssertTree(expression: IrExpression): RootNode {
} }
override fun visitExpression(expression: IrExpression, data: Node) { override fun visitExpression(expression: IrExpression, data: Node) {
val node = data as? ExpressionNode ?: ExpressionNode(data) val node = data as? ExpressionNode ?: ExpressionNode().also { data.addChild(it) }
node.add(expression) node.add(expression)
expression.acceptChildren(this, node) expression.acceptChildren(this, node)
} }
override fun visitContainerExpression(expression: IrContainerExpression, data: Node) { override fun visitContainerExpression(expression: IrContainerExpression, data: Node) {
if (expression.origin in INCREMENT_DECREMENT_OPERATORS) { if (expression.origin in INCREMENT_DECREMENT_OPERATORS) {
val node = data as? ExpressionNode ?: ExpressionNode(data) val node = data as? ExpressionNode ?: ExpressionNode().also { data.addChild(it) }
node.add(expression) node.add(expression)
return // Skip the internals of increment/decrement operations return // Skip the internals of increment/decrement operations
} }
@@ -136,7 +116,7 @@ fun buildAssertTree(expression: IrExpression): RootNode {
when (expression.origin) { when (expression.origin) {
IrStatementOrigin.ANDAND -> { IrStatementOrigin.ANDAND -> {
// flatten `&&` expressions to be at the same level // flatten `&&` expressions to be at the same level
val node = data as? AndNode ?: AndNode(data) val node = data as? AndNode ?: AndNode().also { data.addChild(it) }
require(expression.branches.size == 2) require(expression.branches.size == 2)
val thenBranch = expression.branches[0] val thenBranch = expression.branches[0]
@@ -157,7 +137,7 @@ fun buildAssertTree(expression: IrExpression): RootNode {
} }
IrStatementOrigin.OROR -> { IrStatementOrigin.OROR -> {
// flatten `||` expressions to be at the same level // flatten `||` expressions to be at the same level
val node = data as? OrNode ?: OrNode(data) val node = data as? OrNode ?: OrNode().also { data.addChild(it) }
require(expression.branches.size == 2) require(expression.branches.size == 2)
val thenBranchCondition = expression.branches[0].condition val thenBranchCondition = expression.branches[0].condition
@@ -182,11 +162,11 @@ fun buildAssertTree(expression: IrExpression): RootNode {
else -> { else -> {
// Add as basic expression and terminate // Add as basic expression and terminate
// TODO this has to be broken and not work in all cases... // TODO this has to be broken and not work in all cases...
ExpressionNode(data) ExpressionNode().also { data.addChild(it) }
} }
} }
} }
}, tree) }, tree)
return tree return tree.children.singleOrNull()
} }
@@ -14,8 +14,9 @@
* limitations under the License. * limitations under the License.
*/ */
package com.bnorm.power package com.bnorm.power.diagram
import com.bnorm.power.irString
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.SourceRangeInfo import org.jetbrains.kotlin.ir.SourceRangeInfo
import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope
@@ -31,38 +32,26 @@ import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.IrMemberAccessExpression import org.jetbrains.kotlin.ir.expressions.IrMemberAccessExpression
import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin
data class IrStackVariable( fun IrBuilderWithScope.irDiagramString(
val temporary: IrVariable,
val original: IrExpression
)
data class ValueDisplay(
val value: IrVariable,
val indent: Int,
val row: Int,
val source: String
)
fun IrBuilderWithScope.buildMessage(
file: IrFile, file: IrFile,
fileSource: String, fileSource: String,
title: IrExpression, prefix: IrExpression? = null,
expression: IrExpression, original: IrExpression,
stack: List<IrStackVariable> variables: List<IrTemporaryVariable>
): IrExpression { ): IrExpression {
val originalInfo = file.info(expression) val originalInfo = file.info(original)
val callIndent = originalInfo.startColumnNumber val callIndent = originalInfo.startColumnNumber
val stackValues = stack.map { it.toValueDisplay(fileSource, callIndent, file, originalInfo) } val stackValues = variables.map { it.toValueDisplay(fileSource, callIndent, file, originalInfo) }
val valuesByRow = stackValues.groupBy { it.row } val valuesByRow = stackValues.groupBy { it.row }
val rows = fileSource.substring(expression) val rows = fileSource.substring(original)
.replace("\n" + " ".repeat(callIndent), "\n") // Remove additional indentation .replace("\n" + " ".repeat(callIndent), "\n") // Remove additional indentation
.split("\n") .split("\n")
return irConcat().apply { return irConcat().apply {
addArgument(title) if (prefix != null) addArgument(prefix)
for ((row, rowSource) in rows.withIndex()) { for ((row, rowSource) in rows.withIndex()) {
val rowValues = valuesByRow[row]?.let { values -> values.sortedBy { it.indent } } ?: emptyList() val rowValues = valuesByRow[row]?.let { values -> values.sortedBy { it.indent } } ?: emptyList()
@@ -70,7 +59,7 @@ fun IrBuilderWithScope.buildMessage(
addArgument( addArgument(
irString { irString {
appendLine() if (row != 0 || prefix != null) appendLine()
append(rowSource) append(rowSource)
if (indentations.isNotEmpty()) { if (indentations.isNotEmpty()) {
appendLine() appendLine()
@@ -102,7 +91,14 @@ fun IrBuilderWithScope.buildMessage(
} }
} }
private fun IrStackVariable.toValueDisplay( private data class ValueDisplay(
val value: IrVariable,
val indent: Int,
val row: Int,
val source: String
)
private fun IrTemporaryVariable.toValueDisplay(
fileSource: String, fileSource: String,
callIndent: Int, callIndent: Int,
file: IrFile, file: IrFile,
@@ -0,0 +1,33 @@
package com.bnorm.power.diagram
import org.jetbrains.kotlin.ir.builders.IrStatementsBuilder
import org.jetbrains.kotlin.ir.builders.irGet
import org.jetbrains.kotlin.ir.builders.irTemporary
import org.jetbrains.kotlin.ir.declarations.IrVariable
import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.util.deepCopyWithSymbols
import org.jetbrains.kotlin.ir.visitors.IrElementTransformerVoid
data class IrTemporaryVariable(
val temporary: IrVariable,
val original: IrExpression
)
class IrTemporaryExtractionTransformer(
private val builder: IrStatementsBuilder<*>,
private val transform: Set<IrExpression>
) : IrElementTransformerVoid() {
private val _variables = mutableListOf<IrTemporaryVariable>()
val variables: List<IrTemporaryVariable> = _variables
override fun visitExpression(expression: IrExpression): IrExpression {
return if (expression in transform) {
val copy = expression.deepCopyWithSymbols(builder.scope.getLocalDeclarationParent())
val variable = builder.irTemporary(super.visitExpression(expression))
_variables.add(IrTemporaryVariable(variable, copy))
builder.irGet(variable)
} else {
super.visitExpression(expression)
}
}
}
@@ -0,0 +1,44 @@
package com.bnorm.power
import org.jetbrains.kotlin.name.FqName
import org.junit.Test
class AssertBooleanTest {
@Test
fun `test assertTrue transformation`() {
assertMessage(
"""
import kotlin.test.assertTrue
fun main() {
assertTrue(1 != 1)
}""",
"""
Assertion failed
assertTrue(1 != 1)
|
false
""".trimIndent(),
PowerAssertComponentRegistrar(setOf(FqName("kotlin.test.assertTrue")))
)
}
@Test
fun `test assertFalse transformation`() {
assertMessage(
"""
import kotlin.test.assertFalse
fun main() {
assertFalse(1 == 1)
}""",
"""
Assertion failed
assertFalse(1 == 1)
|
true
""".trimIndent(),
PowerAssertComponentRegistrar(setOf(FqName("kotlin.test.assertFalse")))
)
}
}
@@ -0,0 +1,84 @@
package com.bnorm.power
import com.tschuchort.compiletesting.KotlinCompilation
import com.tschuchort.compiletesting.SourceFile
import org.jetbrains.kotlin.name.FqName
import org.junit.Test
import java.io.ByteArrayOutputStream
import java.io.PrintStream
import java.lang.reflect.InvocationTargetException
import kotlin.test.assertEquals
class DebugFunctionTest {
@Test
fun `debug function transformation`() {
val actual = executeMainDebug(
"""
dbg(1 + 2 + 3)
""".trimIndent()
)
assertEquals(
"""
dbg(1 + 2 + 3)
| |
| 6
3
""".trimIndent(),
actual.trim()
)
}
@Test
fun `debug function transformation with message`() {
val actual = executeMainDebug(
"""
dbg(1 + 2 + 3, "Message:")
""".trimIndent()
)
assertEquals(
"""
Message:
dbg(1 + 2 + 3, "Message:")
| |
| 6
3
""".trimIndent(),
actual.trim()
)
}
}
fun executeMainDebug(mainBody: String): String {
val file = SourceFile.kotlin(
name = "main.kt",
contents = """
fun <T> dbg(value: T): T = value
fun <T> dbg(value: T, msg: String): T {
println(msg)
return value
}
fun main() {
$mainBody
}
""",
trimIndent = false
)
val result = compile(listOf(file), PowerAssertComponentRegistrar(setOf(FqName("dbg"))))
assertEquals(KotlinCompilation.ExitCode.OK, result.exitCode)
val kClazz = result.classLoader.loadClass("MainKt")
val main = kClazz.declaredMethods.single { it.name == "main" && it.parameterCount == 0 }
val prevOut = System.out
try {
val out = ByteArrayOutputStream()
System.setOut(PrintStream(out))
main.invoke(null)
return out.toString("UTF-8")
} catch (t: InvocationTargetException) {
throw t.cause!!
} finally {
System.setOut(prevOut)
}
}
+2 -1
View File
@@ -62,6 +62,7 @@ configure<com.bnorm.power.PowerAssertGradleExtension> {
"kotlin.test.assertTrue", "kotlin.test.assertTrue",
"kotlin.require", "kotlin.require",
"com.bnorm.power.AssertScope.assert", "com.bnorm.power.AssertScope.assert",
"com.bnorm.power.assert" "com.bnorm.power.assert",
"com.bnorm.power.dbg"
) )
} }
@@ -0,0 +1,10 @@
package com.bnorm.power
val debugLog = StringBuilder()
fun <T> dbg(value: T): T = value
fun <T> dbg(value: T, msg: String): T {
debugLog.appendLine(msg)
return value
}
@@ -16,6 +16,7 @@
package com.bnorm.power package com.bnorm.power
import kotlin.test.AfterTest
import kotlin.test.Test import kotlin.test.Test
import kotlin.test.assertEquals import kotlin.test.assertEquals
import kotlin.test.assertFailsWith import kotlin.test.assertFailsWith
@@ -23,6 +24,11 @@ import kotlin.test.assertTrue
class PowerAssertTest { class PowerAssertTest {
@AfterTest
fun cleanup() {
debugLog.clear()
}
@Test @Test
fun assertTrue() { fun assertTrue() {
val error = assertFailsWith<AssertionError> { assertTrue(Person.UNKNOWN.size == 1) } val error = assertFailsWith<AssertionError> { assertTrue(Person.UNKNOWN.size == 1) }
@@ -111,4 +117,37 @@ class PowerAssertTest {
""".trimIndent(), """.trimIndent(),
) )
} }
@Test
fun dbgTest() {
val name = "Jane"
val greeting = dbg("Hello, $name")
assert(greeting == "Hello, Jane")
assertEquals(
actual = debugLog.toString().trim(),
expected = """
dbg("Hello, ${"$"}name")
| |
| Jane
Hello, Jane
""".trimIndent()
)
}
@Test
fun dbgMessageTest() {
val name = "Jane"
val greeting = dbg("Hello, $name", "Greeting:")
assert(greeting == "Hello, Jane")
assertEquals(
actual = debugLog.toString().trim(),
expected = """
Greeting:
dbg("Hello, ${"$"}name", "Greeting:")
| |
| Jane
Hello, Jane
""".trimIndent()
)
}
} }