1
2
3
4
5 package inline
6
7
8
9 import (
10 "bytes"
11 "encoding/gob"
12 "fmt"
13 "go/ast"
14 "go/parser"
15 "go/token"
16 "go/types"
17 "slices"
18 "strings"
19
20 "golang.org/x/tools/go/types/typeutil"
21 "golang.org/x/tools/internal/typeparams"
22 "golang.org/x/tools/internal/typesinternal"
23 )
24
25
26 type Callee struct {
27 impl gobCallee
28 }
29
30 func (callee *Callee) String() string { return callee.impl.Name }
31
32 type gobCallee struct {
33 Content []byte
34
35
36 PkgPath string
37 Name string
38 GoVersion string
39 Unexported []string
40 FreeRefs []freeRef
41 FreeObjs []object
42 ValidForCallStmt bool
43 NumResults int
44 Params []*paramInfo
45 TypeParams []*paramInfo
46 Results []*paramInfo
47 Effects []int
48 HasDefer bool
49 HasBareReturn bool
50 Returns [][]returnOperandFlags
51 Labels []string
52 Falcon falconResult
53 }
54
55
56
57 type returnOperandFlags int
58
59 const (
60 nonTrivialResult returnOperandFlags = 1 << iota
61 untypedNilResult
62 )
63
64
65
66 type freeRef struct {
67 Offset int
68 Object int
69 }
70
71
72 type object struct {
73 Name string
74 Kind string
75 PkgPath string
76 PkgName string
77
78
79 ValidPos bool
80 Shadow shadowMap
81 }
82
83
84
85
86
87
88
89
90
91
92
93
94
95 func AnalyzeCallee(logf func(string, ...any), fset *token.FileSet, pkg *types.Package, info *types.Info, decl *ast.FuncDecl, content []byte) (*Callee, error) {
96 checkInfoFields(info)
97
98
99
100 fn := info.Defs[decl.Name].(*types.Func)
101 sig := fn.Type().(*types.Signature)
102
103 logf("analyzeCallee %v @ %v", fn, fset.PositionFor(decl.Pos(), false))
104
105
106 var name string
107 if sig.Recv() == nil {
108 name = fmt.Sprintf("%s.%s", fn.Pkg().Name(), fn.Name())
109 } else {
110 name = fmt.Sprintf("(%s).%s", types.TypeString(sig.Recv().Type(), (*types.Package).Name), fn.Name())
111 }
112
113 if decl.Body == nil {
114 return nil, fmt.Errorf("cannot inline function %s as it has no body", name)
115 }
116
117
118
119
120
121
122
123
124
125
126
127 var goVersion string
128 for file, v := range info.FileVersions {
129 if file.Pos() < decl.Pos() && decl.Pos() < file.End() {
130 goVersion = v
131 break
132 }
133 }
134
135
136
137 var (
138 fieldObjs = fieldObjs(sig)
139 freeObjIndex = make(map[types.Object]int)
140 freeObjs []object
141 freeRefs []freeRef
142 unexported []string
143 )
144 var f func(n ast.Node, stack []ast.Node) bool
145 var stack []ast.Node
146 stack = append(stack, decl.Type)
147 visit := func(n ast.Node, stack []ast.Node) { ast.PreorderStack(n, stack, f) }
148 f = func(n ast.Node, stack []ast.Node) bool {
149 switch n := n.(type) {
150 case *ast.SelectorExpr:
151
152 if sel, ok := info.Selections[n]; ok &&
153 !within(sel.Obj().Pos(), decl) &&
154 !n.Sel.IsExported() {
155 sym := fmt.Sprintf("(%s).%s", info.TypeOf(n.X), n.Sel.Name)
156 unexported = append(unexported, sym)
157 }
158
159
160 visit(n.X, stack)
161 return false
162
163 case *ast.CompositeLit:
164
165
166 litType := typeparams.Deref(info.TypeOf(n))
167 if s, ok := typeparams.CoreType(litType).(*types.Struct); ok {
168 if n.Type != nil {
169 visit(n.Type, stack)
170 }
171 for i, elt := range n.Elts {
172 var field *types.Var
173 var value ast.Expr
174 if kv, ok := elt.(*ast.KeyValueExpr); ok {
175 field = info.Uses[kv.Key.(*ast.Ident)].(*types.Var)
176 value = kv.Value
177 } else {
178 field = s.Field(i)
179 value = elt
180 }
181 if !within(field.Pos(), decl) && !field.Exported() {
182 sym := fmt.Sprintf("(%s).%s", litType, field.Name())
183 unexported = append(unexported, sym)
184 }
185
186
187 visit(value, stack)
188 }
189 return false
190 }
191
192 case *ast.Ident:
193 if obj, ok := info.Uses[n]; ok {
194
195 if isField(obj) || isMethod(obj) {
196 panic(obj)
197 }
198
199
200
201
202 if !n.IsExported() &&
203 obj.Pkg() != nil && obj.Parent() == obj.Pkg().Scope() {
204 unexported = append(unexported, n.Name)
205 }
206
207
208 if obj == fn || !within(obj.Pos(), decl) {
209 objidx, ok := freeObjIndex[obj]
210 if !ok {
211 objidx = len(freeObjIndex)
212 var pkgPath, pkgName string
213 if pn, ok := obj.(*types.PkgName); ok {
214 pkgPath = pn.Imported().Path()
215 pkgName = pn.Imported().Name()
216 } else if obj.Pkg() != nil {
217 pkgPath = obj.Pkg().Path()
218 pkgName = obj.Pkg().Name()
219 }
220 freeObjs = append(freeObjs, object{
221 Name: obj.Name(),
222 Kind: objectKind(obj),
223 PkgName: pkgName,
224 PkgPath: pkgPath,
225 ValidPos: obj.Pos().IsValid(),
226 })
227 freeObjIndex[obj] = objidx
228 }
229
230 freeObjs[objidx].Shadow = freeObjs[objidx].Shadow.add(info, fieldObjs, obj.Name(), stack)
231
232 freeRefs = append(freeRefs, freeRef{
233 Offset: int(n.Pos() - decl.Pos()),
234 Object: objidx,
235 })
236 }
237 }
238 }
239 return true
240 }
241 visit(decl, stack)
242
243
244
245
246 validForCallStmt := false
247 if len(decl.Body.List) != 1 {
248
249 } else if ret, ok := decl.Body.List[0].(*ast.ReturnStmt); ok && len(ret.Results) == 1 {
250 validForCallStmt = func() bool {
251 switch expr := ast.Unparen(ret.Results[0]).(type) {
252 case *ast.CallExpr:
253 callee := typeutil.Callee(info, expr)
254 if callee == nil {
255 return false
256 }
257
258
259
260
261
262 if builtin, ok := callee.(*types.Builtin); ok {
263 return builtin.Name() == "copy" ||
264 builtin.Name() == "recover"
265 }
266
267 return true
268
269 case *ast.UnaryExpr:
270 return expr.Op == token.ARROW
271 }
272
273
274 return false
275 }()
276 }
277
278
279
280 var (
281 hasDefer = false
282 hasBareReturn = false
283 returnInfo [][]returnOperandFlags
284 labels []string
285 )
286 ast.Inspect(decl.Body, func(n ast.Node) bool {
287 switch n := n.(type) {
288 case *ast.FuncLit:
289 return false
290 case *ast.DeferStmt:
291 hasDefer = true
292 case *ast.LabeledStmt:
293 labels = append(labels, n.Label.Name)
294 case *ast.ReturnStmt:
295
296
297
298 var resultInfo []returnOperandFlags
299 if len(n.Results) > 0 {
300 argInfo := func(i int) (ast.Expr, types.Type) {
301 expr := n.Results[i]
302 return expr, info.TypeOf(expr)
303 }
304 if len(n.Results) == 1 && sig.Results().Len() > 1 {
305
306 tuple := info.TypeOf(n.Results[0]).(*types.Tuple)
307 argInfo = func(i int) (ast.Expr, types.Type) {
308 return nil, tuple.At(i).Type()
309 }
310 }
311 for i := range sig.Results().Len() {
312 expr, typ := argInfo(i)
313 var flags returnOperandFlags
314 if typ == types.Typ[types.UntypedNil] {
315 flags |= untypedNilResult
316 }
317 if !trivialConversion(info.Types[expr].Value, typ, sig.Results().At(i).Type()) {
318 flags |= nonTrivialResult
319 }
320 resultInfo = append(resultInfo, flags)
321 }
322 } else if sig.Results().Len() > 0 {
323 hasBareReturn = true
324 }
325 returnInfo = append(returnInfo, resultInfo)
326 }
327 return true
328 })
329
330
331 for _, obj := range freeObjs {
332
333
334 if strings.HasPrefix(obj.Name, "_Cfunc_") ||
335 strings.HasPrefix(obj.Name, "_Ctype_") ||
336 strings.HasPrefix(obj.Name, "_Cvar_") {
337 return nil, fmt.Errorf("cannot inline cgo-generated functions")
338 }
339 }
340
341
342
343
344
345
346
347
348
349
350 content = append([]byte("package _\n"),
351 content[offsetOf(fset, decl.Pos()):offsetOf(fset, decl.End())]...)
352
353 if _, _, err := parseCompact(content); err != nil {
354 return nil, err
355 }
356
357 params, results, effects, falcon := analyzeParams(logf, fset, info, decl)
358 tparams := analyzeTypeParams(logf, fset, info, decl)
359 return &Callee{gobCallee{
360 Content: content,
361 PkgPath: pkg.Path(),
362 Name: name,
363 GoVersion: goVersion,
364 Unexported: unexported,
365 FreeObjs: freeObjs,
366 FreeRefs: freeRefs,
367 ValidForCallStmt: validForCallStmt,
368 NumResults: sig.Results().Len(),
369 Params: params,
370 TypeParams: tparams,
371 Results: results,
372 Effects: effects,
373 HasDefer: hasDefer,
374 HasBareReturn: hasBareReturn,
375 Returns: returnInfo,
376 Labels: labels,
377 Falcon: falcon,
378 }}, nil
379 }
380
381
382
383 func parseCompact(content []byte) (*token.FileSet, *ast.FuncDecl, error) {
384 fset := token.NewFileSet()
385 const mode = parser.ParseComments | parser.SkipObjectResolution | parser.AllErrors
386 f, err := parser.ParseFile(fset, "callee.go", content, mode)
387 if err != nil {
388 return nil, nil, fmt.Errorf("internal error: cannot compact file: %v", err)
389 }
390 return fset, f.Decls[0].(*ast.FuncDecl), nil
391 }
392
393
394 type paramInfo struct {
395 Name string
396 Index int
397 IsResult bool
398 IsInterface bool
399 Assigned bool
400 Escapes bool
401 Refs []refInfo
402 Shadow shadowMap
403 FalconType string
404 }
405
406 type refInfo struct {
407 Offset int
408 Assignable bool
409 IfaceAssignment bool
410 AffectsInference bool
411
412
413
414
415 IsSelectionOperand bool
416 }
417
418
419
420
421
422
423
424
425 func analyzeParams(logf func(string, ...any), fset *token.FileSet, info *types.Info, decl *ast.FuncDecl) (params, results []*paramInfo, effects []int, _ falconResult) {
426 sig := signature(fset, info, decl)
427
428 paramInfos := make(map[*types.Var]*paramInfo)
429 {
430 newParamInfo := func(param *types.Var, isResult bool) *paramInfo {
431 info := ¶mInfo{
432 Name: param.Name(),
433 IsResult: isResult,
434 Index: len(paramInfos),
435 IsInterface: isNonTypeParamInterface(param.Type()),
436 }
437 paramInfos[param] = info
438 return info
439 }
440 if sig.Recv() != nil {
441 params = append(params, newParamInfo(sig.Recv(), false))
442 }
443 for v := range sig.Params().Variables() {
444 params = append(params, newParamInfo(v, false))
445 }
446 for v := range sig.Results().Variables() {
447 results = append(results, newParamInfo(v, true))
448 }
449 }
450
451
452
453 escape(info, decl, func(v *types.Var, escapes bool) {
454 if info := paramInfos[v]; info != nil {
455 if escapes {
456 info.Escapes = true
457 } else {
458 info.Assigned = true
459 }
460 }
461 })
462
463
464
465
466
467
468 fieldObjs := fieldObjs(sig)
469 var stack []ast.Node
470 stack = append(stack, decl.Type)
471 ast.PreorderStack(decl.Body, stack, func(n ast.Node, stack []ast.Node) bool {
472 if id, ok := n.(*ast.Ident); ok {
473 if v, ok := info.Uses[id].(*types.Var); ok {
474 if pinfo, ok := paramInfos[v]; ok {
475
476
477
478
479
480
481
482
483
484
485
486 stack = append(stack, n)
487 assignable, ifaceAssign, affectsInference := analyzeAssignment(info, stack)
488 ref := refInfo{
489 Offset: int(n.Pos() - decl.Pos()),
490 Assignable: assignable,
491 IfaceAssignment: ifaceAssign,
492 AffectsInference: affectsInference,
493 IsSelectionOperand: isSelectionOperand(stack),
494 }
495 pinfo.Refs = append(pinfo.Refs, ref)
496 pinfo.Shadow = pinfo.Shadow.add(info, fieldObjs, pinfo.Name, stack)
497 }
498 }
499 }
500 return true
501 })
502
503
504
505 effects = calleefx(info, decl.Body, paramInfos)
506 logf("effects list = %v", effects)
507
508 falcon := falcon(logf, fset, paramInfos, info, decl)
509
510 return params, results, effects, falcon
511 }
512
513
514 func analyzeTypeParams(_ logger, fset *token.FileSet, info *types.Info, decl *ast.FuncDecl) []*paramInfo {
515 sig := signature(fset, info, decl)
516 paramInfos := make(map[*types.TypeName]*paramInfo)
517 var params []*paramInfo
518 collect := func(tpl *types.TypeParamList) {
519 for tparam := range tpl.TypeParams() {
520 typeName := tparam.Obj()
521 info := ¶mInfo{Name: typeName.Name()}
522 params = append(params, info)
523 paramInfos[typeName] = info
524 }
525 }
526 collect(sig.RecvTypeParams())
527 collect(sig.TypeParams())
528
529
530
531
532
533 visit := func(n ast.Node, stack []ast.Node) bool {
534 if id, ok := n.(*ast.Ident); ok {
535 if v, ok := info.Uses[id].(*types.TypeName); ok {
536 if pinfo, ok := paramInfos[v]; ok {
537 ref := refInfo{Offset: int(n.Pos() - decl.Pos())}
538 pinfo.Refs = append(pinfo.Refs, ref)
539 pinfo.Shadow = pinfo.Shadow.add(info, nil, pinfo.Name, stack)
540 }
541 }
542 }
543 return true
544 }
545 var stack []ast.Node
546 stack = append(stack, decl.Type)
547 if decl.Type.Params != nil {
548 ast.PreorderStack(decl.Type.Params, stack, visit)
549 }
550 if decl.Type.Results != nil {
551 ast.PreorderStack(decl.Type.Results, stack, visit)
552 }
553 ast.PreorderStack(decl.Body, stack, visit)
554 return params
555 }
556
557 func signature(fset *token.FileSet, info *types.Info, decl *ast.FuncDecl) *types.Signature {
558 fnobj, ok := info.Defs[decl.Name]
559 if !ok {
560 panic(fmt.Sprintf("%s: no func object for %q",
561 fset.PositionFor(decl.Name.Pos(), false), decl.Name))
562 }
563 return fnobj.Type().(*types.Signature)
564 }
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590 func analyzeAssignment(info *types.Info, stack []ast.Node) (assignable, ifaceAssign, affectsInference bool) {
591 remaining, parent, expr := exprContext(stack)
592 if parent == nil {
593 return false, false, false
594 }
595
596
597
598
599 if assign, ok := parent.(*ast.AssignStmt); ok {
600 for i, v := range assign.Rhs {
601 if v == expr {
602 if i >= len(assign.Lhs) {
603 return false, false, false
604 }
605
606 if i < len(assign.Lhs) {
607
608
609 if id, _ := assign.Lhs[i].(*ast.Ident); id != nil && info.Defs[id] != nil {
610
611
612 return false, false, false
613 }
614
615 typ := info.TypeOf(assign.Lhs[i])
616 return true, typ == nil || types.IsInterface(typ), false
617 }
618
619 return assign.Tok == token.ASSIGN, true, false
620 }
621 }
622 }
623
624
625 if spec, ok := parent.(*ast.ValueSpec); ok && spec.Type != nil {
626 if slices.Contains(spec.Values, expr) {
627 typ := info.TypeOf(spec.Type)
628 return true, typ == nil || types.IsInterface(typ), false
629 }
630 }
631
632
633 if ix, ok := parent.(*ast.IndexExpr); ok {
634 if ix.Index == expr {
635 typ := info.TypeOf(ix.X)
636 if typ == nil {
637 return true, true, false
638 }
639 m, _ := typeparams.CoreType(typ).(*types.Map)
640 return true, m == nil || types.IsInterface(m.Key()), false
641 }
642 }
643
644
645
646 if kv, ok := parent.(*ast.KeyValueExpr); ok {
647 var under types.Type
648 if len(remaining) > 0 {
649 if complit, ok := remaining[len(remaining)-1].(*ast.CompositeLit); ok {
650 if typ := info.TypeOf(complit); typ != nil {
651
652
653
654 under = typesinternal.Unpointer(typeparams.CoreType(typ))
655 }
656 }
657 }
658 if kv.Key == expr {
659 m, _ := under.(*types.Map)
660 return true, m == nil || types.IsInterface(m.Key()), false
661 }
662 if kv.Value == expr {
663 switch under := under.(type) {
664 case interface{ Elem() types.Type }:
665 return true, types.IsInterface(under.Elem()), false
666 case *types.Struct:
667 if id, _ := kv.Key.(*ast.Ident); id != nil {
668 for field := range under.Fields() {
669 if info.Uses[id] == field {
670 return true, types.IsInterface(field.Type()), false
671 }
672 }
673 }
674 default:
675 return true, true, false
676 }
677 }
678 }
679 if lit, ok := parent.(*ast.CompositeLit); ok {
680 for i, v := range lit.Elts {
681 if v == expr {
682 typ := info.TypeOf(lit)
683 if typ == nil {
684 return true, true, false
685 }
686
687
688 under := typesinternal.Unpointer(typeparams.CoreType(typ))
689 switch under := under.(type) {
690 case interface{ Elem() types.Type }:
691 return true, types.IsInterface(under.Elem()), false
692 case *types.Struct:
693 if i < under.NumFields() {
694 return true, types.IsInterface(under.Field(i).Type()), false
695 }
696 }
697 return true, true, false
698 }
699 }
700 }
701
702
703 if send, ok := parent.(*ast.SendStmt); ok {
704 if send.Value == expr {
705 typ := info.TypeOf(send.Chan)
706 if typ == nil {
707 return true, true, false
708 }
709 ch, _ := typeparams.CoreType(typ).(*types.Chan)
710 return true, ch == nil || types.IsInterface(ch.Elem()), false
711 }
712 }
713
714
715
716
717 if call, ok := parent.(*ast.CallExpr); ok {
718 if _, ok := isConversion(info, call); ok {
719 return false, false, false
720 }
721
722 for i, arg := range call.Args {
723 if arg == expr {
724 typ := info.TypeOf(call.Fun)
725 if typ == nil {
726 return true, true, false
727 }
728 sig, _ := typeparams.CoreType(typ).(*types.Signature)
729 if sig != nil {
730
731 paramType := paramTypeAtIndex(sig, call, i)
732 ifaceAssign := paramType == nil || types.IsInterface(paramType)
733 affectsInference := false
734 switch callee := typeutil.Callee(info, call).(type) {
735 case *types.Builtin:
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767 switch callee.Name() {
768 case "new", "complex", "real", "imag", "min", "max":
769 affectsInference = true
770 }
771
772 case *types.Func:
773
774
775 if sig2 := callee.Signature(); sig2.Recv() == nil {
776 originParamType := paramTypeAtIndex(sig2, call, i)
777 affectsInference = originParamType == nil || new(typeparams.Free).Has(originParamType)
778 }
779 }
780 return true, ifaceAssign, affectsInference
781 }
782 }
783 }
784 }
785
786 return false, false, false
787 }
788
789
790
791 func paramTypeAtIndex(sig *types.Signature, call *ast.CallExpr, index int) types.Type {
792 if plen := sig.Params().Len(); sig.Variadic() && index >= plen-1 && !call.Ellipsis.IsValid() {
793 if s, ok := sig.Params().At(plen - 1).Type().(*types.Slice); ok {
794 return s.Elem()
795 }
796 } else if index < plen {
797 return sig.Params().At(index).Type()
798 }
799 return nil
800 }
801
802
803
804
805
806
807 func exprContext(stack []ast.Node) (remaining []ast.Node, parent ast.Node, expr ast.Expr) {
808 expr, _ = stack[len(stack)-1].(ast.Expr)
809 if expr == nil {
810 return nil, nil, nil
811 }
812 i := len(stack) - 2
813 for ; i >= 0; i-- {
814 if pexpr, ok := stack[i].(*ast.ParenExpr); ok {
815 expr = pexpr
816 } else {
817 parent = stack[i]
818 break
819 }
820 }
821 if parent == nil {
822 return nil, nil, nil
823 }
824
825 return stack[:i], parent, expr
826 }
827
828
829
830 func isSelectionOperand(stack []ast.Node) bool {
831 _, parent, expr := exprContext(stack)
832 if parent == nil {
833 return false
834 }
835 sel, ok := parent.(*ast.SelectorExpr)
836 return ok && sel.X == expr
837 }
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858 type shadowMap map[string]int
859
860
861
862
863
864
865
866
867 func (s shadowMap) add(info *types.Info, paramIndexes map[types.Object]int, exclude string, stack []ast.Node) shadowMap {
868 for _, n := range stack {
869 if scope := scopeFor(info, n); scope != nil {
870 for _, name := range scope.Names() {
871 if name != exclude {
872 if s == nil {
873 s = make(shadowMap)
874 }
875 obj := scope.Lookup(name)
876 if idx, ok := paramIndexes[obj]; ok {
877 s[name] = idx + 1
878 } else {
879 s[name] = -1
880 }
881 }
882 }
883 }
884 }
885 return s
886 }
887
888
889
890
891 func fieldObjs(sig *types.Signature) map[types.Object]int {
892 m := make(map[types.Object]int)
893 for i := range sig.Params().Len() {
894 if p := sig.Params().At(i); p.Name() != "" && p.Name() != "_" {
895 m[p] = i
896 }
897 }
898 return m
899 }
900
901 func isField(obj types.Object) bool {
902 if v, ok := obj.(*types.Var); ok && v.IsField() {
903 return true
904 }
905 return false
906 }
907
908 func isMethod(obj types.Object) bool {
909 if f, ok := obj.(*types.Func); ok && f.Type().(*types.Signature).Recv() != nil {
910 return true
911 }
912 return false
913 }
914
915
916
917 var (
918 _ gob.GobEncoder = (*Callee)(nil)
919 _ gob.GobDecoder = (*Callee)(nil)
920 )
921
922 func (callee *Callee) GobEncode() ([]byte, error) {
923 var out bytes.Buffer
924 if err := gob.NewEncoder(&out).Encode(callee.impl); err != nil {
925 return nil, err
926 }
927 return out.Bytes(), nil
928 }
929
930 func (callee *Callee) GobDecode(data []byte) error {
931 return gob.NewDecoder(bytes.NewReader(data)).Decode(&callee.impl)
932 }
933
View as plain text