1
2
3
4
5 package modernize
6
7 import (
8 "fmt"
9 "go/ast"
10 "go/token"
11 "go/types"
12
13 "golang.org/x/tools/go/analysis"
14 "golang.org/x/tools/go/ast/edge"
15 "golang.org/x/tools/go/ast/inspector"
16 "golang.org/x/tools/go/types/typeutil"
17 "golang.org/x/tools/internal/analysis/analyzerutil"
18 typeindexanalyzer "golang.org/x/tools/internal/analysis/typeindex"
19 "golang.org/x/tools/internal/astutil"
20 "golang.org/x/tools/internal/stdlib"
21 "golang.org/x/tools/internal/typesinternal/typeindex"
22 )
23
24 var StdIteratorsAnalyzer = &analysis.Analyzer{
25 Name: "stditerators",
26 Doc: analyzerutil.MustExtractDoc(doc, "stditerators"),
27 Requires: []*analysis.Analyzer{
28 typeindexanalyzer.Analyzer,
29 },
30 Run: stditerators,
31 URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/modernize#stditerators",
32 }
33
34
35
36
37 var stditeratorsTable = [...]struct {
38 pkgpath, typename, lenmethod, atmethod, itermethod, elemname string
39
40 seqn int
41 }{
42
43
44
45 {"go/types", "Interface", "NumEmbeddeds", "EmbeddedType", "EmbeddedTypes", "etyp", 1},
46 {"go/types", "Interface", "NumExplicitMethods", "ExplicitMethod", "ExplicitMethods", "method", 1},
47 {"go/types", "Interface", "NumMethods", "Method", "Methods", "method", 1},
48 {"go/types", "MethodSet", "Len", "At", "Methods", "method", 1},
49 {"go/types", "Named", "NumMethods", "Method", "Methods", "method", 1},
50 {"go/types", "Scope", "NumChildren", "Child", "Children", "child", 1},
51 {"go/types", "Struct", "NumFields", "Field", "Fields", "field", 1},
52 {"go/types", "Tuple", "Len", "At", "Variables", "v", 1},
53 {"go/types", "TypeList", "Len", "At", "Types", "t", 1},
54 {"go/types", "TypeParamList", "Len", "At", "TypeParams", "tparam", 1},
55 {"go/types", "Union", "Len", "Term", "Terms", "term", 1},
56 {"reflect", "Type", "NumField", "Field", "Fields", "field", 1},
57 {"reflect", "Type", "NumMethod", "Method", "Methods", "method", 1},
58 {"reflect", "Type", "NumIn", "In", "Ins", "in", 1},
59 {"reflect", "Type", "NumOut", "Out", "Outs", "out", 1},
60 {"reflect", "Value", "NumField", "Field", "Fields", "field", 2},
61 {"reflect", "Value", "NumMethod", "Method", "Methods", "method", 2},
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 func stditerators(pass *analysis.Pass) (any, error) {
102 var (
103 index = pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
104 info = pass.TypesInfo
105 )
106
107 for _, row := range stditeratorsTable {
108
109
110 if within(pass, row.pkgpath) {
111 continue
112 }
113
114 var (
115 lenMethod = index.Selection(row.pkgpath, row.typename, row.lenmethod)
116 atMethod = index.Selection(row.pkgpath, row.typename, row.atmethod)
117 )
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138 chooseName := func(curBody inspector.Cursor, x ast.Expr, i *types.Var) (string, *types.Var) {
139
140
141
142 isVarAssign := func(stmt ast.Stmt) *types.Var {
143 if assign, ok := stmt.(*ast.AssignStmt); ok &&
144 assign.Tok == token.DEFINE &&
145 len(assign.Lhs) == 1 &&
146 len(assign.Rhs) == 1 &&
147 is[*ast.Ident](assign.Lhs[0]) {
148
149 if call, ok := assign.Rhs[0].(*ast.CallExpr); ok &&
150 typeutil.Callee(info, call) == atMethod &&
151 astutil.EqualSyntax(ast.Unparen(call.Fun).(*ast.SelectorExpr).X, x) &&
152 is[*ast.Ident](call.Args[0]) &&
153 info.Uses[call.Args[0].(*ast.Ident)] == i {
154
155 id := assign.Lhs[0].(*ast.Ident)
156 return info.Defs[id].(*types.Var)
157 }
158 }
159 return nil
160 }
161
162 body := curBody.Node().(*ast.BlockStmt)
163 if len(body.List) > 0 {
164
165 if v := isVarAssign(body.List[0]); v != nil {
166 return v.Name(), v
167 }
168
169
170 if ifstmt, ok := body.List[0].(*ast.IfStmt); ok && ifstmt.Init != nil {
171 if v := isVarAssign(ifstmt.Init); v != nil {
172 return v.Name(), v
173 }
174 }
175 }
176
177 loop := curBody.Parent().Node()
178
179
180 name := freshName(info, index, info.Scopes[loop], loop.Pos(), curBody, curBody, token.NoPos, row.elemname)
181 return name, nil
182 }
183
184
185 nextCall:
186 for curLenCall := range index.Calls(lenMethod) {
187 lenSel, ok := ast.Unparen(curLenCall.Node().(*ast.CallExpr).Fun).(*ast.SelectorExpr)
188 if !ok {
189 continue
190 }
191
192
193 var (
194 rng analysis.Range
195 curBody inspector.Cursor
196 indexVar *types.Var
197 elemVar *types.Var
198 elem string
199 edits []analysis.TextEdit
200 )
201
202
203 switch curLenCall.ParentEdgeKind() {
204 case edge.BinaryExpr_Y:
205
206 var (
207 curCmp = curLenCall.Parent()
208 cmp = curCmp.Node().(*ast.BinaryExpr)
209 )
210 if cmp.Op != token.LSS ||
211 curCmp.ParentEdgeKind() != edge.ForStmt_Cond {
212 continue
213 }
214 if id, ok := cmp.X.(*ast.Ident); ok {
215
216 var (
217 v = info.Uses[id].(*types.Var)
218 curFor = curCmp.Parent()
219 loop = curFor.Node().(*ast.ForStmt)
220 )
221 if v != isIncrementLoop(info, loop) {
222 continue
223 }
224
225
226 rng = astutil.RangeOf(loop.For, loop.Post.End())
227 indexVar = v
228 curBody = curFor.ChildAt(edge.ForStmt_Body, -1)
229 elem, elemVar = chooseName(curBody, lenSel.X, indexVar)
230 elemPrefix := cond(row.seqn == 2, "_, ", "")
231
232
233
234
235
236 edits = []analysis.TextEdit{
237 {
238 Pos: v.Pos(),
239 End: v.Pos() + token.Pos(len(v.Name())),
240 NewText: []byte(elemPrefix + elem),
241 },
242 {
243 Pos: loop.Init.(*ast.AssignStmt).Rhs[0].Pos(),
244 End: cmp.Y.Pos(),
245 NewText: []byte("range "),
246 },
247 {
248 Pos: lenSel.Sel.Pos(),
249 End: lenSel.Sel.End(),
250 NewText: []byte(row.itermethod),
251 },
252 {
253 Pos: curLenCall.Node().End(),
254 End: loop.Post.End(),
255 },
256 }
257 }
258
259 case edge.RangeStmt_X:
260
261 var (
262 curRange = curLenCall.Parent()
263 loop = curRange.Node().(*ast.RangeStmt)
264 )
265 if id, ok := loop.Key.(*ast.Ident); ok &&
266 loop.Value == nil &&
267 loop.Tok == token.DEFINE {
268
269
270
271 rng = astutil.RangeOf(loop.Range, loop.X.End())
272 indexVar = info.Defs[id].(*types.Var)
273 curBody = curRange.ChildAt(edge.RangeStmt_Body, -1)
274 elem, elemVar = chooseName(curBody, lenSel.X, indexVar)
275 elemPrefix := cond(row.seqn == 2, "_, ", "")
276
277
278
279
280 edits = []analysis.TextEdit{
281 {
282 Pos: loop.Key.Pos(),
283 End: loop.Key.End(),
284 NewText: []byte(elemPrefix + elem),
285 },
286 {
287 Pos: lenSel.Sel.Pos(),
288 End: lenSel.Sel.End(),
289 NewText: []byte(row.itermethod),
290 },
291 }
292 }
293 }
294
295 if indexVar == nil {
296 continue
297 }
298
299
300
301
302
303
304
305 for curUse := range index.Uses(indexVar) {
306 if !curBody.Contains(curUse) {
307 continue
308 }
309 if ek, argidx := curUse.ParentEdge(); ek != edge.CallExpr_Args || argidx != 0 {
310 continue nextCall
311 }
312 curAtCall := curUse.Parent()
313 atCall := curAtCall.Node().(*ast.CallExpr)
314 if typeutil.Callee(info, atCall) != atMethod {
315 continue nextCall
316 }
317 atSel := ast.Unparen(atCall.Fun).(*ast.SelectorExpr)
318
319
320 if !astutil.EqualSyntax(lenSel.X, atSel.X) {
321 continue nextCall
322 }
323
324
325
326
327
328
329 if obj := lookup(info, curAtCall, elem); obj != nil && obj != elemVar && obj.Pos() > indexVar.Pos() {
330
331
332 continue nextCall
333 }
334
335
336
337
338 edits = append(edits, analysis.TextEdit{
339 Pos: atCall.Pos(),
340 End: atCall.End(),
341 NewText: []byte(elem),
342 })
343 }
344
345
346
347
348
349 if v, err := methodGoVersion(row.pkgpath, row.typename, row.itermethod); err != nil {
350 panic(err)
351 } else if !analyzerutil.FileUsesGoVersion(pass, astutil.EnclosingFile(curLenCall), v.String()) {
352 continue nextCall
353 }
354
355 pass.Report(analysis.Diagnostic{
356 Pos: rng.Pos(),
357 End: rng.End(),
358 Message: fmt.Sprintf("%s/%s loop can simplified using %s.%s iteration",
359 row.lenmethod, row.atmethod, row.typename, row.itermethod),
360 SuggestedFixes: []analysis.SuggestedFix{{
361 Message: fmt.Sprintf(
362 "Replace %s/%s loop with %s.%s iteration",
363 row.lenmethod, row.atmethod, row.typename, row.itermethod),
364 TextEdits: edits,
365 }},
366 })
367 }
368 }
369 return nil, nil
370 }
371
372
373
374
375
376 func methodGoVersion(pkgpath, recvtype, method string) (stdlib.Version, error) {
377
378
379
380 for _, sym := range stdlib.PackageSymbols[pkgpath] {
381 if sym.Kind == stdlib.Method {
382 _, recv, name := sym.SplitMethod()
383 if recv == recvtype && name == method {
384 return sym.Version, nil
385 }
386 }
387 }
388 return 0, fmt.Errorf("methodGoVersion: %s.%s.%s missing from stdlib manifest", pkgpath, recvtype, method)
389 }
390
View as plain text