From e66f646b2594cca238474c49a88ca1b35dc0c513 Mon Sep 17 00:00:00 2001 From: sergeych Date: Tue, 28 Jul 2026 17:20:02 +0400 Subject: [PATCH] Improve type inference for default initializers and add support for assign operator methods --- .../kotlin/net/sergeych/lyng/Compiler.kt | 29 +++++++++++++- lynglib/src/commonTest/kotlin/TypesTest.kt | 27 +++++++++++++ .../sergeych/lyng/OperatorOverloadingTest.kt | 40 +++++++++++++++++++ 3 files changed, 95 insertions(+), 1 deletion(-) diff --git a/lynglib/src/commonMain/kotlin/net/sergeych/lyng/Compiler.kt b/lynglib/src/commonMain/kotlin/net/sergeych/lyng/Compiler.kt index 8ff4dbd..995670e 100644 --- a/lynglib/src/commonMain/kotlin/net/sergeych/lyng/Compiler.kt +++ b/lynglib/src/commonMain/kotlin/net/sergeych/lyng/Compiler.kt @@ -2895,10 +2895,34 @@ class Compiler( private suspend fun parseExpression(): Statement? { val pos = cc.currentPos() 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? { if (level == lastLevel) return parseTerm() @@ -4299,6 +4323,9 @@ class Compiler( cc.ifNextIs(Token.Type.ASSIGN) { assignment -> val expr = parseExpression() ?: throw ScriptError(cc.current().pos, "Expected default value expression") + if (typeInfo == TypeDecl.TypeAny) { + inferTypeDeclFromInitializer(expr)?.let { typeInfo = it } + } defaultValue = wrapBytecode(expr) defaultSource = extractDefaultArgumentSource(assignment.pos) } diff --git a/lynglib/src/commonTest/kotlin/TypesTest.kt b/lynglib/src/commonTest/kotlin/TypesTest.kt index 08d50e7..c528e3e 100644 --- a/lynglib/src/commonTest/kotlin/TypesTest.kt +++ b/lynglib/src/commonTest/kotlin/TypesTest.kt @@ -768,6 +768,33 @@ class TypesTest { 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 { // val s = Script.newScope() // s.eval(""" diff --git a/lynglib/src/commonTest/kotlin/net/sergeych/lyng/OperatorOverloadingTest.kt b/lynglib/src/commonTest/kotlin/net/sergeych/lyng/OperatorOverloadingTest.kt index f02a2f8..4a050ac 100644 --- a/lynglib/src/commonTest/kotlin/net/sergeych/lyng/OperatorOverloadingTest.kt +++ b/lynglib/src/commonTest/kotlin/net/sergeych/lyng/OperatorOverloadingTest.kt @@ -256,6 +256,46 @@ class OperatorOverloadingTest { """.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 fun testBuiltinListPlusAssignOnVal() = runTest { eval("""