1
2
3
4
5 package modernize
6
7 import (
8 "fmt"
9 "go/ast"
10 "go/constant"
11 "go/token"
12 "go/types"
13 "iter"
14 "strconv"
15
16 "golang.org/x/tools/go/analysis"
17 "golang.org/x/tools/go/analysis/passes/inspect"
18 "golang.org/x/tools/go/ast/edge"
19 "golang.org/x/tools/go/ast/inspector"
20 "golang.org/x/tools/go/types/typeutil"
21 "golang.org/x/tools/internal/analysis/analyzerutil"
22 typeindexanalyzer "golang.org/x/tools/internal/analysis/typeindex"
23 "golang.org/x/tools/internal/astutil"
24 "golang.org/x/tools/internal/moreiters"
25 "golang.org/x/tools/internal/typesinternal"
26 "golang.org/x/tools/internal/typesinternal/typeindex"
27 "golang.org/x/tools/internal/versions"
28 )
29
30 var StringsCutAnalyzer = &analysis.Analyzer{
31 Name: "stringscut",
32 Doc: analyzerutil.MustExtractDoc(doc, "stringscut"),
33 Requires: []*analysis.Analyzer{
34 inspect.Analyzer,
35 typeindexanalyzer.Analyzer,
36 },
37 Run: stringscut,
38 URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/modernize#stringscut",
39 }
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110 func stringscut(pass *analysis.Pass) (any, error) {
111 var (
112 index = pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
113 info = pass.TypesInfo
114
115 stringsIndex = index.Object("strings", "Index")
116 stringsIndexByte = index.Object("strings", "IndexByte")
117 bytesIndex = index.Object("bytes", "Index")
118 bytesIndexByte = index.Object("bytes", "IndexByte")
119 )
120
121 stringsplitCut(pass, index)
122
123 scopeFixCount := make(map[*types.Scope]int)
124
125 for _, obj := range []types.Object{
126 stringsIndex,
127 stringsIndexByte,
128 bytesIndex,
129 bytesIndexByte,
130 } {
131
132 nextcall:
133 for curCall := range index.Calls(obj) {
134
135 if !analyzerutil.FileUsesGoVersion(pass, astutil.EnclosingFile(curCall), versions.Go1_18) {
136 continue
137 }
138 indexCall := curCall.Node().(*ast.CallExpr)
139 obj := typeutil.Callee(info, indexCall)
140 if obj == nil {
141 continue
142 }
143
144 var iIdent *ast.Ident
145 switch ek, idx := curCall.ParentEdge(); ek {
146 case edge.ValueSpec_Values:
147
148
149 spec := curCall.Parent().Node().(*ast.ValueSpec)
150 if len(spec.Names) != 1 {
151 continue
152 }
153 curName := curCall.Parent().ChildAt(edge.ValueSpec_Names, idx)
154 iIdent = curName.Node().(*ast.Ident)
155 case edge.AssignStmt_Rhs:
156
157
158 assign := curCall.Parent().Node().(*ast.AssignStmt)
159 if len(assign.Lhs) != 1 {
160 continue
161 }
162 curLhs := curCall.Parent().ChildAt(edge.AssignStmt_Lhs, idx)
163 iIdent, _ = curLhs.Node().(*ast.Ident)
164 }
165
166 if iIdent == nil {
167 continue
168 }
169
170
171 iObj := info.ObjectOf(iIdent)
172 if iObj == nil {
173 continue
174 }
175
176 var (
177 s = indexCall.Args[0]
178 substr = indexCall.Args[1]
179 )
180
181
182
183 if !indexArgValid(info, index, s, indexCall.Pos()) ||
184 !indexArgValid(info, index, substr, indexCall.Pos()) {
185 continue nextcall
186 }
187
188
189
190
191
192
193 negative, nonnegative, beforeSlice, afterSlice := checkIdxUses(pass.TypesInfo, index.Uses(iObj), s, substr, iObj)
194
195
196
197 if negative == nil && nonnegative == nil && beforeSlice == nil && afterSlice == nil {
198 continue
199 }
200
201
202 isContains := (len(negative) > 0 || len(nonnegative) > 0) && len(beforeSlice) == 0 && len(afterSlice) == 0
203
204 enclosingBlock, ok := moreiters.First(curCall.Enclosing((*ast.BlockStmt)(nil)))
205 if !ok {
206 continue
207 }
208 scope := iObj.Parent()
209
210
211
212
213
214
215
216 lastStmtCur, _ := enclosingBlock.LastChild()
217 lastStmt := lastStmtCur.Node()
218
219 fresh := func(preferred string) string {
220 return freshName(info, index, scope, lastStmt.End(), lastStmtCur, enclosingBlock, iIdent.Pos(), preferred)
221 }
222
223 var okVarName, beforeVarName, afterVarName, foundVarName string
224 if isContains {
225 foundVarName = fresh("found")
226 } else {
227 okVarName = fresh("ok")
228 beforeVarName = fresh("before")
229 afterVarName = fresh("after")
230 }
231
232
233
234
235
236 if scopeFixCount[scope] > 0 {
237 suffix := scopeFixCount[scope] - 1
238 if isContains {
239 foundVarName = fresh(fmt.Sprintf("%s%d", foundVarName, suffix))
240 } else {
241 okVarName = fresh(fmt.Sprintf("%s%d", okVarName, suffix))
242 beforeVarName = fresh(fmt.Sprintf("%s%d", beforeVarName, suffix))
243 afterVarName = fresh(fmt.Sprintf("%s%d", afterVarName, suffix))
244 }
245 }
246
247
248
249 if len(negative) == 0 && len(nonnegative) == 0 {
250 okVarName = "_"
251 }
252 if len(beforeSlice) == 0 {
253 beforeVarName = "_"
254 }
255 if len(afterSlice) == 0 {
256 afterVarName = "_"
257 }
258
259 var edits []analysis.TextEdit
260 replace := func(exprs []ast.Expr, new string) {
261 for _, expr := range exprs {
262 edits = append(edits, analysis.TextEdit{
263 Pos: expr.Pos(),
264 End: expr.End(),
265 NewText: []byte(new),
266 })
267 }
268 }
269
270
271 indexCallId := typesinternal.UsedIdent(info, indexCall.Fun)
272 replacedFunc := "Cut"
273 if isContains {
274 replacedFunc = "Contains"
275 replace(negative, "!"+foundVarName)
276 replace(nonnegative, foundVarName)
277
278
279
280
281
282
283 edits = append(edits, analysis.TextEdit{
284 Pos: iIdent.Pos(),
285 End: iIdent.End(),
286 NewText: []byte(foundVarName),
287 }, analysis.TextEdit{
288 Pos: indexCallId.Pos(),
289 End: indexCallId.End(),
290 NewText: []byte("Contains"),
291 })
292 } else {
293 replace(negative, "!"+okVarName)
294 replace(nonnegative, okVarName)
295 replace(beforeSlice, beforeVarName)
296 replace(afterSlice, afterVarName)
297
298
299
300
301
302
303 edits = append(edits, analysis.TextEdit{
304 Pos: iIdent.Pos(),
305 End: iIdent.End(),
306 NewText: fmt.Appendf(nil, "%s, %s, %s", beforeVarName, afterVarName, okVarName),
307 }, analysis.TextEdit{
308 Pos: indexCallId.Pos(),
309 End: indexCallId.End(),
310 NewText: []byte("Cut"),
311 })
312 }
313
314
315
316 if obj.Name() == "IndexByte" {
317 switch obj.Pkg().Name() {
318 case "strings":
319 searchByteVal := info.Types[substr].Value
320 if searchByteVal == nil {
321
322
323 edits = append(edits, []analysis.TextEdit{
324 {
325 Pos: substr.Pos(),
326 NewText: []byte("string("),
327 },
328 {
329 Pos: substr.End(),
330 NewText: []byte(")"),
331 },
332 }...)
333 } else {
334
335 val, _ := constant.Int64Val(searchByteVal)
336
337 edits = append(edits, analysis.TextEdit{
338 Pos: substr.Pos(),
339 End: substr.End(),
340 NewText: strconv.AppendQuote(nil, string(byte(val))),
341 })
342 }
343 case "bytes":
344
345 edits = append(edits, []analysis.TextEdit{
346 {
347 Pos: substr.Pos(),
348 NewText: []byte("[]byte{"),
349 },
350 {
351 Pos: substr.End(),
352 NewText: []byte("}"),
353 },
354 }...)
355 }
356 }
357 scopeFixCount[scope]++
358 pass.Report(analysis.Diagnostic{
359 Pos: indexCall.Fun.Pos(),
360 End: indexCall.Fun.End(),
361 Message: fmt.Sprintf("%s.%s can be simplified using %s.%s",
362 obj.Pkg().Name(), obj.Name(), obj.Pkg().Name(), replacedFunc),
363 Category: "stringscut",
364 SuggestedFixes: []analysis.SuggestedFix{{
365 Message: fmt.Sprintf("Simplify %s.%s call using %s.%s", obj.Pkg().Name(), obj.Name(), obj.Pkg().Name(), replacedFunc),
366 TextEdits: edits,
367 }},
368 })
369 }
370 }
371
372 return nil, nil
373 }
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390 func stringsplitCut(pass *analysis.Pass, index *typeindex.Index) {
391 info := pass.TypesInfo
392
393 stringsSplit := index.Object("strings", "Split")
394 stringsSplitN := index.Object("strings", "SplitN")
395
396 for _, obj := range []types.Object{stringsSplit, stringsSplitN} {
397 for curCall := range index.Calls(obj) {
398 callExpr := curCall.Node().(*ast.CallExpr)
399
400
401 if obj.Name() == "SplitN" && !isIntLiteral(info, callExpr.Args[2], 2) {
402 continue
403 }
404
405
406
407
408
409
410 sepTV := info.Types[callExpr.Args[1]]
411 if sepTV.Value == nil || constant.StringVal(sepTV.Value) == "" {
412 continue
413 }
414
415
416 if curCall.ParentEdgeKind() != edge.IndexExpr_X {
417 continue
418 }
419 parent := curCall.Parent()
420 indexExpr := parent.Node().(*ast.IndexExpr)
421
422
423 if !isZeroIntConst(info, indexExpr.Index) {
424 continue
425 }
426
427
428 if parent.ParentEdgeKind() != edge.AssignStmt_Rhs {
429 continue
430 }
431 assign := parent.Parent().Node().(*ast.AssignStmt)
432 if assign.Tok != token.DEFINE || len(assign.Lhs) != 1 {
433 continue
434 }
435
436
437 lhsIdent, ok := assign.Lhs[0].(*ast.Ident)
438 if !ok || lhsIdent.Name == "_" {
439 continue
440 }
441
442
443 if !analyzerutil.FileUsesGoVersion(pass, astutil.EnclosingFile(curCall), versions.Go1_18) {
444 continue
445 }
446
447
448
449
450
451
452 callFunIdent := typesinternal.UsedIdent(info, callExpr.Fun)
453
454 var edits []analysis.TextEdit
455
456
457 edits = append(edits, analysis.TextEdit{
458 Pos: lhsIdent.End(),
459 End: lhsIdent.End(),
460 NewText: []byte(", _, _"),
461 })
462
463
464 edits = append(edits, analysis.TextEdit{
465 Pos: callFunIdent.Pos(),
466 End: callFunIdent.End(),
467 NewText: []byte("Cut"),
468 })
469
470
471 if obj.Name() == "SplitN" {
472 edits = append(edits, analysis.TextEdit{
473 Pos: callExpr.Args[1].End(),
474 End: callExpr.Rparen,
475 })
476 }
477
478
479 edits = append(edits, analysis.TextEdit{
480 Pos: indexExpr.Lbrack,
481 End: indexExpr.End(),
482 })
483
484 pass.Report(analysis.Diagnostic{
485 Pos: callExpr.Fun.Pos(),
486 End: callExpr.Fun.End(),
487 Message: fmt.Sprintf("strings.%s call can be simplified using strings.Cut", obj.Name()),
488 Category: "stringscut",
489 SuggestedFixes: []analysis.SuggestedFix{{
490 Message: fmt.Sprintf("Simplify strings.%s call using strings.Cut", obj.Name()),
491 TextEdits: edits,
492 }},
493 })
494 }
495 }
496 }
497
498
499
500
501
502
503
504
505 func indexArgValid(info *types.Info, index *typeindex.Index, expr ast.Expr, afterPos token.Pos) bool {
506 tv := info.Types[expr]
507 if tv.Value != nil {
508 return true
509 }
510 switch expr := expr.(type) {
511 case *ast.CallExpr:
512 return types.Identical(tv.Type, byteSliceType) &&
513 info.Types[expr.Fun].IsType() &&
514 indexArgValid(info, index, expr.Args[0], afterPos)
515 case *ast.Ident:
516 for use := range index.Uses(info.Uses[expr]) {
517 if typesinternal.IsAssignedOrAddressTaken(info, use) {
518 return false
519 }
520 }
521 return true
522 default:
523
524
525
526
527
528
529
530
531
532 return false
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 func checkIdxUses(info *types.Info, uses iter.Seq[inspector.Cursor], s, substr ast.Expr, iObj types.Object) (negative, nonnegative, beforeSlice, afterSlice []ast.Expr) {
558 requireGuard := true
559 if l := constSubstrLen(info, substr); l != -1 && l != 1 {
560 requireGuard = false
561 }
562
563 use := func(cur inspector.Cursor) bool {
564 ek := cur.ParentEdgeKind()
565 n := cur.Parent().Node()
566 switch ek {
567 case edge.BinaryExpr_X, edge.BinaryExpr_Y:
568 check := n.(*ast.BinaryExpr)
569 switch checkIdxComparison(info, check, iObj) {
570 case -1:
571 negative = append(negative, check)
572 return true
573 case 1:
574 nonnegative = append(nonnegative, check)
575 return true
576 }
577
578
579
580
581
582 if slice, ok := cur.Parent().Parent().Node().(*ast.SliceExpr); ok &&
583 sameObject(info, s, slice.X) &&
584 slice.Max == nil {
585 if isBeforeSlice(info, ek, slice) && (!requireGuard || isSliceIndexGuarded(info, cur, iObj)) {
586 beforeSlice = append(beforeSlice, slice)
587 return true
588 } else if isAfterSlice(info, ek, slice, substr) && (!requireGuard || isSliceIndexGuarded(info, cur, iObj)) {
589 afterSlice = append(afterSlice, slice)
590 return true
591 }
592 }
593 case edge.SliceExpr_Low, edge.SliceExpr_High:
594 slice := n.(*ast.SliceExpr)
595
596
597 if sameObject(info, s, slice.X) && slice.Max == nil {
598 if isBeforeSlice(info, ek, slice) && (!requireGuard || isSliceIndexGuarded(info, cur, iObj)) {
599 beforeSlice = append(beforeSlice, slice)
600 return true
601 } else if isAfterSlice(info, ek, slice, substr) && (!requireGuard || isSliceIndexGuarded(info, cur, iObj)) {
602 afterSlice = append(afterSlice, slice)
603 return true
604 }
605 }
606 }
607 return false
608 }
609
610 for curIdent := range uses {
611 if !use(curIdent) {
612 return nil, nil, nil, nil
613 }
614 }
615 return negative, nonnegative, beforeSlice, afterSlice
616 }
617
618
619
620
621
622
623
624
625
626 func checkIdxComparison(info *types.Info, check *ast.BinaryExpr, iObj types.Object) int {
627 isI := func(e ast.Expr) bool {
628 id, ok := e.(*ast.Ident)
629 return ok && info.Uses[id] == iObj
630 }
631 if !isI(check.X) && !isI(check.Y) {
632 return 0
633 }
634
635
636 x, op, y := check.X, check.Op, check.Y
637 if info.Types[x].Value != nil {
638 x, op, y = y, flip(op), x
639 }
640
641 yIsInt := func(k int64) bool {
642 return isIntLiteral(info, y, k)
643 }
644
645 if op == token.LSS && yIsInt(0) ||
646 op == token.EQL && yIsInt(-1) ||
647 op == token.LEQ && yIsInt(-1) {
648 return -1
649 }
650
651 if op == token.GEQ && yIsInt(0) ||
652 op == token.NEQ && yIsInt(-1) ||
653 op == token.GTR && yIsInt(-1) {
654 return +1
655 }
656
657 return 0
658 }
659
660
661
662 func flip(op token.Token) token.Token {
663 switch op {
664 case token.EQL:
665 return token.EQL
666 case token.GEQ:
667 return token.LEQ
668 case token.GTR:
669 return token.LSS
670 case token.LEQ:
671 return token.GEQ
672 case token.LSS:
673 return token.GTR
674 }
675 return op
676 }
677
678
679 func isBeforeSlice(info *types.Info, ek edge.Kind, slice *ast.SliceExpr) bool {
680 return ek == edge.SliceExpr_High && (slice.Low == nil || isZeroIntConst(info, slice.Low))
681 }
682
683
684 func constSubstrLen(info *types.Info, substr ast.Expr) int {
685
686 if call, ok := substr.(*ast.CallExpr); ok {
687 tv := info.Types[call.Fun]
688 if tv.IsType() && types.Identical(tv.Type, byteSliceType) {
689
690 substr = call.Args[0]
691 }
692 }
693 substrVal := info.Types[substr].Value
694 if substrVal != nil {
695 switch substrVal.Kind() {
696 case constant.String:
697 return len(constant.StringVal(substrVal))
698 case constant.Int:
699
700
701
702 return 1
703 }
704 }
705 return -1
706 }
707
708
709
710 func isAfterSlice(info *types.Info, ek edge.Kind, slice *ast.SliceExpr, substr ast.Expr) bool {
711 lowExpr, ok := slice.Low.(*ast.BinaryExpr)
712 if !ok || slice.High != nil {
713 return false
714 }
715
716 isLenCall := func(expr ast.Expr) bool {
717 call, ok := expr.(*ast.CallExpr)
718 if !ok || len(call.Args) != 1 {
719 return false
720 }
721 return sameObject(info, substr, call.Args[0]) && typeutil.Callee(info, call) == builtinLen
722 }
723
724 substrLen := constSubstrLen(info, substr)
725
726 switch ek {
727 case edge.BinaryExpr_X:
728 kVal := info.Types[lowExpr.Y].Value
729 if kVal == nil {
730
731 return lowExpr.Op == token.ADD && isLenCall(lowExpr.Y)
732 } else {
733
734 kInt, ok := constant.Int64Val(kVal)
735 return ok && substrLen == int(kInt)
736 }
737 case edge.BinaryExpr_Y:
738 kVal := info.Types[lowExpr.X].Value
739 if kVal == nil {
740
741 return lowExpr.Op == token.ADD && isLenCall(lowExpr.X)
742 } else {
743
744 kInt, ok := constant.Int64Val(kVal)
745 return ok && substrLen == int(kInt)
746 }
747 }
748 return false
749 }
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765 func isSliceIndexGuarded(info *types.Info, cur inspector.Cursor, iObj types.Object) bool {
766 for anc := range cur.Enclosing() {
767 switch anc.ParentEdgeKind() {
768 case edge.IfStmt_Body, edge.IfStmt_Else:
769 ifStmt := anc.Parent().Node().(*ast.IfStmt)
770 check := condChecksIdx(info, ifStmt.Cond, iObj)
771 if anc.ParentEdgeKind() == edge.IfStmt_Else {
772 check = -check
773 }
774 if check > 0 {
775 return true
776 }
777 if check < 0 {
778 return false
779 }
780 case edge.BlockStmt_List:
781
782 for sib, ok := anc.PrevSibling(); ok; sib, ok = sib.PrevSibling() {
783 ifStmt, ok := sib.Node().(*ast.IfStmt)
784 if ok && condChecksIdx(info, ifStmt.Cond, iObj) < 0 && bodyTerminates(ifStmt.Body) {
785 return true
786 }
787 }
788 case edge.FuncDecl_Body, edge.FuncLit_Body:
789 return false
790 }
791 }
792 return false
793 }
794
795
796
797
798 func condChecksIdx(info *types.Info, cond ast.Expr, iObj types.Object) int {
799 binExpr, ok := cond.(*ast.BinaryExpr)
800 if !ok {
801 return 0
802 }
803 return checkIdxComparison(info, binExpr, iObj)
804 }
805
806
807
808 func bodyTerminates(block *ast.BlockStmt) bool {
809 if len(block.List) == 0 {
810 return false
811 }
812 last := block.List[len(block.List)-1]
813 switch last.(type) {
814 case *ast.ReturnStmt, *ast.BranchStmt:
815 return true
816 }
817 return false
818 }
819
820
821 func sameObject(info *types.Info, expr1, expr2 ast.Expr) bool {
822 if ident1, ok := expr1.(*ast.Ident); ok {
823 if ident2, ok := expr2.(*ast.Ident); ok {
824 uses1, ok1 := info.Uses[ident1]
825 uses2, ok2 := info.Uses[ident2]
826 return ok1 && ok2 && uses1 == uses2
827 }
828 }
829 return false
830 }
831
View as plain text