Source file src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/slicescontains.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  
    13  	"golang.org/x/tools/go/analysis"
    14  	"golang.org/x/tools/go/analysis/passes/inspect"
    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/refactor"
    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 SlicesContainsAnalyzer = &analysis.Analyzer{
    28  	Name: "slicescontains",
    29  	Doc:  analyzerutil.MustExtractDoc(doc, "slicescontains"),
    30  	Requires: []*analysis.Analyzer{
    31  		inspect.Analyzer,
    32  		typeindexanalyzer.Analyzer,
    33  	},
    34  	Run: slicescontains,
    35  	URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/modernize#slicescontains",
    36  }
    37  
    38  // The slicescontains pass identifies loops that can be replaced by a
    39  // call to slices.Contains{,Func}. For example:
    40  //
    41  //	for i, elem := range s {
    42  //		if elem == needle {
    43  //			...
    44  //			break
    45  //		}
    46  //	}
    47  //
    48  // =>
    49  //
    50  //	if slices.Contains(s, needle) { ... }
    51  //
    52  // Variants:
    53  //   - if the if-condition is f(elem), the replacement
    54  //     uses slices.ContainsFunc(s, f).
    55  //   - if the if-body is "return true" and the fallthrough
    56  //     statement is "return false" (or vice versa), the
    57  //     loop becomes "return [!]slices.Contains(...)".
    58  //   - if the if-body is "found = true" and the previous
    59  //     statement is "found = false" (or vice versa), the
    60  //     loop becomes "found = [!]slices.Contains(...)".
    61  //
    62  // It rejects candidates whose needle/predicate expression from the if-statement
    63  // has side effects to avoid changes in program behavior.
    64  func slicescontains(pass *analysis.Pass) (any, error) {
    65  	// Skip the analyzer in packages where its
    66  	// fixes would create an import cycle.
    67  	if within(pass, "slices", "runtime") {
    68  		return nil, nil
    69  	}
    70  
    71  	var (
    72  		index = pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
    73  		info  = pass.TypesInfo
    74  	)
    75  
    76  	// check is called for each RangeStmt of this form:
    77  	//   for i, elem := range s { if cond { ... } }
    78  	check := func(file *ast.File, curRange inspector.Cursor) {
    79  		rng := curRange.Node().(*ast.RangeStmt)
    80  		ifStmt := rng.Body.List[0].(*ast.IfStmt)
    81  
    82  		// isSliceElem reports whether e denotes the
    83  		// current slice element (elem or s[i]).
    84  		isSliceElem := func(e ast.Expr) bool {
    85  			if rng.Value != nil && astutil.EqualSyntax(e, rng.Value) {
    86  				return true // "elem"
    87  			}
    88  			if x, ok := e.(*ast.IndexExpr); ok &&
    89  				astutil.EqualSyntax(x.X, rng.X) &&
    90  				astutil.EqualSyntax(x.Index, rng.Key) {
    91  				return true // "s[i]"
    92  			}
    93  			return false
    94  		}
    95  
    96  		// Examine the condition for one of these forms:
    97  		//
    98  		// - if elem or s[i] == needle  { ... } => Contains
    99  		// - if predicate(s[i] or elem) { ... } => ContainsFunc
   100  		var (
   101  			funcName string   // "Contains" or "ContainsFunc"
   102  			arg2     ast.Expr // second argument to func (needle or predicate)
   103  		)
   104  		switch cond := ifStmt.Cond.(type) {
   105  		case *ast.BinaryExpr:
   106  			if cond.Op == token.EQL {
   107  				var elem ast.Expr
   108  				if isSliceElem(cond.X) {
   109  					funcName = "Contains"
   110  					elem = cond.X
   111  					arg2 = cond.Y // "if elem == needle"
   112  				} else if isSliceElem(cond.Y) {
   113  					funcName = "Contains"
   114  					elem = cond.Y
   115  					arg2 = cond.X // "if needle == elem"
   116  				}
   117  
   118  				// Reject if elem and needle have different types.
   119  				if elem != nil {
   120  					tElem := info.TypeOf(elem)
   121  					tNeedle := info.TypeOf(arg2)
   122  					if !types.Identical(tElem, tNeedle) {
   123  						// Avoid ill-typed slices.Contains([]error, any).
   124  						if !types.AssignableTo(tNeedle, tElem) {
   125  							return
   126  						}
   127  						// TODO(adonovan): relax this check to allow
   128  						//   slices.Contains([]error, error(any)),
   129  						// inserting an explicit widening conversion
   130  						// around the needle.
   131  						return
   132  					}
   133  				}
   134  			}
   135  
   136  		case *ast.CallExpr:
   137  			if len(cond.Args) == 1 &&
   138  				isSliceElem(cond.Args[0]) &&
   139  				typeutil.Callee(info, cond) != nil { // not a conversion
   140  
   141  				// Attempt to get signature
   142  				sig, isSignature := info.TypeOf(cond.Fun).(*types.Signature)
   143  				if isSignature {
   144  					// skip variadic functions
   145  					if sig.Variadic() {
   146  						return
   147  					}
   148  
   149  					// Slice element type must match function parameter type.
   150  					var (
   151  						tElem  = typeparams.CoreType(info.TypeOf(rng.X)).(*types.Slice).Elem()
   152  						tParam = sig.Params().At(0).Type()
   153  					)
   154  					if !types.Identical(tElem, tParam) {
   155  						return
   156  					}
   157  				}
   158  
   159  				funcName = "ContainsFunc"
   160  				arg2 = cond.Fun // "if predicate(elem)"
   161  			}
   162  		}
   163  		if funcName == "" {
   164  			return // not a candidate for Contains{,Func}
   165  		}
   166  
   167  		// body is the "true" body.
   168  		body := ifStmt.Body
   169  		if len(body.List) == 0 {
   170  			// (We could perhaps delete the loop entirely.)
   171  			return
   172  		}
   173  
   174  		// Reject if needle/predicate expression has side effects.
   175  		if !typesinternal.NoEffects(info, arg2) {
   176  			return
   177  		}
   178  
   179  		// Reject if the body, needle or predicate references either range variable.
   180  		usesRangeVar := func(n ast.Node) bool {
   181  			cur, ok := curRange.FindNode(n)
   182  			if !ok {
   183  				panic(fmt.Sprintf("FindNode(%T) failed", n))
   184  			}
   185  			return uses(index, cur, info.Defs[rng.Key.(*ast.Ident)]) ||
   186  				rng.Value != nil && uses(index, cur, info.Defs[rng.Value.(*ast.Ident)])
   187  		}
   188  		if usesRangeVar(body) {
   189  			// Body uses range var "i" or "elem".
   190  			//
   191  			// (The check for "i" could be relaxed when we
   192  			// generalize this to support slices.Index;
   193  			// and the check for "elem" could be relaxed
   194  			// if "elem" can safely be replaced in the
   195  			// body by "needle".)
   196  			return
   197  		}
   198  		if usesRangeVar(arg2) {
   199  			return
   200  		}
   201  
   202  		// Prepare slices.Contains{,Func} call.
   203  		prefix, importEdits := refactor.AddImport(info, file, "slices", "slices", funcName, rng.Pos())
   204  		contains := fmt.Sprintf("%s%s(%s, %s)",
   205  			prefix,
   206  			funcName,
   207  			astutil.Format(pass.Fset, rng.X),
   208  			astutil.Format(pass.Fset, arg2))
   209  
   210  		report := func(edits []analysis.TextEdit) {
   211  			pass.Report(analysis.Diagnostic{
   212  				Pos:     rng.Pos(),
   213  				End:     rng.End(),
   214  				Message: fmt.Sprintf("Loop can be simplified using slices.%s", funcName),
   215  				SuggestedFixes: []analysis.SuggestedFix{{
   216  					Message:   "Replace loop by call to slices." + funcName,
   217  					TextEdits: append(edits, importEdits...),
   218  				}},
   219  			})
   220  		}
   221  
   222  		// Last statement of body must return/break out of the loop.
   223  		//
   224  		// TODO(adonovan): opt:consider avoiding FindNode with new API of form:
   225  		//    curRange.Get(edge.RangeStmt_Body, -1).
   226  		//             Get(edge.BodyStmt_List, 0).
   227  		//             Get(edge.IfStmt_Body)
   228  		curBody, _ := curRange.FindNode(body)
   229  		curLastStmt, _ := curBody.LastChild()
   230  
   231  		// Reject if any statement in the body except the
   232  		// last has a free continuation (continue or break)
   233  		// that might affected by melting down the loop.
   234  		//
   235  		// TODO(adonovan): relax check by analyzing branch target.
   236  		numBodyStmts := 0
   237  		for curBodyStmt := range curBody.Children() {
   238  			numBodyStmts += 1
   239  			if curBodyStmt != curLastStmt {
   240  				for range curBodyStmt.Preorder((*ast.BranchStmt)(nil), (*ast.ReturnStmt)(nil)) {
   241  					return
   242  				}
   243  			}
   244  		}
   245  
   246  		switch lastStmt := curLastStmt.Node().(type) {
   247  		case *ast.ReturnStmt:
   248  			// Have: for ... range seq { if ... { stmts; return x } }
   249  
   250  			// Special case:
   251  			// body={ return true } next="return false"   (or negation)
   252  			// => return [!]slices.Contains(...)
   253  			if curNext, ok := curRange.NextSibling(); ok {
   254  				nextStmt := curNext.Node().(ast.Stmt)
   255  				tval := isReturnTrueOrFalse(info, lastStmt)
   256  				fval := isReturnTrueOrFalse(info, nextStmt)
   257  				if len(body.List) == 1 && tval*fval < 0 {
   258  					//    for ... { if ... { return true/false } }
   259  					// => return [!]slices.Contains(...)
   260  					report([]analysis.TextEdit{
   261  						// Delete the range statement and following space.
   262  						{
   263  							Pos: rng.Pos(),
   264  							End: nextStmt.Pos(),
   265  						},
   266  						// Change return to [!]slices.Contains(...).
   267  						{
   268  							Pos: nextStmt.Pos(),
   269  							End: nextStmt.End(),
   270  							NewText: fmt.Appendf(nil, "return %s%s",
   271  								cond(tval > 0, "", "!"),
   272  								contains),
   273  						},
   274  					})
   275  					return
   276  				}
   277  			}
   278  
   279  			// General case:
   280  			// => if slices.Contains(...) { stmts; return x }
   281  			report([]analysis.TextEdit{
   282  				// Replace "for ... { if ... " with "if slices.Contains(...)".
   283  				{
   284  					Pos:     rng.Pos(),
   285  					End:     ifStmt.Body.Pos(),
   286  					NewText: fmt.Appendf(nil, "if %s ", contains),
   287  				},
   288  				// Delete '}' of range statement and preceding space.
   289  				{
   290  					Pos: ifStmt.Body.End(),
   291  					End: rng.End(),
   292  				},
   293  			})
   294  			return
   295  
   296  		case *ast.BranchStmt:
   297  			if lastStmt.Tok == token.BREAK && lastStmt.Label == nil { // unlabeled break
   298  				// Have: for ... { if ... { stmts; break } }
   299  				if numBodyStmts == 1 {
   300  					// If the only stmt in the body is an unlabeled "break" that
   301  					// will get deleted in the fix, don't suggest a fix, as it
   302  					// produces confusing code:
   303  					//    if slices.Contains(slice, f) {}
   304  					// Explicitly discarding the result isn't much better:
   305  					//    _ = slices.Contains(slice, f) // just for effects
   306  					// See https://go.dev/issue/77677.
   307  					return
   308  				}
   309  				var prevStmt ast.Stmt // previous statement to range (if any)
   310  				if curPrev, ok := curRange.PrevSibling(); ok {
   311  					// If the RangeStmt's previous sibling is a Stmt,
   312  					// the RangeStmt must be among the Body list of
   313  					// a BlockStmt, CauseClause, or CommClause.
   314  					// In all cases, the prevStmt is the immediate
   315  					// predecessor of the RangeStmt during execution.
   316  					//
   317  					// (This is not true for Stmts in general;
   318  					// see [Cursor.Children] and #71074.)
   319  					prevStmt, _ = curPrev.Node().(ast.Stmt)
   320  				}
   321  
   322  				// Special case:
   323  				// prev="lhs = false" body={ lhs = true; break }
   324  				// => lhs = slices.Contains(...) (or its negation)
   325  				if assign, ok := body.List[0].(*ast.AssignStmt); ok &&
   326  					len(body.List) == 2 &&
   327  					assign.Tok == token.ASSIGN &&
   328  					len(assign.Lhs) == 1 &&
   329  					len(assign.Rhs) == 1 {
   330  
   331  					// Have: body={ lhs = rhs; break }
   332  					assignBool := isTrueOrFalse(info, assign.Rhs[0])
   333  					if prevAssign, ok := prevStmt.(*ast.AssignStmt); ok &&
   334  						len(prevAssign.Lhs) == 1 &&
   335  						len(prevAssign.Rhs) == 1 &&
   336  						assignBool != 0 && // non-bool assignments don't apply in this case
   337  						astutil.EqualSyntax(prevAssign.Lhs[0], assign.Lhs[0]) &&
   338  						assignBool == -isTrueOrFalse(info, prevAssign.Rhs[0]) {
   339  
   340  						// Have:
   341  						//    lhs = false
   342  						//    for ... { if ... { lhs = true; break } }
   343  						//  =>
   344  						//    lhs = slices.Contains(...)
   345  						//
   346  						// TODO(adonovan):
   347  						// - support "var lhs bool = false" and variants.
   348  						// - allow the break to be omitted.
   349  						neg := cond(assignBool < 0, "!", "")
   350  						report([]analysis.TextEdit{
   351  							// Replace "rhs" of previous assignment by [!]slices.Contains(...)
   352  							{
   353  								Pos:     prevAssign.Rhs[0].Pos(),
   354  								End:     prevAssign.Rhs[0].End(),
   355  								NewText: []byte(neg + contains),
   356  							},
   357  							// Delete the loop and preceding space.
   358  							{
   359  								Pos: prevAssign.Rhs[0].End(),
   360  								End: rng.End(),
   361  							},
   362  						})
   363  						return
   364  					}
   365  				}
   366  
   367  				// General case:
   368  				//    for ... { if ...        { stmts; break } }
   369  				// => if slices.Contains(...) { stmts        }
   370  				report([]analysis.TextEdit{
   371  					// Replace "for ... { if ... " with "if slices.Contains(...)".
   372  					{
   373  						Pos:     rng.Pos(),
   374  						End:     ifStmt.Body.Pos(),
   375  						NewText: fmt.Appendf(nil, "if %s ", contains),
   376  					},
   377  					// Delete break statement and preceding space.
   378  					{
   379  						Pos: func() token.Pos {
   380  							if len(body.List) > 1 {
   381  								beforeBreak, _ := curLastStmt.PrevSibling()
   382  								return beforeBreak.Node().End()
   383  							}
   384  							return lastStmt.Pos()
   385  						}(),
   386  						End: lastStmt.End(),
   387  					},
   388  					// Delete '}' of range statement and preceding space.
   389  					{
   390  						Pos: ifStmt.Body.End(),
   391  						End: rng.End(),
   392  					},
   393  				})
   394  				return
   395  			}
   396  		}
   397  	}
   398  
   399  	for curFile := range filesUsingGoVersion(pass, versions.Go1_21) {
   400  		file := curFile.Node().(*ast.File)
   401  
   402  		for curRange := range curFile.Preorder((*ast.RangeStmt)(nil)) {
   403  			rng := curRange.Node().(*ast.RangeStmt)
   404  
   405  			if is[*ast.Ident](rng.Key) &&
   406  				rng.Tok == token.DEFINE &&
   407  				len(rng.Body.List) == 1 &&
   408  				is[*types.Slice](typeparams.CoreType(info.TypeOf(rng.X))) {
   409  
   410  				// Have:
   411  				// - for _, elem := range s { S }
   412  				// - for i       := range s { S }
   413  
   414  				if ifStmt, ok := rng.Body.List[0].(*ast.IfStmt); ok &&
   415  					ifStmt.Init == nil && ifStmt.Else == nil {
   416  
   417  					// Have: for i, elem := range s { if cond { ... } }
   418  					check(file, curRange)
   419  				}
   420  			}
   421  		}
   422  	}
   423  	return nil, nil
   424  }
   425  
   426  // -- helpers --
   427  
   428  // isReturnTrueOrFalse returns nonzero if stmt returns true (+1) or false (-1).
   429  func isReturnTrueOrFalse(info *types.Info, stmt ast.Stmt) int {
   430  	if ret, ok := stmt.(*ast.ReturnStmt); ok && len(ret.Results) == 1 {
   431  		return isTrueOrFalse(info, ret.Results[0])
   432  	}
   433  	return 0
   434  }
   435  
   436  // isTrueOrFalse returns nonzero if expr is literally true (+1) or false (-1).
   437  func isTrueOrFalse(info *types.Info, expr ast.Expr) int {
   438  	if id, ok := expr.(*ast.Ident); ok {
   439  		switch info.Uses[id] {
   440  		case builtinTrue:
   441  			return +1
   442  		case builtinFalse:
   443  			return -1
   444  		}
   445  	}
   446  	return 0
   447  }
   448  

View as plain text