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

     1  // Copyright 2025 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/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  // stditeratorsTable records std types that have legacy T.{Len,At}
    35  // iteration methods as well as a newer T.All method that returns an
    36  // iter.Seq.
    37  var stditeratorsTable = [...]struct {
    38  	pkgpath, typename, lenmethod, atmethod, itermethod, elemname string
    39  
    40  	seqn int // 1 or 2 => "for x" or "for _, x"
    41  }{
    42  	// Example: in go/types, (*Tuple).Variables returns an
    43  	// iterator that replaces a loop over (*Tuple).{Len,At}.
    44  	// The loop variable is named "v".
    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  // stditerators suggests fixes to replace loops using Len/At-style
    65  // iterator APIs by a range loop over an iterator. The set of
    66  // participating types and methods is defined by [iteratorsTable].
    67  //
    68  // Pattern:
    69  //
    70  //	for i := 0; i < x.Len(); i++ {
    71  //		use(x.At(i))
    72  //	}
    73  //
    74  // =>
    75  //
    76  //	for elem := range x.All() {
    77  //		use(elem)
    78  //	}
    79  //
    80  // Variant:
    81  //
    82  //	for i := range x.Len() { ... }
    83  //
    84  // Note: Iterators have a dynamic cost. How do we know that
    85  // the user hasn't intentionally chosen not to use an
    86  // iterator for that reason? We don't want to go fix to
    87  // undo optimizations. Do we need a suppression mechanism?
    88  //
    89  // TODO(adonovan): recognize the more complex patterns that
    90  // could make full use of both components of an iter.Seq2, e.g.
    91  //
    92  //	for i := 0; i < v.NumField(); i++ {
    93  //		use(v.Field(i), v.Type().Field(i))
    94  //	}
    95  //
    96  // =>
    97  //
    98  //	for structField, field := range v.Fields() {
    99  //		use(structField, field)
   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  		// Don't offer fixes within the package
   109  		// that defines the iterator in question.
   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  		// chooseName returns an appropriate fresh name
   120  		// for the index variable of the iterator loop
   121  		// whose body is specified.
   122  		//
   123  		// If the loop body starts with
   124  		//
   125  		//     for ... { e := x.At(i); use(e) }
   126  		//
   127  		// or
   128  		//
   129  		//     for ... { if e := x.At(i); cond { use(e) } }
   130  		//
   131  		// then chooseName prefers the name e and additionally
   132  		// returns the var's symbol. We'll transform this to:
   133  		//
   134  		//     for e := range x.Len() { e := e; use(e) }
   135  		//
   136  		// which leaves a redundant assignment that a
   137  		// subsequent 'forvar' pass will eliminate.
   138  		chooseName := func(curBody inspector.Cursor, x ast.Expr, i *types.Var) (string, *types.Var) {
   139  
   140  			// isVarAssign reports whether stmt has the form v := x.At(i)
   141  			// and returns the variable if so.
   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  					// call to x.At(i)?
   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  						// Have: elem := x.At(i)
   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  				// Is body { elem := x.At(i); ... } ?
   165  				if v := isVarAssign(body.List[0]); v != nil {
   166  					return v.Name(), v
   167  				}
   168  
   169  				// Or { if elem := x.At(i); cond { ... } } ?
   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  			// We generate a new name only if the preferred name is already declared here
   179  			// and is used within the loop body.
   180  			name := freshName(info, index, info.Scopes[loop], loop.Pos(), curBody, curBody, token.NoPos, row.elemname)
   181  			return name, nil
   182  		}
   183  
   184  		// Process each call of x.Len().
   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  			// lenSel is "x.Len"
   192  
   193  			var (
   194  				rng      analysis.Range   // where to report diagnostic
   195  				curBody  inspector.Cursor // loop body
   196  				indexVar *types.Var       // old loop index var
   197  				elemVar  *types.Var       // existing "elem := x.At(i)" var, if present
   198  				elem     string           // name for new loop var
   199  				edits    []analysis.TextEdit
   200  			)
   201  
   202  			// Analyze enclosing loop.
   203  			switch curLenCall.ParentEdgeKind() {
   204  			case edge.BinaryExpr_Y:
   205  				// pattern 1: for i := 0; i < x.Len(); i++ { ... }
   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  					// Have: for _; i < x.Len(); _ { ... }
   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  					// Have: for i := 0; i < x.Len(); i++ { ... }.
   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  					//	for i       := 0; i < x.Len(); i++ {
   233  					//          ----       -------  ---  -----
   234  					//	for elem    := range  x.All()      {
   235  					// or   for _, elem := ...
   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  				// pattern 2: for i := range x.Len() { ... }
   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  					// Have: for i := range x.Len() { ... }
   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  					//	for i    := range x.Len() {
   278  					//          ----            ---
   279  					//	for elem := range x.All() {
   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 // no loop of the required form
   297  			}
   298  
   299  			// TODO(adonovan): what about possible
   300  			// modifications of x within the loop?
   301  			// Aliasing seems to make a conservative
   302  			// treatment impossible.
   303  
   304  			// Check that all uses of var i within loop body are x.At(i).
   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 // use is not arg of call
   311  				}
   312  				curAtCall := curUse.Parent()
   313  				atCall := curAtCall.Node().(*ast.CallExpr)
   314  				if typeutil.Callee(info, atCall) != atMethod {
   315  					continue nextCall // use is not arg of call to T.At
   316  				}
   317  				atSel := ast.Unparen(atCall.Fun).(*ast.SelectorExpr)
   318  
   319  				// Check receivers of Len, At calls match (syntactically).
   320  				if !astutil.EqualSyntax(lenSel.X, atSel.X) {
   321  					continue nextCall
   322  				}
   323  
   324  				// At each point of use, check that
   325  				// the fresh variable is not shadowed
   326  				// by an intervening local declaration
   327  				// (or by the idiomatic elemVar optionally
   328  				// found by chooseName).
   329  				if obj := lookup(info, curAtCall, elem); obj != nil && obj != elemVar && obj.Pos() > indexVar.Pos() {
   330  					// (Ideally, instead of giving up, we would
   331  					// embellish the name and try again.)
   332  					continue nextCall
   333  				}
   334  
   335  				// use(x.At(i))
   336  				//     -------
   337  				// use(elem   )
   338  				edits = append(edits, analysis.TextEdit{
   339  					Pos:     atCall.Pos(),
   340  					End:     atCall.End(),
   341  					NewText: []byte(elem),
   342  				})
   343  			}
   344  
   345  			// Check file Go version is new enough for the iterator method.
   346  			// (In the long run, version filters are not highly selective,
   347  			// so there's no need to do them first, especially as this check
   348  			// may be somewhat expensive.)
   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  // -- helpers --
   373  
   374  // methodGoVersion reports the version at which the method
   375  // (pkgpath.recvtype).method appeared in the standard library.
   376  func methodGoVersion(pkgpath, recvtype, method string) (stdlib.Version, error) {
   377  	// TODO(adonovan): opt: this might be inefficient for large packages
   378  	// like go/types. If so, memoize using a map (and kill two birds with
   379  	// one stone by also memoizing the 'within' check above).
   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