1
2
3
4
5 package modernize
6
7 import (
8 "fmt"
9 "go/ast"
10 "go/token"
11 "go/types"
12 "strings"
13
14 "golang.org/x/tools/go/analysis"
15 "golang.org/x/tools/go/analysis/passes/inspect"
16 "golang.org/x/tools/go/ast/edge"
17 "golang.org/x/tools/go/ast/inspector"
18 "golang.org/x/tools/internal/analysis/analyzerutil"
19 typeindexanalyzer "golang.org/x/tools/internal/analysis/typeindex"
20 "golang.org/x/tools/internal/astutil"
21 "golang.org/x/tools/internal/typeparams"
22 "golang.org/x/tools/internal/typesinternal"
23 "golang.org/x/tools/internal/typesinternal/typeindex"
24 "golang.org/x/tools/internal/versions"
25 )
26
27 var MinMaxAnalyzer = &analysis.Analyzer{
28 Name: "minmax",
29 Doc: analyzerutil.MustExtractDoc(doc, "minmax"),
30 Requires: []*analysis.Analyzer{
31 inspect.Analyzer,
32 typeindexanalyzer.Analyzer,
33 },
34 Run: minmax,
35 URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/modernize#minmax",
36 }
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58 func minmax(pass *analysis.Pass) (any, error) {
59 var (
60 inspect = pass.ResultOf[inspect.Analyzer].(*inspector.Inspector)
61 info = pass.TypesInfo
62 )
63
64 checkUserDefinedMinMax(pass)
65
66
67
68 check := func(file *ast.File, curIfStmt inspector.Cursor, compare *ast.BinaryExpr) {
69 var (
70 ifStmt = curIfStmt.Node().(*ast.IfStmt)
71 tassign = ifStmt.Body.List[0].(*ast.AssignStmt)
72 a = compare.X
73 b = compare.Y
74 lhs = tassign.Lhs[0]
75 rhs = tassign.Rhs[0]
76 sign = isInequality(compare.Op)
77
78
79 callArg = func(arg ast.Expr, start, end token.Pos) string {
80 comments := allComments(file, start, end)
81 return cond(arg == b, ", ", "") +
82 cond(comments != "", "\n", "") +
83 comments +
84 astutil.Format(pass.Fset, arg)
85 }
86 )
87
88 if fblock, ok := ifStmt.Else.(*ast.BlockStmt); ok && isAssignBlock(fblock) {
89 fassign := fblock.List[0].(*ast.AssignStmt)
90
91
92 lhs2 := fassign.Lhs[0]
93 rhs2 := fassign.Rhs[0]
94
95
96
97
98 if astutil.EqualSyntax(lhs, lhs2) {
99 if astutil.EqualSyntax(rhs, a) && astutil.EqualSyntax(rhs2, b) {
100 sign = +sign
101 } else if astutil.EqualSyntax(rhs2, a) && astutil.EqualSyntax(rhs, b) {
102 sign = -sign
103 } else {
104 return
105 }
106
107 sym := cond(sign < 0, "min", "max")
108
109 if !is[*types.Builtin](lookup(pass.TypesInfo, curIfStmt, sym)) {
110 return
111 }
112 if !analyzerutil.FileUsesGoVersion(pass, file, versions.Go1_21) {
113 return
114 }
115
116
117
118
119
120 pass.Report(analysis.Diagnostic{
121
122 Pos: compare.Pos(),
123 End: compare.End(),
124 Message: fmt.Sprintf("if/else statement can be modernized using %s", sym),
125 SuggestedFixes: []analysis.SuggestedFix{{
126 Message: fmt.Sprintf("Replace if statement with %s", sym),
127 TextEdits: []analysis.TextEdit{{
128
129 Pos: ifStmt.Pos(),
130 End: ifStmt.End(),
131 NewText: fmt.Appendf(nil, "%s = %s(%s%s)",
132 astutil.Format(pass.Fset, lhs),
133 sym,
134 callArg(a, ifStmt.Pos(), ifStmt.Else.Pos()),
135 callArg(b, ifStmt.Else.Pos(), ifStmt.End()),
136 ),
137 }},
138 }},
139 })
140 }
141
142 } else if prev, ok := curIfStmt.PrevSibling(); ok && isSimpleAssign(prev.Node()) && ifStmt.Else == nil {
143 fassign := prev.Node().(*ast.AssignStmt)
144
145
146
147
148
149
150
151
152
153
154
155 lhs0 := fassign.Lhs[0]
156 rhs0 := fassign.Rhs[0]
157
158
159
160
161 if prev.ParentEdgeKind() == edge.CommClause_Comm {
162 return
163 }
164
165 if astutil.EqualSyntax(lhs, lhs0) {
166 if astutil.EqualSyntax(rhs, a) && (astutil.EqualSyntax(rhs0, b) || astutil.EqualSyntax(lhs0, b)) {
167 sign = +sign
168 } else if (astutil.EqualSyntax(rhs0, a) || astutil.EqualSyntax(lhs0, a)) && astutil.EqualSyntax(rhs, b) {
169 sign = -sign
170 } else {
171 return
172 }
173 sym := cond(sign < 0, "min", "max")
174
175 if !is[*types.Builtin](lookup(pass.TypesInfo, curIfStmt, sym)) {
176 return
177 }
178
179
180
181
182 if astutil.EqualSyntax(lhs0, a) {
183 a = rhs0
184 } else if astutil.EqualSyntax(lhs0, b) {
185 b = rhs0
186 }
187
188 if !analyzerutil.FileUsesGoVersion(pass, file, versions.Go1_21) {
189 return
190 }
191
192
193 pass.Report(analysis.Diagnostic{
194
195 Pos: compare.Pos(),
196 End: compare.End(),
197 Message: fmt.Sprintf("if statement can be modernized using %s", sym),
198 SuggestedFixes: []analysis.SuggestedFix{{
199 Message: fmt.Sprintf("Replace if/else with %s", sym),
200 TextEdits: []analysis.TextEdit{{
201 Pos: fassign.Pos(),
202 End: ifStmt.End(),
203
204 NewText: fmt.Appendf(nil, "%s %s %s(%s%s)",
205 astutil.Format(pass.Fset, lhs),
206 fassign.Tok.String(),
207 sym,
208 callArg(a, fassign.Pos(), ifStmt.Pos()),
209 callArg(b, ifStmt.Pos(), ifStmt.End()),
210 ),
211 }},
212 }},
213 })
214 }
215 }
216 }
217
218
219 for curIfStmt := range inspect.Root().Preorder((*ast.IfStmt)(nil)) {
220 ifStmt := curIfStmt.Node().(*ast.IfStmt)
221
222
223
224
225
226
227 if curIfStmt.ParentEdgeKind() == edge.IfStmt_Else {
228 continue
229 }
230
231 if compare, ok := ifStmt.Cond.(*ast.BinaryExpr); ok &&
232 ifStmt.Init == nil &&
233 isInequality(compare.Op) != 0 &&
234 typesinternal.NoEffects(info, compare) &&
235 isAssignBlock(ifStmt.Body) {
236
237 if tLHS := info.TypeOf(ifStmt.Body.List[0].(*ast.AssignStmt).Lhs[0]); tLHS != nil && !maybeNaN(tLHS) {
238
239 check(astutil.EnclosingFile(curIfStmt), curIfStmt, compare)
240 }
241 }
242 }
243 return nil, nil
244 }
245
246
247 func allComments(file *ast.File, start, end token.Pos) string {
248 var buf strings.Builder
249 for co := range astutil.Comments(file, start, end) {
250 _, _ = fmt.Fprintf(&buf, "%s\n", co.Text)
251 }
252 return buf.String()
253 }
254
255
256
257 func isInequality(tok token.Token) int {
258 switch tok {
259 case token.LEQ, token.LSS:
260 return -1
261 case token.GEQ, token.GTR:
262 return +1
263 }
264 return 0
265 }
266
267
268 func isAssignBlock(b *ast.BlockStmt) bool {
269 if len(b.List) != 1 {
270 return false
271 }
272
273 return isSimpleAssign(b.List[0])
274 }
275
276
277 func isSimpleAssign(n ast.Node) bool {
278 assign, ok := n.(*ast.AssignStmt)
279 return ok &&
280 (assign.Tok == token.ASSIGN || assign.Tok == token.DEFINE) &&
281 len(assign.Lhs) == 1 &&
282 len(assign.Rhs) == 1
283 }
284
285
286 func maybeNaN(t types.Type) bool {
287
288
289
290 t = typeparams.CoreType(t)
291 if t == nil {
292 return true
293 }
294 if basic, ok := t.(*types.Basic); ok && basic.Info()&types.IsFloat != 0 {
295 return true
296 }
297 return false
298 }
299
300
301
302 func checkUserDefinedMinMax(pass *analysis.Pass) {
303 index := pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
304
305
306 for _, funcName := range []string{"min", "max"} {
307 if fn, ok := pass.Pkg.Scope().Lookup(funcName).(*types.Func); ok {
308
309 if def, ok := index.Def(fn); ok {
310 decl := def.Parent().Node().(*ast.FuncDecl)
311
312
313 if canUseBuiltinMinMax(fn, decl.Body) &&
314 analyzerutil.FileUsesGoVersion(pass, astutil.EnclosingFile(def), versions.Go1_21) {
315
316 pos := decl.Pos()
317 if docs := astutil.DocComment(decl); docs != nil {
318 pos = docs.Pos()
319 }
320
321 pass.Report(analysis.Diagnostic{
322 Pos: decl.Pos(),
323 End: decl.End(),
324 Message: fmt.Sprintf("user-defined %s function is equivalent to built-in %s and can be removed", funcName, funcName),
325 SuggestedFixes: []analysis.SuggestedFix{{
326 Message: fmt.Sprintf("Remove user-defined %s function", funcName),
327 TextEdits: []analysis.TextEdit{{
328 Pos: pos,
329 End: decl.End(),
330 }},
331 }},
332 })
333 }
334 }
335 }
336 }
337 }
338
339
340
341 func canUseBuiltinMinMax(fn *types.Func, body *ast.BlockStmt) bool {
342 sig := fn.Type().(*types.Signature)
343
344
345 if sig.Params().Len() != 2 {
346 return false
347 }
348
349
350 for param := range sig.Params().Variables() {
351 if maybeNaN(param.Type()) {
352 return false
353 }
354 }
355
356
357 if sig.Results().Len() != 1 {
358 return false
359 }
360
361
362 if body == nil {
363 return false
364 }
365
366 return hasMinMaxLogic(body, fn.Name(), sig.Params().At(0).Name(), sig.Params().At(1).Name())
367 }
368
369
370 func hasMinMaxLogic(body *ast.BlockStmt, funcName, param0, param1 string) bool {
371
372 if len(body.List) == 1 {
373 if ifStmt, ok := body.List[0].(*ast.IfStmt); ok {
374
375 if elseBlock, ok := ifStmt.Else.(*ast.BlockStmt); ok && len(elseBlock.List) == 1 {
376 if elseRet, ok := elseBlock.List[0].(*ast.ReturnStmt); ok && len(elseRet.Results) == 1 {
377 return checkMinMaxPattern(ifStmt, elseRet.Results[0], funcName, param0, param1)
378 }
379 }
380 }
381 }
382
383
384 if len(body.List) == 2 {
385 if ifStmt, ok := body.List[0].(*ast.IfStmt); ok && ifStmt.Else == nil {
386 if retStmt, ok := body.List[1].(*ast.ReturnStmt); ok && len(retStmt.Results) == 1 {
387 return checkMinMaxPattern(ifStmt, retStmt.Results[0], funcName, param0, param1)
388 }
389 }
390 }
391
392 return false
393 }
394
395
396
397
398
399
400 func checkMinMaxPattern(ifStmt *ast.IfStmt, falseResult ast.Expr, funcName, param0, param1 string) bool {
401
402 cmp, ok := ifStmt.Cond.(*ast.BinaryExpr)
403 if !ok {
404 return false
405 }
406
407
408 if len(ifStmt.Body.List) != 1 {
409 return false
410 }
411
412 thenRet, ok := ifStmt.Body.List[0].(*ast.ReturnStmt)
413 if !ok || len(thenRet.Results) != 1 {
414 return false
415 }
416
417
418 sign := isInequality(cmp.Op)
419 if sign == 0 {
420 return false
421 }
422
423 t := thenRet.Results[0]
424 f := falseResult
425 x, ok := cmp.X.(*ast.Ident)
426 if !ok {
427 return false
428 }
429 y, ok := cmp.Y.(*ast.Ident)
430 if !ok {
431 return false
432 }
433
434
435
436
437 if !(param0 == x.Name && param1 == y.Name ||
438 param0 == y.Name && param1 == x.Name) {
439 return false
440 }
441
442
443 if astutil.EqualSyntax(t, x) && astutil.EqualSyntax(f, y) {
444 sign = +sign
445 } else if astutil.EqualSyntax(t, y) && astutil.EqualSyntax(f, x) {
446 sign = -sign
447 } else {
448 return false
449 }
450
451
452 return cond(sign < 0, "min", "max") == funcName
453 }
454
View as plain text