Infer final expression returns in lambdas

This commit is contained in:
pavel 2026-08-27 23:31:46 +02:00
commit 72606e7818
6 changed files with 300 additions and 7 deletions

View file

@ -613,7 +613,7 @@ func (g *goGenerator) stmt(stmt Stmt, tail []Stmt) error {
} else {
value, err := g.expr(s.Value, g.currentFunc.ReturnType)
if err != nil {
return err
return fmt.Errorf("returning %s: %w", g.currentFunc.ReturnType, err)
}
value = g.passthroughValue(s.Value, value)
g.line(fmt.Sprintf("return %s", value))
@ -1210,12 +1210,15 @@ func (g *goGenerator) expr(expr Expr, expectedType string) (string, error) {
}
args := make([]string, 0, len(e.Args))
argTypes := g.callArgTypes(e.Callee, len(e.Args))
argTypes := g.callArgTypes(e)
for i, arg := range e.Args {
argType := ""
if i < len(argTypes) {
argType = argTypes[i]
}
if _, isNull := arg.(NullExpr); isNull && g.goCallArgumentIsNilable(e, i) {
argType = ""
}
value, err := g.expr(arg, argType)
if err != nil {
return "", err
@ -2162,7 +2165,23 @@ func (g *goGenerator) lambda(lambda LambdaExpr, expectedType string) (string, er
for _, param := range params {
sub.defineType(param.Name, param.Type)
}
if err := sub.block(lambda.Body); err != nil {
if returnType != "" && returnType != "Unit" && !lambdaHasValueReturn(lambda.Body) {
if len(lambda.Body) == 0 {
return "", fmt.Errorf("value lambda requires a final expression")
}
last, ok := lambda.Body[len(lambda.Body)-1].(ExprStmt)
if !ok {
return "", fmt.Errorf("value lambda requires a final expression or explicit return")
}
if err := sub.block(lambda.Body[:len(lambda.Body)-1]); err != nil {
return "", err
}
value, err := sub.expr(last.Value, returnType)
if err != nil {
return "", err
}
sub.line("return " + value)
} else if err := sub.block(lambda.Body); err != nil {
return "", err
}
b.Write(sub.buf.Bytes())
@ -2170,7 +2189,50 @@ func (g *goGenerator) lambda(lambda LambdaExpr, expectedType string) (string, er
return b.String(), nil
}
func (g *goGenerator) callArgTypes(callee Expr, argCount int) []string {
func (g *goGenerator) callArgTypes(call CallExpr) []string {
argCount := len(call.Args)
if resolved := exprMeta(call); resolved != nil {
switch node := resolved.Node.(type) {
case HIRGoCall:
if len(node.Params) > 0 {
params := node.Params
if node.InjectContext {
params = params[1:]
}
return semanticParamStrings(params, argCount, node.Variadic)
}
case HIRGotlinCall:
if node.Target != nil {
if function, ok := node.Target.Type.(FunctionType); ok {
bindings := map[string]Type{}
for index, argument := range call.TypeArgs {
if index < len(function.TypeParams) {
typ, _ := g.semantic.ResolveType(argument)
bindings[function.TypeParams[index]] = typ
}
}
if selector, ok := call.Callee.(SelectorExpr); ok {
if _, receiverBindings := classInstance(ResolvedType(selector.Receiver)); receiverBindings != nil {
for name, typ := range receiverBindings {
bindings[name] = typ
}
}
}
for index, argument := range call.Args {
if index < len(function.Params) {
inferTypeBindings(function.Params[index], ResolvedType(argument), bindings)
}
}
params := make([]Type, len(function.Params))
for index, param := range function.Params {
params[index] = substituteType(param, bindings)
}
return semanticParamStrings(params, argCount, function.Variadic)
}
}
}
}
callee := call.Callee
ident, ok := callee.(IdentExpr)
if ok {
if ident.Name == "runCatching" {
@ -2192,7 +2254,6 @@ func (g *goGenerator) callArgTypes(callee Expr, argCount int) []string {
}
return argTypes
}
return nil
}
selector, ok := selectorPath(callee)
@ -2209,6 +2270,63 @@ func (g *goGenerator) callArgTypes(callee Expr, argCount int) []string {
}
}
func semanticParamStrings(params []Type, argCount int, variadic bool) []string {
result := make([]string, 0, argCount)
for index := 0; index < argCount; index++ {
paramIndex := index
if paramIndex >= len(params) {
if !variadic || len(params) == 0 {
break
}
paramIndex = len(params) - 1
}
param := params[paramIndex]
if variadic && paramIndex == len(params)-1 {
if list, ok := param.(GenericType); ok && (list.Base.String() == "List" || list.Base.String() == "MutableList") && len(list.Args) == 1 {
param = list.Args[0]
}
}
result = append(result, param.String())
}
return result
}
func (g *goGenerator) goCallArgumentIsNilable(call CallExpr, index int) bool {
resolved := exprMeta(call)
if resolved == nil {
return false
}
goCall, ok := resolved.Node.(HIRGoCall)
if !ok {
return false
}
params := goCall.Params
if goCall.InjectContext && len(params) > 0 {
params = params[1:]
}
if len(params) == 0 {
return false
}
paramIndex := index
if paramIndex >= len(params) {
if !goCall.Variadic {
return false
}
paramIndex = len(params) - 1
}
param := params[paramIndex]
if goCall.Variadic && paramIndex == len(params)-1 {
if list, ok := param.(GenericType); ok && len(list.Args) == 1 {
param = list.Args[0]
}
}
switch param.(type) {
case GoInterfaceType, GoPointerType, NullableType:
return true
}
return false
}
func nameUsedInStmts(name string, stmts []Stmt) bool {
return nameUsedInStmtsWithShadow(name, stmts, false)
}