Infer final expression returns in lambdas
This commit is contained in:
parent
773c34f3f4
commit
72606e7818
6 changed files with 300 additions and 7 deletions
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue