1
2
3
4
5 package inline
6
7 import (
8 "bytes"
9 "fmt"
10 "go/ast"
11 "go/constant"
12 "go/format"
13 "go/parser"
14 "go/token"
15 "go/types"
16 "maps"
17 pathpkg "path"
18 "reflect"
19 "slices"
20 "strings"
21
22 "golang.org/x/tools/go/ast/astutil"
23 "golang.org/x/tools/go/types/typeutil"
24 internalastutil "golang.org/x/tools/internal/astutil"
25 "golang.org/x/tools/internal/astutil/free"
26 "golang.org/x/tools/internal/packagepath"
27 "golang.org/x/tools/internal/refactor"
28 "golang.org/x/tools/internal/typeparams"
29 "golang.org/x/tools/internal/typesinternal"
30 "golang.org/x/tools/internal/versions"
31 )
32
33
34
35
36 type Caller struct {
37 Fset *token.FileSet
38 Types *types.Package
39 Info *types.Info
40 File *ast.File
41 Call *ast.CallExpr
42
43
44
45 CountUses func(pkgname *types.PkgName) int
46
47 path []ast.Node
48 enclosingFunc *ast.FuncDecl
49 }
50
51 type logger = func(string, ...any)
52
53
54
55 type Options struct {
56 Logf logger
57 IgnoreEffects bool
58 }
59
60
61 type Result struct {
62 Edits []refactor.Edit
63 Literalized bool
64 BindingDecl bool
65 }
66
67
68
69
70
71 func Inline(caller *Caller, callee *Callee, opts *Options) (*Result, error) {
72 copy := *opts
73 opts = ©
74
75 if opts.Logf == nil {
76 opts.Logf = func(string, ...any) {}
77 }
78
79 st := &state{
80 caller: caller,
81 callee: callee,
82 opts: opts,
83 }
84 return st.inline()
85 }
86
87
88 type state struct {
89 caller *Caller
90 callee *Callee
91 opts *Options
92 }
93
94 func (st *state) inline() (*Result, error) {
95 logf, caller, callee := st.opts.Logf, st.caller, st.callee
96
97 logf("inline %s @ %v",
98 debugFormatNode(caller.Fset, caller.Call),
99 caller.Fset.PositionFor(caller.Call.Lparen, false))
100
101 if ast.IsGenerated(caller.File) {
102 return nil, fmt.Errorf("cannot inline calls from generated files")
103 }
104
105 res, err := st.inlineCall()
106 if err != nil {
107 return nil, err
108 }
109
110
111 assert(res.old != nil, "old is nil")
112 assert(res.new != nil, "new is nil")
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135 if new, ok := res.new.(ast.Expr); ok {
136 parent := caller.path[slices.Index(caller.path, res.old)+1]
137 res.new = internalastutil.MaybeParenthesize(parent, res.old.(ast.Expr), new)
138 }
139
140
141
142
143
144
145
146
147
148
149
150
151
152 elideBraces := res.elideBraces
153 if !elideBraces {
154 if newBlock, ok := res.new.(*ast.BlockStmt); ok {
155 i := slices.Index(caller.path, res.old)
156 parent := caller.path[i+1]
157 var body []ast.Stmt
158 switch parent := parent.(type) {
159 case *ast.BlockStmt:
160 body = parent.List
161 case *ast.CommClause:
162 body = parent.Body
163 case *ast.CaseClause:
164 body = parent.Body
165 }
166 if body != nil {
167 callerNames := declares(body)
168
169
170
171 addFieldNames := func(fields *ast.FieldList) {
172 if fields != nil {
173 for _, field := range fields.List {
174 for _, id := range field.Names {
175 callerNames[id.Name] = true
176 }
177 }
178 }
179 }
180 switch f := caller.path[i+2].(type) {
181 case *ast.FuncDecl:
182 addFieldNames(f.Recv)
183 addFieldNames(f.Type.Params)
184 addFieldNames(f.Type.Results)
185 case *ast.FuncLit:
186 addFieldNames(f.Type.Params)
187 addFieldNames(f.Type.Results)
188 }
189
190 if len(callerLabels(caller.path)) > 0 {
191
192
193 logf("keeping block braces: caller uses control labels")
194 } else if intersects(declares(newBlock.List), callerNames) {
195 logf("keeping block braces: avoids name conflict")
196 } else {
197 elideBraces = true
198 }
199 }
200 }
201 }
202
203 var edits []refactor.Edit
204
205
206 {
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222 var out bytes.Buffer
223 if elideBraces {
224 for i, stmt := range res.new.(*ast.BlockStmt).List {
225 if i > 0 {
226 out.WriteByte('\n')
227 }
228 if err := format.Node(&out, caller.Fset, stmt); err != nil {
229 return nil, err
230 }
231 }
232 } else {
233 if err := format.Node(&out, caller.Fset, res.new); err != nil {
234 return nil, err
235 }
236 }
237
238 edits = append(edits, refactor.Edit{
239 Pos: res.old.Pos(),
240 End: res.old.End(),
241 NewText: out.Bytes(),
242 })
243 }
244
245
246
247
248
249
250
251 for _, imp := range res.newImports {
252
253 if !packagepath.CanImport(caller.Types.Path(), imp.path) {
254 return nil, fmt.Errorf("can't inline function %v as its body refers to inaccessible package %q", callee, imp.path)
255 }
256
257
258
259 name := ""
260 if imp.explicit {
261 name = imp.name
262 }
263 edits = append(edits, refactor.AddImportEdits(caller.File, name, imp.path)...)
264 }
265
266 literalized := false
267 if call, ok := res.new.(*ast.CallExpr); ok && is[*ast.FuncLit](call.Fun) {
268 literalized = true
269 }
270
271
272
273
274
275
276
277
278
279
280
281
282 for _, oldImport := range res.oldImports {
283 spec := oldImport.spec
284
285
286 pos := spec.Pos()
287 if doc := spec.Doc; doc != nil {
288 pos = doc.Pos()
289 }
290 end := spec.End()
291 if doc := spec.Comment; doc != nil {
292 end = doc.End()
293 }
294
295
296
297 for _, decl := range caller.File.Decls {
298 decl, ok := decl.(*ast.GenDecl)
299 if !(ok && decl.Tok == token.IMPORT) {
300 break
301 }
302 if internalastutil.NodeContainsPos(decl, spec.Pos()) && !decl.Rparen.IsValid() {
303
304 pos = decl.Pos()
305 if doc := decl.Doc; doc != nil {
306 pos = doc.Pos()
307 }
308 end = decl.End()
309 break
310 }
311 }
312
313 edits = append(edits, refactor.Edit{
314 Pos: pos,
315 End: end,
316 })
317 }
318
319 return &Result{
320 Edits: edits,
321 Literalized: literalized,
322 BindingDecl: res.bindingDecl,
323 }, nil
324 }
325
326
327 type oldImport struct {
328 pkgName *types.PkgName
329 spec *ast.ImportSpec
330 }
331
332
333 type newImport struct {
334 name string
335 path string
336 explicit bool
337 }
338
339
340 type importState struct {
341 logf func(string, ...any)
342 caller *Caller
343 importMap map[string][]string
344 newImports []newImport
345 oldImports []oldImport
346 }
347
348
349 func newImportState(logf func(string, ...any), caller *Caller, callee *gobCallee) *importState {
350
351
352
353 ist := &importState{
354 logf: logf,
355 caller: caller,
356 importMap: make(map[string][]string),
357 }
358
359
360
361 countUses := caller.CountUses
362 if countUses == nil {
363 uses := make(map[*types.PkgName]int)
364 for _, obj := range caller.Info.Uses {
365 if pkgname, ok := obj.(*types.PkgName); ok {
366 uses[pkgname]++
367 }
368 }
369 countUses = func(pkgname *types.PkgName) int {
370 return uses[pkgname]
371 }
372 }
373
374 for _, imp := range caller.File.Imports {
375 if pkgName, ok := importedPkgName(caller.Info, imp); ok &&
376 pkgName.Name() != "." &&
377 pkgName.Name() != "_" {
378
379
380
381
382
383
384
385
386
387
388
389
390 needed := true
391 if sel, ok := ast.Unparen(caller.Call.Fun).(*ast.SelectorExpr); ok &&
392 is[*ast.Ident](sel.X) &&
393 caller.Info.Uses[sel.X.(*ast.Ident)] == pkgName &&
394 countUses(pkgName) == 1 {
395 needed = false
396
397 for _, obj := range callee.FreeObjs {
398 if obj.PkgPath == pkgName.Imported().Path() && obj.Shadow[pkgName.Name()] == 0 {
399 needed = true
400 break
401 }
402 }
403 }
404
405
406
407 if needed {
408 path := pkgName.Imported().Path()
409 ist.importMap[path] = append(ist.importMap[path], pkgName.Name())
410 } else {
411 ist.oldImports = append(ist.oldImports, oldImport{pkgName: pkgName, spec: imp})
412 }
413 }
414 }
415 return ist
416 }
417
418
419
420
421
422 func (i *importState) importName(pkgPath string, shadow shadowMap) string {
423 for _, name := range i.importMap[pkgPath] {
424
425
426
427 if shadow[name] == 0 {
428 found := i.caller.lookup(name)
429 if is[*types.PkgName](found) || found == nil {
430 return name
431 }
432 }
433 }
434 return ""
435 }
436
437
438
439
440 func (i *importState) findNewLocalName(pkgName, calleePkgName string, shadow shadowMap) string {
441 newlyAdded := func(name string) bool {
442 return slices.ContainsFunc(i.newImports, func(n newImport) bool { return n.name == name })
443 }
444
445
446
447 shadowedInCaller := func(name string) bool {
448 obj := i.caller.lookup(name)
449 if obj == nil {
450 return false
451 }
452
453 return !slices.ContainsFunc(i.oldImports, func(o oldImport) bool { return o.pkgName == obj })
454 }
455
456
457
458
459
460
461
462
463
464 if shadow[calleePkgName] == 0 && !shadowedInCaller(calleePkgName) && !newlyAdded(calleePkgName) && calleePkgName != "init" {
465 return calleePkgName
466 }
467
468 base := pkgName
469 name := base
470 for n := 0; shadow[name] != 0 || shadowedInCaller(name) || newlyAdded(name) || name == "init"; n++ {
471 name = fmt.Sprintf("%s%d", base, n)
472 }
473
474 return name
475 }
476
477
478
479 func (i *importState) localName(pkgPath, pkgName, calleePkgName string, shadow shadowMap) string {
480
481 if name := i.importName(pkgPath, shadow); name != "" {
482 return name
483 }
484
485 name := i.findNewLocalName(pkgName, calleePkgName, shadow)
486 i.logf("adding import %s %q", name, pkgPath)
487
488
489 i.newImports = append(i.newImports, newImport{
490 name: name,
491 path: pkgPath,
492 explicit: name != pkgName || name != pathpkg.Base(pkgPath),
493 })
494 i.importMap[pkgPath] = append(i.importMap[pkgPath], name)
495 return name
496 }
497
498 type inlineCallResult struct {
499 newImports []newImport
500 oldImports []oldImport
501
502
503
504
505
506
507
508
509
510
511
512
513 elideBraces bool
514 bindingDecl bool
515 old, new ast.Node
516 }
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559 func (st *state) inlineCall() (*inlineCallResult, error) {
560 logf, caller, callee := st.opts.Logf, st.caller, &st.callee.impl
561
562 checkInfoFields(caller.Info)
563
564
565
566 calleeSymbol := typeutil.StaticCallee(caller.Info, caller.Call)
567 if calleeSymbol == nil {
568
569 return nil, fmt.Errorf("cannot inline: not a static function call")
570 }
571
572
573
574 samePkg := caller.Types.Path() == callee.PkgPath
575 if !samePkg && len(callee.Unexported) > 0 {
576 return nil, fmt.Errorf("cannot inline call to %s because body refers to non-exported %s",
577 callee.Name, callee.Unexported[0])
578 }
579
580
581
582
583 callerGoVersion := caller.Info.FileVersions[caller.File]
584 if callerGoVersion != "" && callee.GoVersion != "" && versions.Before(callerGoVersion, callee.GoVersion) {
585 return nil, fmt.Errorf("cannot inline call to %s (declared using %s) into a file using %s",
586 callee.Name, callee.GoVersion, callerGoVersion)
587 }
588
589
590
591
592
593 caller.path, _ = astutil.PathEnclosingInterval(caller.File, caller.Call.Pos(), caller.Call.End())
594 for _, n := range caller.path {
595 if decl, ok := n.(*ast.FuncDecl); ok {
596 caller.enclosingFunc = decl
597 break
598 }
599 }
600
601
602
603
604 var assign1 func(v *types.Var) bool
605 {
606 updatedLocals := make(map[*types.Var]bool)
607 if caller.enclosingFunc != nil {
608 escape(caller.Info, caller.enclosingFunc, func(v *types.Var, _ bool) {
609 updatedLocals[v] = true
610 })
611 logf("multiple-assignment vars: %v", updatedLocals)
612 }
613 assign1 = func(v *types.Var) bool { return !updatedLocals[v] }
614 }
615
616
617 istate := newImportState(logf, caller, callee)
618
619
620 objRenames, err := st.renameFreeObjs(istate)
621 if err != nil {
622 return nil, err
623 }
624
625 res := &inlineCallResult{
626 newImports: istate.newImports,
627 oldImports: istate.oldImports,
628 }
629
630
631 calleeFset, calleeDecl, err := parseCompact(callee.Content)
632 if err != nil {
633 return nil, err
634 }
635
636
637
638 replaceCalleeID := func(offset int, repl ast.Expr, unpackVariadic bool) {
639 path, id := findIdent(calleeDecl, calleeDecl.Pos()+token.Pos(offset))
640 logf("- replace id %q @ #%d to %q", id.Name, offset, debugFormatNode(calleeFset, repl))
641
642 if lit, ok := repl.(*ast.CompositeLit); ok && unpackVariadic && len(path) > 0 {
643 if call, ok := last(path).(*ast.CallExpr); ok &&
644 call.Ellipsis.IsValid() &&
645 id == last(call.Args) {
646
647 call.Args = append(call.Args[:len(call.Args)-1], lit.Elts...)
648 call.Ellipsis = token.NoPos
649 return
650 }
651 }
652 if len(path) > 0 {
653 repl = internalastutil.MaybeParenthesize(last(path), id, repl)
654 }
655 replaceNode(calleeDecl, id, repl)
656 }
657
658
659
660 for _, ref := range callee.FreeRefs {
661 if repl := objRenames[ref.Object]; repl != nil {
662 replaceCalleeID(ref.Offset, repl, false)
663 }
664 }
665
666
667
668 args, err := st.arguments(caller, calleeDecl, assign1)
669 if err != nil {
670 return nil, err
671 }
672
673
674
675 var params []*parameter
676 {
677 sig := calleeSymbol.Type().(*types.Signature)
678 if sig.Recv() != nil {
679 params = append(params, ¶meter{
680 obj: sig.Recv(),
681 fieldType: calleeDecl.Recv.List[0].Type,
682 info: callee.Params[0],
683 })
684 }
685
686
687 var types []ast.Expr
688 for _, field := range calleeDecl.Type.Params.List {
689 if field.Names == nil {
690 types = append(types, field.Type)
691 } else {
692 for range field.Names {
693 types = append(types, field.Type)
694 }
695 }
696 }
697
698 for i := 0; i < sig.Params().Len(); i++ {
699 params = append(params, ¶meter{
700 obj: sig.Params().At(i),
701 fieldType: types[i],
702 info: callee.Params[len(params)],
703 })
704 }
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719 if sig.Variadic() {
720 lastParam := last(params)
721 if len(args) > 0 && last(args).spread {
722
723 lastParam.variadic = true
724 } else {
725
726
727
728 lastParamField := last(calleeDecl.Type.Params.List)
729 lastParamField.Type = &ast.ArrayType{
730 Elt: lastParamField.Type.(*ast.Ellipsis).Elt,
731 }
732
733 if caller.Call.Ellipsis.IsValid() {
734
735
736 } else {
737
738
739
740
741
742 n := len(params) - 1
743 ordinary, extra := args[:n], args[n:]
744 var elts []ast.Expr
745 freevars := make(map[string]bool)
746 pure, effects := true, false
747 for _, arg := range extra {
748 elts = append(elts, arg.expr)
749 pure = pure && arg.pure
750 effects = effects || arg.effects
751 maps.Copy(freevars, arg.freevars)
752 }
753 args = append(ordinary, &argument{
754 expr: &ast.CompositeLit{
755 Type: lastParamField.Type,
756 Elts: elts,
757 },
758 typ: lastParam.obj.Type(),
759 constant: nil,
760 pure: pure,
761 effects: effects,
762 duplicable: false,
763 freevars: freevars,
764 variadic: true,
765 })
766 }
767 }
768 }
769 }
770
771 typeArgs := st.typeArguments(caller.Call)
772 if len(typeArgs) != len(callee.TypeParams) {
773 return nil, fmt.Errorf("cannot inline: type parameter inference is not yet supported")
774 }
775 if err := substituteTypeParams(logf, callee.TypeParams, typeArgs, params, replaceCalleeID); err != nil {
776 return nil, err
777 }
778
779
780 for i, arg := range args {
781 logf("arg #%d: %s pure=%t effects=%t duplicable=%t free=%v type=%v",
782 i, debugFormatNode(caller.Fset, arg.expr),
783 arg.pure, arg.effects, arg.duplicable, arg.freevars, arg.typ)
784 }
785
786
787
788
789
790
791 substitute(logf, caller, params, args, callee.Effects, callee.Falcon, replaceCalleeID)
792
793
794 updateCalleeParams(calleeDecl, params)
795
796
797 bindingDecl := createBindingDecl(logf, caller, args, calleeDecl, callee.Results)
798
799 var remainingArgs []ast.Expr
800 for _, arg := range args {
801 if arg != nil {
802 remainingArgs = append(remainingArgs, arg.expr)
803 }
804 }
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830 if len(calleeDecl.Body.List) == 0 {
831 logf("strategy: reduce call to empty body")
832
833
834
835
836
837 if stmt := callStmt(caller.path, false); stmt != nil {
838 res.old = stmt
839 if nargs := len(remainingArgs); nargs > 0 {
840
841
842
843
844
845
846
847
848
849 if last := last(args); last != nil && last.spread {
850 nspread := last.typ.(*types.Tuple).Len()
851 if len(args) > 1 {
852
853 res.new = &ast.GenDecl{
854 Tok: token.VAR,
855 Specs: []ast.Spec{
856 &ast.ValueSpec{
857 Names: []*ast.Ident{makeIdent("_")},
858 Values: []ast.Expr{args[0].expr},
859 },
860 &ast.ValueSpec{
861 Names: blanks[*ast.Ident](nspread),
862 Values: []ast.Expr{args[1].expr},
863 },
864 },
865 }
866 return res, nil
867 }
868
869
870 nargs = nspread
871 }
872
873 res.new = &ast.AssignStmt{
874 Lhs: blanks[ast.Expr](nargs),
875 Tok: token.ASSIGN,
876 Rhs: remainingArgs,
877 }
878
879 } else {
880
881 res.new = &ast.EmptyStmt{}
882 }
883 return res, nil
884 }
885 }
886
887
888
889
890 allResultsUnreferenced := forall(callee.Results, func(i int, r *paramInfo) bool { return len(r.Refs) == 0 })
891 needBindingDecl := !allResultsUnreferenced ||
892 exists(params, func(i int, p *parameter) bool { return p != nil })
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918 if len(calleeDecl.Body.List) == 1 &&
919 is[*ast.ReturnStmt](calleeDecl.Body.List[0]) &&
920 len(calleeDecl.Body.List[0].(*ast.ReturnStmt).Results) > 0 {
921 results := calleeDecl.Body.List[0].(*ast.ReturnStmt).Results
922
923 parent, grandparent := callContext(caller.path)
924
925
926 if stmt, ok := parent.(*ast.ExprStmt); ok &&
927 (!needBindingDecl || bindingDecl != nil) {
928 logf("strategy: reduce stmt-context call to { return exprs }")
929 clearPositions(calleeDecl.Body)
930
931 if callee.ValidForCallStmt {
932 logf("callee body is valid as statement")
933
934 if !needBindingDecl {
935
936 res.old = caller.Call
937 res.new = results[0]
938 } else {
939
940 res.bindingDecl = true
941 res.old = stmt
942 res.new = &ast.BlockStmt{
943 List: []ast.Stmt{
944 bindingDecl.stmt,
945 &ast.ExprStmt{X: results[0]},
946 },
947 }
948 }
949 } else {
950 logf("callee body is not valid as statement")
951
952
953
954
955 discard := &ast.AssignStmt{
956 Lhs: blanks[ast.Expr](callee.NumResults),
957 Tok: token.ASSIGN,
958 Rhs: results,
959 }
960 res.old = stmt
961 if !needBindingDecl {
962
963 res.new = discard
964 } else {
965
966 res.bindingDecl = true
967 res.new = &ast.BlockStmt{
968 List: []ast.Stmt{
969 bindingDecl.stmt,
970 discard,
971 },
972 }
973 }
974 }
975 return res, nil
976 }
977
978
979
980
981
982 if stmt, ok := parent.(*ast.AssignStmt); ok &&
983 is[*ast.BlockStmt](grandparent) &&
984 (!needBindingDecl || (bindingDecl != nil && len(bindingDecl.names) == 0)) {
985
986
987 if newStmts, ok := st.assignStmts(stmt, results, istate.importName); ok {
988 logf("strategy: reduce assign-context call to { return exprs }")
989
990 clearPositions(calleeDecl.Body)
991
992 block := &ast.BlockStmt{
993 List: newStmts,
994 }
995 if needBindingDecl {
996 res.bindingDecl = true
997 block.List = prepend(bindingDecl.stmt, block.List...)
998 }
999
1000
1001
1002
1003 res.elideBraces = true
1004 res.old = stmt
1005 res.new = block
1006 return res, nil
1007 }
1008 }
1009
1010
1011 if !needBindingDecl {
1012 clearPositions(calleeDecl.Body)
1013
1014 anyNonTrivialReturns := hasNonTrivialReturn(callee.Returns)
1015
1016 if callee.NumResults == 1 {
1017 logf("strategy: reduce expr-context call to { return expr }")
1018
1019
1020
1021 if anyNonTrivialReturns {
1022 results[0] = convert(calleeDecl.Type.Results.List[0].Type, results[0])
1023 }
1024
1025 res.old = caller.Call
1026 res.new = results[0]
1027 return res, nil
1028
1029 } else if !anyNonTrivialReturns {
1030 logf("strategy: reduce spread-context call to { return expr }")
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047 res.old = parent
1048 switch context := parent.(type) {
1049 case *ast.AssignStmt:
1050
1051 assign := shallowCopy(context)
1052 assign.Rhs = results
1053 res.new = assign
1054 case *ast.ValueSpec:
1055
1056 spec := shallowCopy(context)
1057 spec.Values = results
1058 res.new = spec
1059 case *ast.CallExpr:
1060
1061 call := shallowCopy(context)
1062 call.Args = results
1063 res.new = call
1064 case *ast.ReturnStmt:
1065
1066 ret := shallowCopy(context)
1067 ret.Results = results
1068 res.new = ret
1069 default:
1070 return nil, fmt.Errorf("internal error: unexpected context %T for spread call", context)
1071 }
1072 return res, nil
1073 }
1074 }
1075 }
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101 parent, _ := callContext(caller.path)
1102 if ret, ok := parent.(*ast.ReturnStmt); ok &&
1103 len(ret.Results) == 1 &&
1104 tailCallSafeReturn(caller, calleeSymbol, callee) &&
1105 !callee.HasBareReturn &&
1106 (!needBindingDecl || bindingDecl != nil) &&
1107 !hasLabelConflict(caller.path, callee.Labels) &&
1108 allResultsUnreferenced {
1109 logf("strategy: reduce tail-call")
1110 body := calleeDecl.Body
1111 clearPositions(body)
1112 if needBindingDecl {
1113 res.bindingDecl = true
1114 body.List = prepend(bindingDecl.stmt, body.List...)
1115 }
1116 res.old = ret
1117 res.new = body
1118 return res, nil
1119 }
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137 if stmt := callStmt(caller.path, true); stmt != nil &&
1138 (!needBindingDecl || bindingDecl != nil) &&
1139 !callee.HasDefer &&
1140 !hasLabelConflict(caller.path, callee.Labels) &&
1141 len(callee.Returns) == 0 {
1142 logf("strategy: reduce stmt-context call to { stmts }")
1143 body := calleeDecl.Body
1144 var repl ast.Stmt = body
1145 clearPositions(repl)
1146 if needBindingDecl {
1147 body.List = prepend(bindingDecl.stmt, body.List...)
1148 }
1149 res.old = stmt
1150 res.new = repl
1151 return res, nil
1152 }
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179 if len(args) == 2 && args[0] != nil && args[1] != nil && is[*types.Tuple](args[1].typ) {
1180 return nil, fmt.Errorf("can't yet inline spread call to method")
1181 }
1182
1183
1184
1185
1186
1187 logf("strategy: literalization")
1188 funcLit := &ast.FuncLit{
1189 Type: calleeDecl.Type,
1190 Body: calleeDecl.Body,
1191 }
1192
1193
1194
1195 clearPositions(funcLit)
1196
1197
1198
1199
1200
1201
1202
1203
1204 if bindingDecl != nil && allResultsUnreferenced {
1205 funcLit.Type.Params.List = nil
1206 remainingArgs = nil
1207 res.bindingDecl = true
1208 funcLit.Body.List = prepend(bindingDecl.stmt, funcLit.Body.List...)
1209 }
1210
1211
1212
1213 newCall := &ast.CallExpr{
1214 Fun: funcLit,
1215 Ellipsis: token.NoPos,
1216 Args: remainingArgs,
1217 }
1218 res.old = caller.Call
1219 res.new = newCall
1220 return res, nil
1221 }
1222
1223
1224
1225
1226 func (st *state) renameFreeObjs(istate *importState) ([]ast.Expr, error) {
1227 caller, callee := st.caller, &st.callee.impl
1228 objRenames := make([]ast.Expr, len(callee.FreeObjs))
1229 for i, obj := range callee.FreeObjs {
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250 var newName ast.Expr
1251 if obj.Kind == "pkgname" {
1252
1253 n := istate.localName(obj.PkgPath, obj.PkgName, obj.Name, obj.Shadow)
1254 newName = makeIdent(n)
1255 } else if !obj.ValidPos {
1256
1257
1258 found := caller.lookup(obj.Name)
1259 if found.Pos().IsValid() {
1260 return nil, fmt.Errorf("cannot inline, because the callee refers to built-in %q, which in the caller is shadowed by a %s (declared at line %d)",
1261 obj.Name, objectKind(found),
1262 caller.Fset.PositionFor(found.Pos(), false).Line)
1263 }
1264
1265 } else {
1266
1267
1268 qualify := false
1269 if obj.PkgPath == callee.PkgPath {
1270
1271 if caller.Types.Path() == callee.PkgPath {
1272
1273
1274
1275
1276
1277
1278 found := caller.lookup(obj.Name)
1279 if found != nil && !isPkgLevel(found) {
1280 return nil, fmt.Errorf("cannot inline, because the callee refers to %s %q, which in the caller is shadowed by a %s (declared at line %d)",
1281 obj.Kind, obj.Name,
1282 objectKind(found),
1283 caller.Fset.PositionFor(found.Pos(), false).Line)
1284 }
1285 } else {
1286
1287 qualify = true
1288 }
1289 } else {
1290
1291
1292
1293 qualify = true
1294 }
1295
1296
1297 if qualify {
1298 pkgName := istate.localName(obj.PkgPath, obj.PkgName, obj.PkgName, obj.Shadow)
1299 newName = &ast.SelectorExpr{
1300 X: makeIdent(pkgName),
1301 Sel: makeIdent(obj.Name),
1302 }
1303 }
1304 }
1305 objRenames[i] = newName
1306 }
1307 return objRenames, nil
1308 }
1309
1310 type argument struct {
1311 expr ast.Expr
1312 typ types.Type
1313 constant constant.Value
1314 spread bool
1315 pure bool
1316 effects bool
1317 duplicable bool
1318 freevars map[string]bool
1319 variadic bool
1320 desugaredRecv bool
1321 }
1322
1323
1324
1325
1326 func (st *state) typeArguments(call *ast.CallExpr) []*argument {
1327 var exprs []ast.Expr
1328 switch d := ast.Unparen(call.Fun).(type) {
1329 case *ast.IndexExpr:
1330 exprs = []ast.Expr{d.Index}
1331 case *ast.IndexListExpr:
1332 exprs = d.Indices
1333 default:
1334
1335 return nil
1336 }
1337 var args []*argument
1338 for _, e := range exprs {
1339 arg := &argument{expr: e, freevars: freeVars(st.caller.Info, e)}
1340 args = append(args, arg)
1341 }
1342 return args
1343 }
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375 func (st *state) arguments(caller *Caller, calleeDecl *ast.FuncDecl, assign1 func(*types.Var) bool) ([]*argument, error) {
1376 var args []*argument
1377
1378 callArgs := caller.Call.Args
1379 if calleeDecl.Recv != nil {
1380 if len(st.callee.impl.TypeParams) > 0 {
1381 return nil, fmt.Errorf("cannot inline: generic methods not yet supported")
1382 }
1383 sel := ast.Unparen(caller.Call.Fun).(*ast.SelectorExpr)
1384 seln := caller.Info.Selections[sel]
1385 var recvArg ast.Expr
1386 switch seln.Kind() {
1387 case types.MethodVal:
1388 recvArg = sel.X
1389 case types.MethodExpr:
1390 recvArg = callArgs[0]
1391 callArgs = callArgs[1:]
1392 }
1393 if recvArg != nil {
1394
1395
1396
1397 arg := &argument{
1398 expr: recvArg,
1399 typ: caller.Info.TypeOf(recvArg),
1400 constant: caller.Info.Types[recvArg].Value,
1401 pure: pure(caller.Info, assign1, recvArg),
1402 effects: st.effects(caller.Info, recvArg),
1403 duplicable: duplicable(caller.Info, recvArg),
1404 freevars: freeVars(caller.Info, recvArg),
1405 }
1406 recvArg = nil
1407
1408
1409 args = append(args, arg)
1410
1411
1412
1413 indices := seln.Index()
1414 for _, index := range indices[:len(indices)-1] {
1415 fld := typeparams.CoreType(typeparams.Deref(arg.typ)).(*types.Struct).Field(index)
1416 if fld.Pkg() != caller.Types && !fld.Exported() {
1417 return nil, fmt.Errorf("in %s, implicit reference to unexported field .%s cannot be made explicit",
1418 debugFormatNode(caller.Fset, caller.Call.Fun),
1419 fld.Name())
1420 }
1421 if isPointer(arg.typ) {
1422 arg.pure = false
1423 }
1424 arg.expr = &ast.SelectorExpr{
1425 X: arg.expr,
1426 Sel: makeIdent(fld.Name()),
1427 }
1428 arg.typ = fld.Type()
1429 arg.duplicable = false
1430 }
1431
1432
1433 argIsPtr := isPointer(arg.typ)
1434 paramIsPtr := isPointer(seln.Obj().Type().Underlying().(*types.Signature).Recv().Type())
1435 if !argIsPtr && paramIsPtr {
1436
1437 arg.expr = &ast.UnaryExpr{Op: token.AND, X: arg.expr}
1438 arg.typ = types.NewPointer(arg.typ)
1439 arg.desugaredRecv = true
1440 } else if argIsPtr && !paramIsPtr {
1441
1442 arg.expr = &ast.StarExpr{X: arg.expr}
1443 arg.typ = typeparams.Deref(arg.typ)
1444 arg.duplicable = false
1445 arg.pure = false
1446 arg.desugaredRecv = true
1447 }
1448 }
1449 }
1450 for _, expr := range callArgs {
1451 tv := caller.Info.Types[expr]
1452 args = append(args, &argument{
1453 expr: expr,
1454 typ: tv.Type,
1455 constant: tv.Value,
1456 spread: is[*types.Tuple](tv.Type),
1457 pure: pure(caller.Info, assign1, expr),
1458 effects: st.effects(caller.Info, expr),
1459 duplicable: duplicable(caller.Info, expr),
1460 freevars: freeVars(caller.Info, expr),
1461 })
1462 }
1463
1464
1465
1466
1467
1468
1469
1470
1471
1472
1473
1474
1475
1476
1477
1478
1479
1480
1481
1482 for _, arg := range args {
1483 if arg.constant == nil {
1484 continue
1485 }
1486 info := &types.Info{Types: make(map[ast.Expr]types.TypeAndValue)}
1487 if err := types.CheckExpr(caller.Fset, caller.Types, caller.Call.Pos(), arg.expr, info); err != nil {
1488 return nil, err
1489 }
1490 arg.typ = info.TypeOf(arg.expr)
1491 }
1492
1493 return args, nil
1494 }
1495
1496 type parameter struct {
1497 obj *types.Var
1498 fieldType ast.Expr
1499 info *paramInfo
1500 variadic bool
1501 }
1502
1503
1504
1505
1506
1507 type replacer = func(offset int, repl ast.Expr, unpackVariadic bool)
1508
1509
1510
1511 func substituteTypeParams(logf logger, typeParams []*paramInfo, typeArgs []*argument, params []*parameter, replace replacer) error {
1512 assert(len(typeParams) == len(typeArgs), "mismatched number of type params/args")
1513 for i, paramInfo := range typeParams {
1514 arg := typeArgs[i]
1515
1516 for free := range arg.freevars {
1517 if paramInfo.Shadow[free] != 0 {
1518 return fmt.Errorf("cannot inline: type argument #%d (type parameter %s) is shadowed", i, paramInfo.Name)
1519 }
1520 }
1521 logf("replacing type param %s with %s", paramInfo.Name, debugFormatNode(token.NewFileSet(), arg.expr))
1522 for _, ref := range paramInfo.Refs {
1523 replace(ref.Offset, internalastutil.CloneNode(arg.expr), false)
1524 }
1525
1526
1527
1528 for _, p := range params {
1529 if id, ok := p.fieldType.(*ast.Ident); ok && id.Name == paramInfo.Name {
1530 p.fieldType = arg.expr
1531 } else {
1532 for _, id := range identsNamed(p.fieldType, paramInfo.Name) {
1533 replaceNode(p.fieldType, id, arg.expr)
1534 }
1535 }
1536 }
1537 }
1538 return nil
1539 }
1540
1541 func identsNamed(n ast.Node, name string) []*ast.Ident {
1542 var ids []*ast.Ident
1543 ast.Inspect(n, func(n ast.Node) bool {
1544 if id, ok := n.(*ast.Ident); ok && id.Name == name {
1545 ids = append(ids, id)
1546 }
1547 return true
1548 })
1549 return ids
1550 }
1551
1552
1553
1554
1555
1556
1557
1558
1559
1560
1561
1562
1563
1564
1565
1566
1567
1568
1569
1570
1571 func substitute(logf logger, caller *Caller, params []*parameter, args []*argument, effects []int, falcon falconResult, replace replacer) {
1572
1573
1574
1575
1576
1577
1578
1579 assert(len(args) <= len(params), "too many arguments")
1580
1581
1582
1583
1584
1585
1586
1587
1588
1589
1590
1591
1592 sg := make(substGraph)
1593 next:
1594 for i, param := range params {
1595 arg := args[i]
1596
1597
1598
1599
1600
1601
1602
1603
1604 if arg.spread {
1605
1606 logf("keeping param %q and following ones: argument %s is spread",
1607 param.info.Name, debugFormatNode(caller.Fset, arg.expr))
1608 return
1609 }
1610 assert(!param.variadic, "unsimplified variadic parameter")
1611 if param.info.Escapes {
1612 logf("keeping param %q: escapes from callee", param.info.Name)
1613 continue
1614 }
1615 if param.info.Assigned {
1616 logf("keeping param %q: assigned by callee", param.info.Name)
1617 continue
1618 }
1619 if len(param.info.Refs) > 1 && !arg.duplicable {
1620 logf("keeping param %q: argument is not duplicable", param.info.Name)
1621 continue
1622 }
1623 if len(param.info.Refs) == 0 {
1624 if arg.effects {
1625 logf("keeping param %q: though unreferenced, it has effects", param.info.Name)
1626 continue
1627 }
1628
1629
1630
1631
1632 if caller.enclosingFunc != nil {
1633 for free := range arg.freevars {
1634
1635
1636
1637 if v, ok := caller.lookup(free).(*types.Var); ok && within(v.Pos(), caller.enclosingFunc.Body) && !isUsedOutsideCall(caller, v) {
1638
1639
1640
1641 usedElsewhere := func() bool {
1642 for i, param := range params {
1643 if i < len(args) && len(param.info.Refs) > 0 {
1644 for name := range args[i].freevars {
1645 if caller.lookup(name) == v {
1646 return true
1647 }
1648 }
1649 }
1650 }
1651 return false
1652 }
1653 if !usedElsewhere() {
1654 logf("keeping param %q: arg contains perhaps the last reference to caller local %v @ %v",
1655 param.info.Name, v, caller.Fset.PositionFor(v.Pos(), false))
1656 continue next
1657 }
1658 }
1659 }
1660 }
1661 }
1662
1663
1664
1665
1666
1667
1668
1669
1670
1671
1672
1673
1674
1675
1676 sg[arg] = nil
1677 for free := range arg.freevars {
1678 switch s := param.info.Shadow[free]; {
1679 case s < 0:
1680
1681 delete(sg, arg)
1682 case s > 0:
1683
1684
1685 if s > len(args) {
1686
1687
1688 delete(sg, arg)
1689 }
1690 if edges, ok := sg[arg]; ok {
1691 sg[arg] = append(edges, args[s-1])
1692 }
1693 }
1694 }
1695 }
1696
1697
1698 sg.prune()
1699
1700
1701
1702
1703
1704
1705
1706
1707
1708
1709
1710
1711
1712
1713
1714
1715
1716
1717
1718
1719
1720
1721
1722
1723
1724
1725
1726
1727
1728
1729
1730
1731
1732
1733
1734
1735
1736
1737
1738
1739
1740
1741
1742
1743
1744
1745
1746
1747
1748
1749
1750
1751
1752
1753
1754
1755
1756
1757 for checkFalconConstraints(logf, params, args, falcon, sg) {
1758 sg.prune()
1759 }
1760
1761
1762
1763
1764
1765
1766
1767 for resolveEffects(logf, args, effects, sg) {
1768 sg.prune()
1769 }
1770
1771
1772 for i, param := range params {
1773 if arg := args[i]; sg.has(arg) {
1774
1775
1776
1777
1778
1779
1780 logf("replacing parameter %q by argument %q",
1781 param.info.Name, debugFormatNode(caller.Fset, arg.expr))
1782 for _, ref := range param.info.Refs {
1783
1784 argExpr := arg.expr
1785
1786
1787
1788
1789
1790
1791 if ref.IsSelectionOperand && arg.desugaredRecv {
1792 switch e := argExpr.(type) {
1793 case *ast.UnaryExpr:
1794 argExpr = e.X
1795 case *ast.StarExpr:
1796 argExpr = e.X
1797 }
1798 }
1799
1800
1801
1802
1803
1804
1805
1806
1807
1808
1809
1810
1811
1812
1813
1814
1815
1816
1817
1818
1819
1820
1821
1822
1823
1824
1825
1826
1827
1828
1829
1830
1831
1832
1833
1834 needType := ref.AffectsInference ||
1835 (ref.Assignable && ref.IfaceAssignment && !param.info.IsInterface) ||
1836 (!ref.Assignable && !trivialConversion(arg.constant, arg.typ, param.obj.Type()))
1837
1838 if needType &&
1839 !types.Identical(types.Default(arg.typ), param.obj.Type()) {
1840
1841
1842 if call, ok := argExpr.(*ast.CallExpr); ok && len(call.Args) == 1 {
1843 if typ, ok := isConversion(caller.Info, call); ok && isNonTypeParamInterface(typ) {
1844 argExpr = call.Args[0]
1845 }
1846 }
1847
1848 argExpr = convert(param.fieldType, argExpr)
1849 logf("param %q (offset %d): adding explicit %s -> %s conversion around argument",
1850 param.info.Name, ref.Offset, arg.typ, param.obj.Type())
1851 }
1852 replace(ref.Offset, internalastutil.CloneNode(argExpr).(ast.Expr), arg.variadic)
1853 }
1854 params[i] = nil
1855 args[i] = nil
1856 }
1857 }
1858 }
1859
1860
1861
1862
1863
1864 func isConversion(info *types.Info, call *ast.CallExpr) (types.Type, bool) {
1865 if tv, ok := info.Types[call.Fun]; ok && tv.IsType() {
1866 return tv.Type, true
1867 }
1868 return nil, false
1869 }
1870
1871
1872
1873 func isNonTypeParamInterface(t types.Type) bool {
1874 return !typeparams.IsTypeParam(t) && types.IsInterface(t)
1875 }
1876
1877
1878
1879 func isUsedOutsideCall(caller *Caller, v *types.Var) bool {
1880 used := false
1881 ast.Inspect(caller.enclosingFunc.Body, func(n ast.Node) bool {
1882 if n == caller.Call {
1883 return false
1884 }
1885 switch n := n.(type) {
1886 case *ast.Ident:
1887 if use := caller.Info.Uses[n]; use == v {
1888 used = true
1889 }
1890 case *ast.FuncType:
1891
1892 for _, fld := range n.Params.List {
1893 for _, n := range fld.Names {
1894 if def := caller.Info.Defs[n]; def == v {
1895 used = true
1896 }
1897 }
1898 }
1899 }
1900 return !used
1901 })
1902 return used
1903 }
1904
1905
1906
1907
1908
1909
1910
1911
1912
1913
1914 func checkFalconConstraints(logf logger, params []*parameter, args []*argument, falcon falconResult, sg substGraph) bool {
1915
1916
1917 pkg := types.NewPackage("falcon", "falcon")
1918
1919
1920 for _, typ := range falcon.Types {
1921 logf("falcon env: type %s %s", typ.Name, types.Typ[typ.Kind])
1922 pkg.Scope().Insert(types.NewTypeName(token.NoPos, pkg, typ.Name, types.Typ[typ.Kind]))
1923 }
1924
1925
1926 nconst := 0
1927 for i, param := range params {
1928 name := param.info.Name
1929 if name == "" {
1930 continue
1931 }
1932 arg := args[i]
1933 if arg.constant != nil && sg.has(arg) && param.info.FalconType != "" {
1934 t := pkg.Scope().Lookup(param.info.FalconType).Type()
1935 pkg.Scope().Insert(types.NewConst(token.NoPos, pkg, name, t, arg.constant))
1936 logf("falcon env: const %s %s = %v", name, param.info.FalconType, arg.constant)
1937 nconst++
1938 } else {
1939 v := types.NewVar(token.NoPos, pkg, name, arg.typ)
1940 typesinternal.SetVarKind(v, typesinternal.PackageVar)
1941 pkg.Scope().Insert(v)
1942 logf("falcon env: var %s %s", name, arg.typ)
1943 }
1944 }
1945 if nconst == 0 {
1946 return false
1947 }
1948
1949
1950 fset := token.NewFileSet()
1951 removed := false
1952 for _, falcon := range falcon.Constraints {
1953 expr, err := parser.ParseExprFrom(fset, "falcon", falcon, 0)
1954 if err != nil {
1955 panic(fmt.Sprintf("failed to parse falcon constraint %s: %v", falcon, err))
1956 }
1957 if err := types.CheckExpr(fset, pkg, token.NoPos, expr, nil); err != nil {
1958 logf("falcon: constraint %s violated: %v", falcon, err)
1959 for j, arg := range args {
1960 if arg.constant != nil && sg.has(arg) {
1961 logf("keeping param %q due falcon violation", params[j].info.Name)
1962 removed = sg.remove(arg) || removed
1963 }
1964 }
1965 break
1966 }
1967 logf("falcon: constraint %s satisfied", falcon)
1968 }
1969 return removed
1970 }
1971
1972
1973
1974
1975
1976
1977
1978
1979
1980
1981
1982
1983
1984
1985
1986
1987
1988
1989
1990
1991
1992
1993
1994
1995
1996
1997
1998
1999
2000
2001
2002
2003
2004
2005
2006
2007
2008
2009
2010
2011
2012
2013
2014
2015
2016 func resolveEffects(logf logger, args []*argument, effects []int, sg substGraph) bool {
2017 effectStr := func(effects bool, idx int) string {
2018 i := fmt.Sprint(idx)
2019 if idx == len(args) {
2020 i = "∞"
2021 }
2022 return string("RW"[btoi(effects)]) + i
2023 }
2024 removed := false
2025 for i, argi := range slices.Backward(args) {
2026 if sg.has(argi) && !argi.pure {
2027
2028 idx := slices.Index(effects, i)
2029 if idx >= 0 {
2030 for _, j := range effects[:idx] {
2031 var (
2032 ji int
2033 jw bool
2034 )
2035 if j == winf || j == rinf {
2036 jw = j == winf
2037 ji = len(args)
2038 } else {
2039 jw = args[j].effects
2040 ji = j
2041 }
2042 if ji > i && (jw || argi.effects) {
2043 logf("binding argument %s: preceded by %s",
2044 effectStr(argi.effects, i), effectStr(jw, ji))
2045
2046 removed = sg.remove(argi) || removed
2047 break
2048 }
2049 }
2050 }
2051 }
2052 if !sg.has(argi) {
2053 for j := range i {
2054 argj := args[j]
2055 if argj.pure {
2056 continue
2057 }
2058 if (argi.effects || argj.effects) && sg.has(argj) {
2059 logf("binding argument %s: %s is bound",
2060 effectStr(argj.effects, j), effectStr(argi.effects, i))
2061
2062 removed = sg.remove(argj) || removed
2063 }
2064 }
2065 }
2066 }
2067 return removed
2068 }
2069
2070
2071
2072
2073
2074
2075
2076
2077
2078
2079
2080
2081
2082
2083
2084
2085
2086
2087 type substGraph map[*argument][]*argument
2088
2089
2090 func (g substGraph) has(arg *argument) bool {
2091 _, ok := g[arg]
2092 return ok
2093 }
2094
2095
2096
2097
2098
2099
2100
2101
2102 func (g substGraph) remove(arg *argument) bool {
2103 pre := len(g)
2104 delete(g, arg)
2105 return len(g) < pre
2106 }
2107
2108
2109
2110 func (g substGraph) prune() {
2111
2112
2113
2114
2115
2116
2117
2118
2119
2120
2121
2122
2123
2124
2125
2126
2127 var visit func(*argument, map[*argument]unit) bool
2128 visit = func(arg *argument, seen map[*argument]unit) bool {
2129 deps, ok := g[arg]
2130 if !ok {
2131 return false
2132 }
2133 if _, ok := seen[arg]; !ok {
2134 seen[arg] = unit{}
2135 for _, dep := range deps {
2136 if !visit(dep, seen) {
2137 delete(g, arg)
2138 return false
2139 }
2140 }
2141 }
2142 return true
2143 }
2144 for arg := range g {
2145
2146
2147
2148
2149 visit(arg, make(map[*argument]unit))
2150 }
2151 }
2152
2153
2154
2155
2156 func updateCalleeParams(calleeDecl *ast.FuncDecl, params []*parameter) {
2157
2158
2159
2160
2161
2162
2163
2164
2165
2166
2167 paramIdx := 0
2168 var newParams []*ast.Field
2169 filterParams := func(field *ast.Field) {
2170 var names []*ast.Ident
2171 if field.Names == nil {
2172
2173 if params[paramIdx] != nil {
2174
2175
2176
2177 names = append(names, makeIdent("_"))
2178 }
2179 paramIdx++
2180 } else {
2181
2182
2183
2184 for _, id := range field.Names {
2185 if pinfo := params[paramIdx]; pinfo != nil {
2186
2187
2188
2189
2190 if len(pinfo.info.Refs) == 0 {
2191 id = makeIdent("_")
2192 }
2193 names = append(names, id)
2194 }
2195 paramIdx++
2196 }
2197 }
2198 if names != nil {
2199 newParams = append(newParams, &ast.Field{
2200 Names: names,
2201 Type: field.Type,
2202 })
2203 }
2204 }
2205 if calleeDecl.Recv != nil {
2206 filterParams(calleeDecl.Recv.List[0])
2207 calleeDecl.Recv = nil
2208 }
2209 for _, field := range calleeDecl.Type.Params.List {
2210 filterParams(field)
2211 }
2212 calleeDecl.Type.Params.List = newParams
2213 }
2214
2215
2216
2217 type bindingDeclInfo struct {
2218 names map[string]bool
2219 stmt ast.Stmt
2220 }
2221
2222
2223
2224
2225
2226
2227
2228
2229
2230
2231
2232
2233
2234
2235
2236
2237
2238
2239
2240
2241
2242
2243
2244
2245
2246
2247
2248
2249
2250
2251
2252
2253
2254
2255
2256
2257
2258
2259
2260
2261 func createBindingDecl(logf logger, caller *Caller, args []*argument, calleeDecl *ast.FuncDecl, results []*paramInfo) *bindingDeclInfo {
2262
2263
2264
2265
2266
2267
2268
2269
2270
2271
2272
2273
2274 if lastArg := last(args); lastArg != nil && lastArg.spread {
2275 logf("binding decls not yet supported for spread calls")
2276 return nil
2277 }
2278
2279 var (
2280 specs []ast.Spec
2281 names = make(map[string]bool)
2282 )
2283
2284
2285
2286
2287 shadow := func(spec *ast.ValueSpec) bool {
2288
2289
2290
2291
2292
2293 const includeComplitIdents = true
2294 free := free.Names(spec.Type, includeComplitIdents)
2295 for _, value := range spec.Values {
2296 for name := range freeVars(caller.Info, value) {
2297 free[name] = true
2298 }
2299 }
2300 for name := range free {
2301 if names[name] {
2302 logf("binding decl would shadow free name %q", name)
2303 return true
2304 }
2305 }
2306 for _, id := range spec.Names {
2307 if id.Name != "_" {
2308 names[id.Name] = true
2309 }
2310 }
2311 return false
2312 }
2313
2314
2315
2316
2317
2318
2319 var values []ast.Expr
2320 for _, arg := range args {
2321 if arg != nil {
2322 values = append(values, arg.expr)
2323 }
2324 }
2325 for _, field := range calleeDecl.Type.Params.List {
2326
2327 spec := &ast.ValueSpec{
2328 Names: cleanNodes(field.Names),
2329 Type: cleanNode(field.Type),
2330 Values: values[:len(field.Names)],
2331 }
2332 values = values[len(field.Names):]
2333 if shadow(spec) {
2334 return nil
2335 }
2336 specs = append(specs, spec)
2337 }
2338 assert(len(values) == 0, "args/params mismatch")
2339
2340
2341
2342
2343
2344 if calleeDecl.Type.Results != nil {
2345 resultIdx := 0
2346 for _, field := range calleeDecl.Type.Results.List {
2347 if field.Names == nil {
2348 resultIdx++
2349 continue
2350 }
2351 var names []*ast.Ident
2352 for _, id := range field.Names {
2353 if len(results[resultIdx].Refs) > 0 {
2354 names = append(names, id)
2355 }
2356 resultIdx++
2357 }
2358 if len(names) > 0 {
2359 spec := &ast.ValueSpec{
2360 Names: cleanNodes(names),
2361 Type: cleanNode(field.Type),
2362 }
2363 if shadow(spec) {
2364 return nil
2365 }
2366 specs = append(specs, spec)
2367 }
2368 }
2369 }
2370
2371 if len(specs) == 0 {
2372 logf("binding decl not needed: all parameters substituted")
2373 return nil
2374 }
2375
2376 stmt := &ast.DeclStmt{
2377 Decl: &ast.GenDecl{
2378 Tok: token.VAR,
2379 Specs: specs,
2380 },
2381 }
2382 logf("binding decl: %s", debugFormatNode(caller.Fset, stmt))
2383 return &bindingDeclInfo{names: names, stmt: stmt}
2384 }
2385
2386
2387 func (caller *Caller) lookup(name string) types.Object {
2388 pos := caller.Call.Pos()
2389 for _, n := range caller.path {
2390 if scope := scopeFor(caller.Info, n); scope != nil {
2391 if _, obj := scope.LookupParent(name, pos); obj != nil {
2392 return obj
2393 }
2394 }
2395 }
2396 return nil
2397 }
2398
2399 func scopeFor(info *types.Info, n ast.Node) *types.Scope {
2400
2401
2402 switch fn := n.(type) {
2403 case *ast.FuncDecl:
2404 n = fn.Type
2405 case *ast.FuncLit:
2406 n = fn.Type
2407 }
2408 return info.Scopes[n]
2409 }
2410
2411
2412
2413
2414
2415
2416 func freeVars(info *types.Info, e ast.Expr) map[string]bool {
2417 free := make(map[string]bool)
2418 ast.Inspect(e, func(n ast.Node) bool {
2419 if id, ok := n.(*ast.Ident); ok {
2420
2421 if obj, ok := info.Uses[id]; ok && !within(obj.Pos(), e) && !isField(obj) {
2422 free[obj.Name()] = true
2423 }
2424 }
2425 return true
2426 })
2427 return free
2428 }
2429
2430
2431
2432
2433 func (st *state) effects(info *types.Info, expr ast.Expr) bool {
2434 effects := false
2435 ast.Inspect(expr, func(n ast.Node) bool {
2436 switch n := n.(type) {
2437 case *ast.FuncLit:
2438 return false
2439
2440 case *ast.CallExpr:
2441 if info.Types[n.Fun].IsType() {
2442
2443 } else if !typesinternal.CallsPureBuiltin(info, n) {
2444
2445
2446
2447
2448
2449
2450
2451 effects = true
2452 }
2453
2454 case *ast.UnaryExpr:
2455 if n.Op == token.ARROW {
2456 effects = true
2457 }
2458 }
2459 return true
2460 })
2461
2462
2463
2464 if st.opts.IgnoreEffects && effects {
2465 effects = false
2466 st.opts.Logf("ignoring potential effects of argument %s",
2467 debugFormatNode(st.caller.Fset, expr))
2468 }
2469
2470 return effects
2471 }
2472
2473
2474
2475
2476
2477
2478
2479
2480
2481
2482
2483
2484
2485
2486
2487
2488
2489
2490
2491
2492 func pure(info *types.Info, assign1 func(*types.Var) bool, e ast.Expr) bool {
2493 var pure func(e ast.Expr) bool
2494 pure = func(e ast.Expr) bool {
2495 switch e := e.(type) {
2496 case *ast.ParenExpr:
2497 return pure(e.X)
2498
2499 case *ast.Ident:
2500 if v, ok := info.Uses[e].(*types.Var); ok {
2501
2502
2503
2504
2505
2506
2507
2508
2509
2510 return !isPkgLevel(v) && assign1(v)
2511 }
2512
2513
2514 return true
2515
2516 case *ast.FuncLit:
2517
2518
2519
2520
2521 return true
2522
2523 case *ast.BasicLit:
2524 return true
2525
2526 case *ast.UnaryExpr:
2527 return e.Op != token.ARROW && pure(e.X)
2528
2529 case *ast.BinaryExpr:
2530 return pure(e.X) && pure(e.Y)
2531
2532 case *ast.CallExpr:
2533
2534 if info.Types[e.Fun].IsType() {
2535 return pure(e.Args[0])
2536 }
2537
2538
2539 if typesinternal.CallsPureBuiltin(info, e) {
2540 for _, arg := range e.Args {
2541 if !pure(arg) {
2542 return false
2543 }
2544 }
2545 return true
2546 }
2547
2548
2549
2550
2551
2552
2553
2554
2555
2556
2557 return false
2558
2559 case *ast.CompositeLit:
2560
2561 for _, elt := range e.Elts {
2562 if kv, ok := elt.(*ast.KeyValueExpr); ok {
2563 if !pure(kv.Value) {
2564 return false
2565 }
2566 if id, ok := kv.Key.(*ast.Ident); ok {
2567 if v, ok := info.Uses[id].(*types.Var); ok && v.IsField() {
2568 continue
2569 }
2570 }
2571
2572 if !pure(kv.Key) {
2573 return false
2574 }
2575
2576 } else if !pure(elt) {
2577 return false
2578 }
2579 }
2580 return true
2581
2582 case *ast.SelectorExpr:
2583 if seln, ok := info.Selections[e]; ok {
2584
2585 switch seln.Kind() {
2586 case types.MethodExpr:
2587
2588
2589 return true
2590
2591 case types.MethodVal, types.FieldVal:
2592
2593
2594
2595 return !indirectSelection(seln) && pure(e.X)
2596
2597 default:
2598 panic(seln)
2599 }
2600 } else {
2601
2602
2603 return pure(e.Sel)
2604 }
2605
2606 case *ast.StarExpr:
2607 return false
2608
2609 default:
2610 return false
2611 }
2612 }
2613 return pure(e)
2614 }
2615
2616
2617
2618
2619
2620
2621
2622
2623
2624
2625
2626
2627
2628
2629
2630 func duplicable(info *types.Info, e ast.Expr) bool {
2631 switch e := e.(type) {
2632 case *ast.ParenExpr:
2633 return duplicable(info, e.X)
2634
2635 case *ast.Ident:
2636 return true
2637
2638 case *ast.BasicLit:
2639 v := info.Types[e].Value
2640 switch e.Kind {
2641 case token.INT:
2642 return true
2643 case token.STRING:
2644 return consteq(v, kZeroString)
2645 case token.FLOAT:
2646 return consteq(v, kZeroFloat) || consteq(v, kOneFloat)
2647 }
2648
2649 case *ast.UnaryExpr:
2650 return (e.Op == token.ADD || e.Op == token.SUB) && duplicable(info, e.X)
2651
2652 case *ast.CompositeLit:
2653
2654
2655
2656 if len(e.Elts) == 0 {
2657 switch info.TypeOf(e).Underlying().(type) {
2658 case *types.Struct, *types.Array:
2659 return true
2660 }
2661 }
2662 return false
2663
2664 case *ast.CallExpr:
2665
2666
2667
2668
2669
2670
2671
2672
2673
2674 if !info.Types[e.Fun].IsType() {
2675 return false
2676 }
2677
2678 fun := info.TypeOf(e.Fun)
2679 arg := info.TypeOf(e.Args[0])
2680
2681 switch fun := fun.Underlying().(type) {
2682 case *types.Slice:
2683
2684 elem, ok := fun.Elem().Underlying().(*types.Basic)
2685 if ok && (elem.Kind() == types.Rune || elem.Kind() == types.Byte) {
2686 from, ok := arg.Underlying().(*types.Basic)
2687 isString := ok && from.Info()&types.IsString != 0
2688 return !isString
2689 }
2690 case *types.TypeParam:
2691 return false
2692 }
2693 return true
2694
2695 case *ast.SelectorExpr:
2696 if seln, ok := info.Selections[e]; ok {
2697
2698
2699 return !indirectSelection(seln)
2700 }
2701
2702 return true
2703 }
2704 return false
2705 }
2706
2707 func consteq(x, y constant.Value) bool {
2708 return constant.Compare(x, token.EQL, y)
2709 }
2710
2711 var (
2712 kZeroInt = constant.MakeInt64(0)
2713 kZeroString = constant.MakeString("")
2714 kZeroFloat = constant.MakeFloat64(0.0)
2715 kOneFloat = constant.MakeFloat64(1.0)
2716 )
2717
2718
2719
2720 func assert(cond bool, msg string) {
2721 if !cond {
2722 panic(msg)
2723 }
2724 }
2725
2726
2727 func blanks[E ast.Expr](n int) []E {
2728 if n == 0 {
2729 panic("blanks(0)")
2730 }
2731 res := make([]E, n)
2732 for i := range res {
2733 res[i] = ast.Expr(makeIdent("_")).(E)
2734 }
2735 return res
2736 }
2737
2738 func makeIdent(name string) *ast.Ident {
2739 return &ast.Ident{Name: name}
2740 }
2741
2742
2743
2744 func importedPkgName(info *types.Info, imp *ast.ImportSpec) (*types.PkgName, bool) {
2745 var obj types.Object
2746 if imp.Name != nil {
2747 obj = info.Defs[imp.Name]
2748 } else {
2749 obj = info.Implicits[imp]
2750 }
2751 pkgname, ok := obj.(*types.PkgName)
2752 return pkgname, ok
2753 }
2754
2755 func isPkgLevel(obj types.Object) bool {
2756
2757
2758
2759 return obj.Pkg().Scope().Lookup(obj.Name()) == obj
2760 }
2761
2762
2763
2764 func callContext(callPath []ast.Node) (parent, grandparent ast.Node) {
2765 _ = callPath[0].(*ast.CallExpr)
2766 for _, n := range callPath[1:] {
2767 if !is[*ast.ParenExpr](n) {
2768 if parent == nil {
2769 parent = n
2770 } else {
2771 return parent, n
2772 }
2773 }
2774 }
2775 return parent, nil
2776 }
2777
2778
2779
2780
2781 func hasLabelConflict(callPath []ast.Node, calleeLabels []string) bool {
2782 labels := callerLabels(callPath)
2783 for _, label := range calleeLabels {
2784 if labels[label] {
2785 return true
2786 }
2787 }
2788 return false
2789 }
2790
2791
2792
2793 func callerLabels(callPath []ast.Node) map[string]bool {
2794 var callerBody *ast.BlockStmt
2795 switch f := callerFunc(callPath).(type) {
2796 case *ast.FuncDecl:
2797 callerBody = f.Body
2798 case *ast.FuncLit:
2799 callerBody = f.Body
2800 }
2801 var labels map[string]bool
2802 if callerBody != nil {
2803 ast.Inspect(callerBody, func(n ast.Node) bool {
2804 switch n := n.(type) {
2805 case *ast.FuncLit:
2806 return false
2807 case *ast.LabeledStmt:
2808 if labels == nil {
2809 labels = make(map[string]bool)
2810 }
2811 labels[n.Label.Name] = true
2812 }
2813 return true
2814 })
2815 }
2816 return labels
2817 }
2818
2819
2820
2821 func callerFunc(callPath []ast.Node) ast.Node {
2822 _ = callPath[0].(*ast.CallExpr)
2823 for _, n := range callPath[1:] {
2824 if is[*ast.FuncDecl](n) || is[*ast.FuncLit](n) {
2825 return n
2826 }
2827 }
2828 return nil
2829 }
2830
2831
2832
2833
2834
2835
2836
2837
2838 func callStmt(callPath []ast.Node, unrestricted bool) *ast.ExprStmt {
2839 parent, _ := callContext(callPath)
2840 stmt, ok := parent.(*ast.ExprStmt)
2841 if ok && unrestricted {
2842 switch callPath[slices.Index(callPath, ast.Node(stmt))+1].(type) {
2843 case *ast.LabeledStmt,
2844 *ast.BlockStmt,
2845 *ast.CaseClause,
2846 *ast.CommClause:
2847
2848 default:
2849
2850
2851
2852
2853
2854 return nil
2855 }
2856 }
2857 return stmt
2858 }
2859
2860
2861
2862
2863
2864
2865
2866
2867
2868
2869
2870
2871
2872
2873
2874
2875
2876
2877
2878
2879
2880
2881
2882
2883
2884
2885
2886
2887
2888
2889
2890
2891
2892
2893
2894
2895
2896
2897
2898
2899
2900
2901
2902
2903 func replaceNode(root ast.Node, from, to ast.Node) {
2904 if from == nil {
2905 panic("from == nil")
2906 }
2907 if reflect.ValueOf(from).IsNil() {
2908 panic(fmt.Sprintf("from == (%T)(nil)", from))
2909 }
2910 if from == root {
2911 panic("from == root")
2912 }
2913 found := false
2914 var parent reflect.Value
2915 var visit func(reflect.Value)
2916 visit = func(v reflect.Value) {
2917 switch v.Kind() {
2918 case reflect.Pointer:
2919 if v.Interface() == from {
2920 found = true
2921
2922
2923
2924
2925
2926
2927
2928
2929
2930 if !v.CanAddr() {
2931 v = parent
2932 }
2933
2934
2935 var toV reflect.Value
2936 if to != nil {
2937 toV = reflect.ValueOf(to)
2938 } else {
2939 toV = reflect.Zero(v.Type())
2940 }
2941 v.Set(toV)
2942
2943 } else if !v.IsNil() {
2944 switch v.Interface().(type) {
2945 case *ast.Object, *ast.Scope:
2946
2947 default:
2948 visit(v.Elem())
2949 }
2950 }
2951
2952 case reflect.Struct:
2953 for i := range v.Type().NumField() {
2954 visit(v.Field(i))
2955 }
2956
2957 case reflect.Slice:
2958 compact := false
2959 for i := range v.Len() {
2960 visit(v.Index(i))
2961 if v.Index(i).IsNil() {
2962 compact = true
2963 }
2964 }
2965 if compact {
2966
2967
2968
2969 j := 0
2970 for i := range v.Len() {
2971 if !v.Index(i).IsNil() {
2972 v.Index(j).Set(v.Index(i))
2973 j++
2974 }
2975 }
2976 v.SetLen(j)
2977 }
2978 case reflect.Interface:
2979 parent = v
2980 visit(v.Elem())
2981
2982 case reflect.Array, reflect.Chan, reflect.Func, reflect.Map, reflect.UnsafePointer:
2983 panic(v)
2984 default:
2985
2986 }
2987 parent = reflect.Value{}
2988 }
2989 visit(reflect.ValueOf(root))
2990 if !found {
2991 panic(fmt.Sprintf("%T not found", from))
2992 }
2993 }
2994
2995
2996
2997
2998
2999 func cleanNode[T ast.Node](node T) T {
3000 clone := internalastutil.CloneNode(node)
3001 clearPositions(clone)
3002 return clone
3003 }
3004
3005 func cleanNodes[T ast.Node](nodes []T) []T {
3006 var clean []T
3007 for _, node := range nodes {
3008 clean = append(clean, cleanNode(node))
3009 }
3010 return clean
3011 }
3012
3013
3014
3015
3016
3017
3018
3019
3020
3021
3022
3023
3024 func clearPositions(root ast.Node) {
3025 posType := reflect.TypeFor[token.Pos]()
3026 ast.Inspect(root, func(n ast.Node) bool {
3027 if n != nil {
3028 v := reflect.ValueOf(n).Elem()
3029 fields := v.Type().NumField()
3030 for i := range fields {
3031 f := v.Field(i)
3032
3033
3034
3035
3036
3037
3038
3039
3040
3041 if f.Type() == posType {
3042 if f.Interface() != token.NoPos {
3043 f.Set(reflect.ValueOf(token.Pos(1)))
3044 }
3045 }
3046 }
3047 }
3048 return true
3049 })
3050 }
3051
3052
3053
3054
3055
3056 func findIdent(root ast.Node, pos token.Pos) ([]ast.Node, *ast.Ident) {
3057
3058 var (
3059 path []ast.Node
3060 found *ast.Ident
3061 )
3062 ast.Inspect(root, func(n ast.Node) bool {
3063 if found != nil {
3064 return false
3065 }
3066 if n == nil {
3067 path = path[:len(path)-1]
3068 return false
3069 }
3070 if id, ok := n.(*ast.Ident); ok {
3071 if id.Pos() == pos {
3072 found = id
3073 return true
3074 }
3075 }
3076 path = append(path, n)
3077 return true
3078 })
3079 if found == nil {
3080 panic(fmt.Sprintf("findIdent %d not found in %s",
3081 pos, debugFormatNode(token.NewFileSet(), root)))
3082 }
3083 return path, found
3084 }
3085
3086 func prepend[T any](elem T, slice ...T) []T {
3087 return append([]T{elem}, slice...)
3088 }
3089
3090
3091
3092 func debugFormatNode(fset *token.FileSet, n ast.Node) string {
3093 var out strings.Builder
3094 if err := format.Node(&out, fset, n); err != nil {
3095 out.WriteString(err.Error())
3096 }
3097 return out.String()
3098 }
3099
3100 func shallowCopy[T any](ptr *T) *T {
3101 copy := *ptr
3102 return ©
3103 }
3104
3105
3106 func forall[T any](list []T, f func(i int, x T) bool) bool {
3107 for i, x := range list {
3108 if !f(i, x) {
3109 return false
3110 }
3111 }
3112 return true
3113 }
3114
3115
3116 func exists[T any](list []T, f func(i int, x T) bool) bool {
3117 for i, x := range list {
3118 if f(i, x) {
3119 return true
3120 }
3121 }
3122 return false
3123 }
3124
3125
3126 func last[T any](slice []T) T {
3127 n := len(slice)
3128 if n > 0 {
3129 return slice[n-1]
3130 }
3131 return *new(T)
3132 }
3133
3134
3135
3136
3137 func declares(stmts []ast.Stmt) map[string]bool {
3138 names := make(map[string]bool)
3139 for _, stmt := range stmts {
3140 switch stmt := stmt.(type) {
3141 case *ast.DeclStmt:
3142 for _, spec := range stmt.Decl.(*ast.GenDecl).Specs {
3143 switch spec := spec.(type) {
3144 case *ast.ValueSpec:
3145 for _, id := range spec.Names {
3146 names[id.Name] = true
3147 }
3148 case *ast.TypeSpec:
3149 names[spec.Name.Name] = true
3150 }
3151 }
3152
3153 case *ast.AssignStmt:
3154 if stmt.Tok == token.DEFINE {
3155 for _, lhs := range stmt.Lhs {
3156 names[lhs.(*ast.Ident).Name] = true
3157 }
3158 }
3159 }
3160 }
3161 delete(names, "_")
3162 return names
3163 }
3164
3165
3166
3167
3168
3169
3170
3171
3172 type importNameFunc = func(pkgPath string, shadow shadowMap) string
3173
3174
3175
3176
3177
3178
3179
3180
3181
3182
3183
3184
3185
3186
3187
3188
3189
3190
3191
3192
3193
3194
3195
3196
3197
3198
3199
3200
3201
3202
3203
3204
3205
3206
3207
3208
3209
3210
3211
3212
3213
3214
3215
3216 func (st *state) assignStmts(callerStmt *ast.AssignStmt, returnOperands []ast.Expr, importName importNameFunc) ([]ast.Stmt, bool) {
3217 logf, caller, callee := st.opts.Logf, st.caller, &st.callee.impl
3218
3219 assert(len(callee.Returns) == 1, "unexpected multiple returns")
3220 resultInfo := callee.Returns[0]
3221
3222
3223
3224
3225
3226
3227
3228
3229
3230
3231
3232 var (
3233 lhs []ast.Expr
3234 defs = make([]*ast.Ident, len(callerStmt.Lhs))
3235 blanks = make([]bool, len(callerStmt.Lhs))
3236 byType typeutil.Map
3237 )
3238 for i, expr := range callerStmt.Lhs {
3239 lhs = append(lhs, expr)
3240 if name, ok := expr.(*ast.Ident); ok {
3241 if name.Name == "_" {
3242 blanks[i] = true
3243 continue
3244 }
3245
3246 if obj, isDef := caller.Info.Defs[name]; isDef {
3247 defs[i] = name
3248 typ := obj.Type()
3249 idxs, _ := byType.At(typ).([]int)
3250 idxs = append(idxs, i)
3251 byType.Set(typ, idxs)
3252 }
3253 }
3254 }
3255
3256
3257
3258
3259
3260 var (
3261 rhs []ast.Expr
3262 callIdx = -1
3263 nilBlankAssigns = make(map[int]unit)
3264 freeNames = make(map[string]bool)
3265 nonTrivial = make(map[int]bool)
3266 )
3267 const includeComplitIdents = true
3268
3269 for i, expr := range callerStmt.Rhs {
3270 if expr == caller.Call {
3271 assert(callIdx == -1, "malformed (duplicative) AST")
3272 callIdx = i
3273 for j, returnOperand := range returnOperands {
3274 maps.Copy(freeNames, free.Names(returnOperand, includeComplitIdents))
3275 rhs = append(rhs, returnOperand)
3276 if resultInfo[j]&nonTrivialResult != 0 {
3277 nonTrivial[i+j] = true
3278 }
3279 if blanks[i+j] && resultInfo[j]&untypedNilResult != 0 {
3280 nilBlankAssigns[i+j] = unit{}
3281 }
3282 }
3283 } else {
3284
3285 expr = internalastutil.CloneNode(expr)
3286 clearPositions(expr)
3287 maps.Copy(freeNames, free.Names(expr, includeComplitIdents))
3288 rhs = append(rhs, expr)
3289 }
3290 }
3291 assert(callIdx >= 0, "failed to find call in RHS")
3292
3293
3294
3295
3296
3297
3298
3299
3300
3301
3302
3303
3304
3305
3306 if len(nonTrivial) == 0 {
3307
3308 logf("substrategy: splice assignment")
3309 return []ast.Stmt{&ast.AssignStmt{
3310 Lhs: lhs,
3311 Tok: callerStmt.Tok,
3312 TokPos: callerStmt.TokPos,
3313 Rhs: rhs,
3314 }}, true
3315 }
3316
3317
3318
3319
3320
3321
3322
3323
3324
3325
3326 universeAny := types.Universe.Lookup("any")
3327 typeExpr := func(typ types.Type, shadow shadowMap) ast.Expr {
3328 var (
3329 typeName string
3330 obj *types.TypeName
3331 )
3332 if tname := typesinternal.TypeNameFor(typ); tname != nil {
3333 obj = tname
3334 typeName = tname.Name()
3335 }
3336
3337
3338
3339 if typ == universeAny.Type() {
3340 typeName = "any"
3341 }
3342
3343 if typeName == "" {
3344 return nil
3345 }
3346
3347 if obj == nil || obj.Pkg() == nil || obj.Pkg() == caller.Types {
3348 if shadow[typeName] != 0 {
3349 logf("cannot write shadowed type name %q", typeName)
3350 return nil
3351 }
3352 obj, _ := caller.lookup(typeName).(*types.TypeName)
3353 if obj != nil && types.Identical(obj.Type(), typ) {
3354 return ast.NewIdent(typeName)
3355 }
3356 } else if pkgName := importName(obj.Pkg().Path(), shadow); pkgName != "" {
3357 return &ast.SelectorExpr{
3358 X: ast.NewIdent(pkgName),
3359 Sel: ast.NewIdent(typeName),
3360 }
3361 }
3362 return nil
3363 }
3364
3365
3366
3367
3368
3369
3370
3371
3372
3373
3374
3375
3376
3377
3378
3379 if len(rhs) != len(lhs) {
3380 assert(len(rhs) == 1 && len(returnOperands) == 1, "expected spread call")
3381
3382 for _, id := range defs {
3383 if id != nil && freeNames[id.Name] {
3384
3385
3386 return nil, false
3387 }
3388 }
3389
3390
3391
3392 var (
3393 specs []ast.Spec
3394 specIdxs []int
3395 shadow = make(shadowMap)
3396 )
3397 failed := false
3398 byType.Iterate(func(typ types.Type, v any) {
3399 if failed {
3400 return
3401 }
3402 idxs := v.([]int)
3403 specIdxs = append(specIdxs, idxs[0])
3404 texpr := typeExpr(typ, shadow)
3405 if texpr == nil {
3406 failed = true
3407 return
3408 }
3409 spec := &ast.ValueSpec{
3410 Type: texpr,
3411 }
3412 for _, idx := range idxs {
3413 spec.Names = append(spec.Names, ast.NewIdent(defs[idx].Name))
3414 }
3415 specs = append(specs, spec)
3416 })
3417 if failed {
3418 return nil, false
3419 }
3420 logf("substrategy: spread assignment")
3421 return []ast.Stmt{
3422 &ast.DeclStmt{
3423 Decl: &ast.GenDecl{
3424 Tok: token.VAR,
3425 Specs: specs,
3426 },
3427 },
3428 &ast.AssignStmt{
3429 Lhs: callerStmt.Lhs,
3430 Tok: token.ASSIGN,
3431 Rhs: returnOperands,
3432 },
3433 }, true
3434 }
3435
3436 assert(len(lhs) == len(rhs), "mismatching LHS and RHS")
3437
3438
3439
3440
3441
3442
3443
3444
3445
3446
3447
3448
3449
3450
3451 var origIdxs []int
3452 i := 0
3453 for j := range lhs {
3454 if _, ok := nilBlankAssigns[j]; !ok {
3455 lhs[i] = lhs[j]
3456 rhs[i] = rhs[j]
3457 origIdxs = append(origIdxs, j)
3458 i++
3459 }
3460 }
3461 lhs = lhs[:i]
3462 rhs = rhs[:i]
3463
3464 if len(lhs) == 0 {
3465 logf("trivial assignment after pruning nil blanks assigns")
3466
3467
3468 return nil, true
3469 }
3470
3471
3472
3473
3474
3475 for i, expr := range rhs {
3476 idx := origIdxs[i]
3477 if nonTrivial[idx] && defs[idx] != nil {
3478 typ := caller.Info.TypeOf(lhs[i])
3479 texpr := typeExpr(typ, nil)
3480 if texpr == nil {
3481 return nil, false
3482 }
3483 if _, ok := texpr.(*ast.StarExpr); ok {
3484
3485 texpr = &ast.ParenExpr{X: texpr}
3486 }
3487 rhs[i] = &ast.CallExpr{
3488 Fun: texpr,
3489 Args: []ast.Expr{expr},
3490 }
3491 }
3492 }
3493 logf("substrategy: convert assignment")
3494 return []ast.Stmt{&ast.AssignStmt{
3495 Lhs: lhs,
3496 Tok: callerStmt.Tok,
3497 Rhs: rhs,
3498 }}, true
3499 }
3500
3501
3502
3503 func tailCallSafeReturn(caller *Caller, calleeSymbol *types.Func, callee *gobCallee) bool {
3504
3505 if !hasNonTrivialReturn(callee.Returns) {
3506 return true
3507 }
3508
3509 var callerType types.Type
3510
3511
3512 loop:
3513 for _, n := range caller.path {
3514 switch f := n.(type) {
3515 case *ast.FuncDecl:
3516 callerType = caller.Info.ObjectOf(f.Name).Type()
3517 break loop
3518 case *ast.FuncLit:
3519 callerType = caller.Info.TypeOf(f)
3520 break loop
3521 }
3522 }
3523
3524
3525
3526
3527 callerResults := callerType.(*types.Signature).Results()
3528 calleeResults := calleeSymbol.Type().(*types.Signature).Results()
3529 return types.Identical(callerResults, calleeResults)
3530 }
3531
3532
3533
3534 func hasNonTrivialReturn(returnInfo [][]returnOperandFlags) bool {
3535 for _, resultInfo := range returnInfo {
3536 for _, r := range resultInfo {
3537 if r&nonTrivialResult != 0 {
3538 return true
3539 }
3540 }
3541 }
3542 return false
3543 }
3544
3545 type unit struct{}
3546
View as plain text