1
2
3
4
5 package modernize
6
7
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
47
48 expr := call.Args[0]
49
50
51
52
53 if info.Types[expr].IsNil() {
54 continue
55 }
56
57 if !typesinternal.NoEffects(info, expr) {
58 continue
59 }
60
61 t := info.TypeOf(expr)
62 var edits []analysis.TextEdit
63
64
65
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)
70 obj := typeutil.Callee(info, call2)
71 if typesinternal.IsMethodNamed(obj, "reflect", "Type", "Elem") {
72
73
74
75 if typ, hasElem := t.(interface{ Elem() types.Type }); hasElem {
76
77 t = typ.Elem()
78 edits = []analysis.TextEdit{{
79 Pos: call.End(),
80 End: call2.End(),
81 }}
82 }
83 }
84 }
85 }
86
87
88
89
90 if types.IsInterface(t) && edits == nil {
91 continue
92 }
93
94
95
96
97
98
99
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
107 }
108
109
110
111
112
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
119 }
120
121
122
123
124 if isComplicatedType(t) {
125 continue
126 }
127
128
129
130
131
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
144
145
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
154 {
155 Pos: call.Lparen + 1,
156 End: call.Rparen,
157 },
158 }, edits...),
159 }},
160 })
161 }
162
163 return nil, nil
164 }
165
166
167
168
169
170
171
172
173
174
175
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
185
186 default:
187 return true
188 }
189 }
190 return false
191 }
192
193
194
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
208
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
226 return true
227 }
228 }
229
230 return check(t)
231 }
232
View as plain text