Source file src/cmd/vendor/golang.org/x/tools/go/analysis/passes/modernize/reflect.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  // This file defines modernizers that use the "reflect" package.
     8  
     9  import (
    10  	"go/ast"
    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/edge"
    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/typesinternal"
    21  	"golang.org/x/tools/internal/typesinternal/typeindex"
    22  	"golang.org/x/tools/internal/versions"
    23  )
    24  
    25  var ReflectTypeForAnalyzer = &analysis.Analyzer{
    26  	Name: "reflecttypefor",
    27  	Doc:  analyzerutil.MustExtractDoc(doc, "reflecttypefor"),
    28  	Requires: []*analysis.Analyzer{
    29  		inspect.Analyzer,
    30  		typeindexanalyzer.Analyzer,
    31  	},
    32  	Run: reflecttypefor,
    33  	URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/modernize#reflecttypefor",
    34  }
    35  
    36  func reflecttypefor(pass *analysis.Pass) (any, error) {
    37  	var (
    38  		index = pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
    39  		info  = pass.TypesInfo
    40  
    41  		reflectTypeOf = index.Object("reflect", "TypeOf")
    42  	)
    43  
    44  	for curCall := range index.Calls(reflectTypeOf) {
    45  		call := curCall.Node().(*ast.CallExpr)
    46  		// Have: reflect.TypeOf(expr)
    47  
    48  		expr := call.Args[0]
    49  
    50  		// reflect.TypeFor cannot be instantiated with an untyped nil.
    51  		// We use type information rather than checking the identifier name
    52  		// to correctly handle edge cases where "nil" is shadowed (e.g. nil := "nil").
    53  		if info.Types[expr].IsNil() {
    54  			continue
    55  		}
    56  
    57  		if !typesinternal.NoEffects(info, expr) {
    58  			continue // don't eliminate operand: may have effects
    59  		}
    60  
    61  		t := info.TypeOf(expr)
    62  		var edits []analysis.TextEdit
    63  
    64  		// Special cases for TypeOf((*T)(nil)).Elem(), and
    65  		// TypeOf([]T(nil)).Elem(), needed when T is an interface type.
    66  		if curCall.ParentEdgeKind() == edge.SelectorExpr_X {
    67  			curSel := astutil.UnparenEnclosingCursor(curCall).Parent()
    68  			if curSel.ParentEdgeKind() == edge.CallExpr_Fun {
    69  				call2 := astutil.UnparenEnclosingCursor(curSel).Parent().Node().(*ast.CallExpr) // potentially .Elem()
    70  				obj := typeutil.Callee(info, call2)
    71  				if typesinternal.IsMethodNamed(obj, "reflect", "Type", "Elem") {
    72  					// reflect.TypeOf(expr).Elem()
    73  					//                     -------
    74  					// reflect.TypeOf(expr)
    75  					if typ, hasElem := t.(interface{ Elem() types.Type }); hasElem {
    76  						// Have: TypeOf(expr).Elem() where expr is *T, []T, [k]T, chan T, map[K]T, etc.
    77  						t = typ.Elem()
    78  						edits = []analysis.TextEdit{{
    79  							Pos: call.End(),
    80  							End: call2.End(),
    81  						}}
    82  					}
    83  				}
    84  			}
    85  		}
    86  
    87  		// TypeOf(x) where x has an interface type is a
    88  		// dynamic operation; don't transform it to TypeFor.
    89  		// (edits == nil means "not the Elem() special case".)
    90  		if types.IsInterface(t) && edits == nil {
    91  			continue
    92  		}
    93  
    94  		// Don't offer the fix if it would erase a reference to a
    95  		// (non-type) symbol, as this may break intended coupling.
    96  		// Examples:
    97  		//   TypeOf(0)       -> TypeFor[int]()  // ok
    98  		//   TypeOf(pkg.Var) -> TypeFor[int]()  // bad: loses connection to pkg.Var
    99  		//   TypeOf(uint(0)) -> TypeFor[uint]() // ok: (most) type symbols are preserved
   100  		if usesNonTypeSymbol(info, expr) {
   101  			continue
   102  		}
   103  
   104  		file := astutil.EnclosingFile(curCall)
   105  		if !analyzerutil.FileUsesGoVersion(pass, file, versions.Go1_22) {
   106  			continue // TypeFor requires go1.22
   107  		}
   108  
   109  		// Format the type as valid Go syntax.
   110  		// TODO(adonovan): FileQualifier needs to respect
   111  		// visibility at the current point, and either fail
   112  		// or edit the imports as needed.
   113  		qual := typesinternal.FileQualifier(file, pass.Pkg)
   114  		tstr := types.TypeString(t, qual)
   115  
   116  		sel, ok := call.Fun.(*ast.SelectorExpr)
   117  		if !ok {
   118  			continue // e.g. reflect was dot-imported
   119  		}
   120  
   121  		// Don't offer a fix if the type contains an unnamed struct or unnamed
   122  		// interface because the replacement would be significantly more verbose.
   123  		// (See golang/go#76698)
   124  		if isComplicatedType(t) {
   125  			continue
   126  		}
   127  
   128  		// Don't offer the fix if the type string is too long. We define "too
   129  		// long" as more than three times the length of the original expression
   130  		// and at least 16 characters (a 3x length increase of a very
   131  		// short expression should not be cause for skipping the fix).
   132  		oldLen := int(expr.End() - expr.Pos())
   133  		newLen := len(tstr)
   134  		if newLen >= 16 && newLen > 3*oldLen {
   135  			continue
   136  		}
   137  
   138  		pass.Report(analysis.Diagnostic{
   139  			Pos:     call.Fun.Pos(),
   140  			End:     call.Fun.End(),
   141  			Message: "reflect.TypeOf call can be simplified using TypeFor",
   142  			SuggestedFixes: []analysis.SuggestedFix{{
   143  				// reflect.TypeOf    (...T value...)
   144  				//         ------     -------------
   145  				// reflect.TypeFor[T](             )
   146  				Message: "Replace TypeOf by TypeFor",
   147  				TextEdits: append([]analysis.TextEdit{
   148  					{
   149  						Pos:     sel.Sel.Pos(),
   150  						End:     sel.Sel.End(),
   151  						NewText: []byte("TypeFor[" + tstr + "]"),
   152  					},
   153  					// delete (pure) argument
   154  					{
   155  						Pos: call.Lparen + 1,
   156  						End: call.Rparen,
   157  					},
   158  				}, edits...),
   159  			}},
   160  		})
   161  	}
   162  
   163  	return nil, nil
   164  }
   165  
   166  // usesNonTypeSymbol reports whether expr uses a non-type symbol:
   167  // a value-level object (a var, const, or func) or any other named entity
   168  // whose identifier would disappear in a TypeFor replacement. We suppress
   169  // the fix in that case so the rewrite does not erase a symbol the author
   170  // named on purpose, for example reflect.TypeOf(f), reflect.TypeOf(x.Field),
   171  // or a named array length such as [arrayLen]byte.
   172  //
   173  // Type names, package names, nil, and builtins are not such symbols: they
   174  // either reappear in the type argument (e.g. T in reflect.TypeOf(T{})) or
   175  // are irrelevant to it, so the rewrite is allowed.
   176  func usesNonTypeSymbol(info *types.Info, expr ast.Expr) bool {
   177  	for n := range ast.Preorder(expr) {
   178  		id, ok := n.(*ast.Ident)
   179  		if !ok {
   180  			continue
   181  		}
   182  		switch info.Uses[id].(type) {
   183  		case *types.TypeName, *types.PkgName, *types.Nil, *types.Builtin:
   184  			// Type-level names reappear in (or are irrelevant to) the
   185  			// type argument, so they may be erased safely.
   186  		default:
   187  			return true
   188  		}
   189  	}
   190  	return false
   191  }
   192  
   193  // isComplicatedType reports whether type t is complicated, e.g. it is or contains an
   194  // unnamed struct, interface, or function signature.
   195  func isComplicatedType(t types.Type) bool {
   196  	var check func(typ types.Type) bool
   197  	check = func(typ types.Type) bool {
   198  		switch t := typ.(type) {
   199  		case typesinternal.NamedOrAlias:
   200  			for ta := range t.TypeArgs().Types() {
   201  				if check(ta) {
   202  					return true
   203  				}
   204  			}
   205  			return false
   206  		case *types.Struct, *types.Interface, *types.Signature:
   207  			// These are complex types with potentially many elements
   208  			// so we should avoid duplicating their definition.
   209  			return true
   210  		case *types.Pointer:
   211  			return check(t.Elem())
   212  		case *types.Slice:
   213  			return check(t.Elem())
   214  		case *types.Array:
   215  			return check(t.Elem())
   216  		case *types.Chan:
   217  			return check(t.Elem())
   218  		case *types.Map:
   219  			return check(t.Key()) || check(t.Elem())
   220  		case *types.Basic:
   221  			return false
   222  		case *types.TypeParam:
   223  			return false
   224  		default:
   225  			// Includes types.Union
   226  			return true
   227  		}
   228  	}
   229  
   230  	return check(t)
   231  }
   232  

View as plain text