Source file src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/minmax.go

     1  // Copyright 2024 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     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  // The minmax pass replaces if/else statements with calls to min or max,
    39  // and removes user-defined min/max functions that are equivalent to built-ins.
    40  //
    41  // If/else replacement patterns:
    42  //
    43  //  1. if a < b { x = a } else { x = b }        =>      x = min(a, b)
    44  //  2. x = a; if a < b { x = b }                =>      x = max(a, b)
    45  //
    46  // Pattern 1 requires that a is not NaN, and pattern 2 requires that b
    47  // is not Nan. Since this is hard to prove, we reject floating-point
    48  // numbers.
    49  //
    50  // Function removal:
    51  // User-defined min/max functions are suggested for removal if they may
    52  // be safely replaced by their built-in namesake.
    53  //
    54  // Variants:
    55  // - all four ordered comparisons
    56  // - "x := a" or "x = a" or "var x = a" in pattern 2
    57  // - "x < b" or "a < b" in pattern 2
    58  func minmax(pass *analysis.Pass) (any, error) {
    59  	var (
    60  		inspect = pass.ResultOf[inspect.Analyzer].(*inspector.Inspector)
    61  		info    = pass.TypesInfo
    62  	)
    63  	// Check for user-defined min/max functions that can be removed
    64  	checkUserDefinedMinMax(pass)
    65  
    66  	// check is called for all statements of this form:
    67  	//   if a < b { lhs = rhs }
    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  			// callArg formats a call argument, preserving comments from [start-end).
    79  			callArg = func(arg ast.Expr, start, end token.Pos) string {
    80  				comments := allComments(file, start, end)
    81  				return cond(arg == b, ", ", "") + // second argument needs a comma
    82  					cond(comments != "", "\n", "") + // comments need their own line
    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  			// Have: if a < b { lhs = rhs } else { lhs2 = rhs2 }
    92  			lhs2 := fassign.Lhs[0]
    93  			rhs2 := fassign.Rhs[0]
    94  
    95  			// For pattern 1, check that:
    96  			// - lhs = lhs2
    97  			// - {rhs,rhs2} = {a,b}
    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 // min/max function is shadowed
   111  				}
   112  				if !analyzerutil.FileUsesGoVersion(pass, file, versions.Go1_21) {
   113  					return // min/max is too new
   114  				}
   115  
   116  				// pattern 1
   117  				//
   118  				// TODO(adonovan): if lhs is declared "var lhs T" on preceding line,
   119  				// simplify the whole thing to "lhs := min(a, b)".
   120  				pass.Report(analysis.Diagnostic{
   121  					// Highlight the condition a < b.
   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  							// Replace IfStmt with lhs = min(a, b).
   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  			// Have: lhs0 = rhs0; if a < b { lhs = rhs }
   146  			//
   147  			// For pattern 2, check that
   148  			// - lhs = lhs0
   149  			// - {a,b} = {rhs,rhs0} or {rhs,lhs0}
   150  			//   The replacement must use rhs0 not lhs0 though.
   151  			//   For example, we accept this variant:
   152  			//     lhs = x; if lhs < y { lhs = y }   =>   lhs = min(x, y), not min(lhs, y)
   153  			//
   154  			// TODO(adonovan): accept "var lhs0 = rhs0" form too.
   155  			lhs0 := fassign.Lhs[0]
   156  			rhs0 := fassign.Rhs[0]
   157  
   158  			// If the assignment occurs within a select
   159  			// comms clause (like "case lhs0 := <-rhs0:"),
   160  			// there's no way of rewriting it into a min/max call.
   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 // min/max function is shadowed
   177  				}
   178  
   179  				// Permit lhs0 to stand for rhs0 in the matching,
   180  				// but don't actually reduce to lhs0 = min(lhs0, rhs)
   181  				// since the "=" could be a ":=". Use min(rhs0, rhs).
   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 // min/max is too new
   190  				}
   191  
   192  				// pattern 2
   193  				pass.Report(analysis.Diagnostic{
   194  					// Highlight the condition a < b.
   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  							// Replace "x := a; if ... {}" with "x = min(...)", preserving comments.
   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  	// Find all "if a < b { lhs = rhs }" statements.
   219  	for curIfStmt := range inspect.Root().Preorder((*ast.IfStmt)(nil)) {
   220  		ifStmt := curIfStmt.Node().(*ast.IfStmt)
   221  		// Don't bother handling "if a < b { lhs = rhs }" when it appears
   222  		// as the "else" branch of another if-statement.
   223  		//    if cond { ... } else if a < b { lhs = rhs }
   224  		// (This case would require introducing another block
   225  		//    if cond { ... } else { if a < b { lhs = rhs } }
   226  		// and checking that there is no following "else".)
   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  			// a blank var has no type.
   237  			if tLHS := info.TypeOf(ifStmt.Body.List[0].(*ast.AssignStmt).Lhs[0]); tLHS != nil && !maybeNaN(tLHS) {
   238  				// Have: if a < b { lhs = rhs }
   239  				check(astutil.EnclosingFile(curIfStmt), curIfStmt, compare)
   240  			}
   241  		}
   242  	}
   243  	return nil, nil
   244  }
   245  
   246  // allComments collects all the comments from start to end.
   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  // isInequality reports non-zero if tok is one of < <= => >:
   256  // +1 for > and -1 for <.
   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  // isAssignBlock reports whether b is a block of the form { lhs = rhs }.
   268  func isAssignBlock(b *ast.BlockStmt) bool {
   269  	if len(b.List) != 1 {
   270  		return false
   271  	}
   272  	// Inv: the sole statement cannot be { lhs := rhs }.
   273  	return isSimpleAssign(b.List[0])
   274  }
   275  
   276  // isSimpleAssign reports whether n has the form "lhs = rhs" or "lhs := rhs".
   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  // maybeNaN reports whether t is (or may be) a floating-point type.
   286  func maybeNaN(t types.Type) bool {
   287  	// For now, we rely on core types.
   288  	// TODO(adonovan): In the post-core-types future,
   289  	// follow the approach of types.Checker.applyTypeFunc.
   290  	t = typeparams.CoreType(t)
   291  	if t == nil {
   292  		return true // fail safe
   293  	}
   294  	if basic, ok := t.(*types.Basic); ok && basic.Info()&types.IsFloat != 0 {
   295  		return true
   296  	}
   297  	return false
   298  }
   299  
   300  // checkUserDefinedMinMax looks for user-defined min/max functions that are
   301  // equivalent to the built-in functions and suggests removing them.
   302  func checkUserDefinedMinMax(pass *analysis.Pass) {
   303  	index := pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
   304  
   305  	// Look up min and max functions by name in package scope
   306  	for _, funcName := range []string{"min", "max"} {
   307  		if fn, ok := pass.Pkg.Scope().Lookup(funcName).(*types.Func); ok {
   308  			// Use typeindex to get the FuncDecl directly
   309  			if def, ok := index.Def(fn); ok {
   310  				decl := def.Parent().Node().(*ast.FuncDecl)
   311  				// Check if this function matches the built-in min/max signature
   312  				// and behavior, and verify that we have go1.21.
   313  				if canUseBuiltinMinMax(fn, decl.Body) &&
   314  					analyzerutil.FileUsesGoVersion(pass, astutil.EnclosingFile(def), versions.Go1_21) {
   315  					// Expand to include leading doc comment
   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  // canUseBuiltinMinMax reports whether it is safe to replace a call
   340  // to this min or max function by its built-in namesake.
   341  func canUseBuiltinMinMax(fn *types.Func, body *ast.BlockStmt) bool {
   342  	sig := fn.Type().(*types.Signature)
   343  
   344  	// Only consider the most common case: exactly 2 parameters
   345  	if sig.Params().Len() != 2 {
   346  		return false
   347  	}
   348  
   349  	// Check if any parameter might be floating-point
   350  	for param := range sig.Params().Variables() {
   351  		if maybeNaN(param.Type()) {
   352  			return false // Don't suggest removal for float types due to NaN handling
   353  		}
   354  	}
   355  
   356  	// Must have exactly one return value
   357  	if sig.Results().Len() != 1 {
   358  		return false
   359  	}
   360  
   361  	// Check that the function body implements the expected min/max logic
   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  // hasMinMaxLogic checks if the function body implements simple min/max logic.
   370  func hasMinMaxLogic(body *ast.BlockStmt, funcName, param0, param1 string) bool {
   371  	// Pattern 1: Single if/else statement
   372  	if len(body.List) == 1 {
   373  		if ifStmt, ok := body.List[0].(*ast.IfStmt); ok {
   374  			// Get the "false" result from the else block
   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  	// Pattern 2: if statement followed by return
   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  // checkMinMaxPattern checks if an if statement implements min/max logic.
   396  // ifStmt: the if statement to check
   397  // falseResult: the expression returned when the condition is false
   398  // funcName: "min" or "max"
   399  // param0, param1: the two parameter names for the function.
   400  func checkMinMaxPattern(ifStmt *ast.IfStmt, falseResult ast.Expr, funcName, param0, param1 string) bool {
   401  	// Must have condition with comparison
   402  	cmp, ok := ifStmt.Cond.(*ast.BinaryExpr)
   403  	if !ok {
   404  		return false
   405  	}
   406  
   407  	// Check if then branch returns one of the compared values
   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  	// Use the same logic as the existing minmax analyzer
   418  	sign := isInequality(cmp.Op)
   419  	if sign == 0 {
   420  		return false // Not a comparison operator
   421  	}
   422  
   423  	t := thenRet.Results[0]     // "true" result
   424  	f := falseResult            // "false" result
   425  	x, ok := cmp.X.(*ast.Ident) // left operand
   426  	if !ok {
   427  		return false // Not a basic min/max comparison
   428  	}
   429  	y, ok := cmp.Y.(*ast.Ident) // right operand
   430  	if !ok {
   431  		return false // Not a basic min/max comparison
   432  	}
   433  
   434  	// Check that the min max algorithm uses the function's params
   435  	// Which param corresponds to which part of the operation doesn't matter,
   436  	// so we have to try both.
   437  	if !(param0 == x.Name && param1 == y.Name ||
   438  		param0 == y.Name && param1 == x.Name) {
   439  		return false
   440  	}
   441  
   442  	// Check operand order and adjust sign accordingly
   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  	// Check if the sign matches the function name
   452  	return cond(sign < 0, "min", "max") == funcName
   453  }
   454  

View as plain text