Improve type inference for default initializers and add support for assign operator methods

This commit is contained in:
Sergey Chernov 2026-07-28 17:20:02 +04:00
parent 2f118a1fef
commit e66f646b25
3 changed files with 95 additions and 1 deletions

View File

@ -2895,10 +2895,34 @@ class Compiler(
private suspend fun parseExpression(): Statement? { private suspend fun parseExpression(): Statement? {
val pos = cc.currentPos() val pos = cc.currentPos()
return parseExpressionLevel()?.let { ref -> return parseExpressionLevel()?.let { ref ->
ExpressionStatement(ref, pos) ExpressionStatement(lowerTypedAssignOperator(ref, pos), pos)
} }
} }
private fun lowerTypedAssignOperator(ref: ObjRef, pos: Pos): ObjRef {
val assign = ref as? AssignOpRef ?: return ref
val target = assign.target as? FieldRef ?: return ref
val methodName = when (assign.op) {
BinOp.PLUS -> "plusAssign"
BinOp.MINUS -> "minusAssign"
BinOp.STAR -> "mulAssign"
BinOp.SLASH -> "divAssign"
BinOp.PERCENT -> "modAssign"
else -> return ref
}
val targetClass = resolveReceiverTypeDecl(target)?.let { resolveTypeDeclObjClass(it) } ?: return ref
if (targetClass.getInstanceMemberOrNull(methodName, includeAbstract = true, includeStatic = false) == null) {
return ref
}
return MethodCallRef(
target,
methodName,
listOf(ParsedArgument(ExpressionStatement(assign.value, pos), pos)),
tailBlock = false,
isOptional = target.isOptional,
)
}
private suspend fun parseExpressionLevel(level: Int = 0): ObjRef? { private suspend fun parseExpressionLevel(level: Int = 0): ObjRef? {
if (level == lastLevel) if (level == lastLevel)
return parseTerm() return parseTerm()
@ -4299,6 +4323,9 @@ class Compiler(
cc.ifNextIs(Token.Type.ASSIGN) { assignment -> cc.ifNextIs(Token.Type.ASSIGN) { assignment ->
val expr = parseExpression() val expr = parseExpression()
?: throw ScriptError(cc.current().pos, "Expected default value expression") ?: throw ScriptError(cc.current().pos, "Expected default value expression")
if (typeInfo == TypeDecl.TypeAny) {
inferTypeDeclFromInitializer(expr)?.let { typeInfo = it }
}
defaultValue = wrapBytecode(expr) defaultValue = wrapBytecode(expr)
defaultSource = extractDefaultArgumentSource(assignment.pos) defaultSource = extractDefaultArgumentSource(assignment.pos)
} }

View File

@ -768,6 +768,33 @@ class TypesTest {
assertTrue(e.message?.contains("extern variable value cannot have an initializer or delegate") == true) assertTrue(e.message?.contains("extern variable value cannot have an initializer or delegate") == true)
} }
@Test
fun testInferenceFromDefaults() = runTest {
eval("""
fun foo(i = 42, s = "bar",r = 0.0) {
assert( i is Int )
assert( s is String )
assert( r is Real )
r.toInt()
}
class Foobar(val i = 100,val s = "42", r = 1.0) {
fun test(amount: Int) {
assert( i is Int )
assert( s is String )
assert( r is Real )
(r * amount).toInt()
}
}
foo()
val fb = Foobar()
assert(fb.i is Int)
assert(fb.s is String)
assert(fb.r is Real)
val n = fb.test(15)
assertEquals(n, 15)
""".trimIndent())
}
// @Test fun nonTrivialOperatorsTest() = runTest { // @Test fun nonTrivialOperatorsTest() = runTest {
// val s = Script.newScope() // val s = Script.newScope()
// s.eval(""" // s.eval("""

View File

@ -256,6 +256,46 @@ class OperatorOverloadingTest {
""".trimIndent()) """.trimIndent())
} }
@Test
fun testAssignOperatorMethodsOnValMember() = runTest {
eval("""
class Counter(var n: Int) {
fun plusAssign(x: Int) { n = n + x }
fun minusAssign(x: Int) { n = n - x }
fun mulAssign(x: Int) { n = n * x }
fun divAssign(x: Int) { n = n / x }
fun modAssign(x: Int) { n = n % x }
}
class Holder(val counter: Counter)
val holder = Holder(Counter(10))
holder.counter += 2
assertEquals(12, holder.counter.n)
holder.counter -= 3
assertEquals(9, holder.counter.n)
holder.counter *= 4
assertEquals(36, holder.counter.n)
holder.counter /= 6
assertEquals(6, holder.counter.n)
holder.counter %= 4
assertEquals(2, holder.counter.n)
""".trimIndent())
}
@Test
fun testAssignOperatorFallbackOnMutableMember() = runTest {
eval("""
class Counter(val n: Int) {
fun minus(x: Int) = Counter(n - x)
}
class Holder(var counter: Counter)
val holder = Holder(Counter(10))
holder.counter -= 3
assertEquals(7, holder.counter.n)
""".trimIndent())
}
@Test @Test
fun testBuiltinListPlusAssignOnVal() = runTest { fun testBuiltinListPlusAssignOnVal() = runTest {
eval(""" eval("""