1
2
3
4
5 package modernize
6
7 import (
8 "bytes"
9 "fmt"
10 "go/ast"
11 "go/token"
12 "go/types"
13 "slices"
14 "strings"
15
16 "golang.org/x/tools/go/analysis"
17 "golang.org/x/tools/go/analysis/passes/inspect"
18 "golang.org/x/tools/go/ast/edge"
19 "golang.org/x/tools/go/ast/inspector"
20 "golang.org/x/tools/internal/analysis/analyzerutil"
21 typeindexanalyzer "golang.org/x/tools/internal/analysis/typeindex"
22 "golang.org/x/tools/internal/astutil"
23 "golang.org/x/tools/internal/moreiters"
24 "golang.org/x/tools/internal/typesinternal/typeindex"
25 "golang.org/x/tools/internal/versions"
26 )
27
28 var EmbedLitAnalyzer = &analysis.Analyzer{
29 Name: "embedlit",
30 Doc: analyzerutil.MustExtractDoc(doc, "embedlit"),
31 Requires: []*analysis.Analyzer{
32 inspect.Analyzer,
33 typeindexanalyzer.Analyzer,
34 },
35 Run: runEmbedLit,
36 URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/modernize#embedlit",
37 }
38
39
40
41
42
43
44 func runEmbedLit(pass *analysis.Pass) (any, error) {
45 var (
46 inspect = pass.ResultOf[inspect.Analyzer].(*inspector.Inspector)
47 index = pass.ResultOf[typeindexanalyzer.Analyzer].(*typeindex.Index)
48 info = pass.TypesInfo
49 )
50 for curLit := range inspect.Root().Preorder((*ast.CompositeLit)(nil)) {
51 if curLit.ParentEdgeKind() != edge.KeyValueExpr_Value {
52
53
54
55
56 if !embedlitUnnest(pass, info, curLit) {
57 err := embedlitCombine(pass, index, info, curLit)
58 if err != nil {
59 return nil, err
60 }
61 }
62 }
63 }
64 return nil, nil
65 }
66
67
68
69
70
71 func embedlitUnnest(pass *analysis.Pass, info *types.Info, curLit inspector.Cursor) bool {
72 var (
73 edits []analysis.TextEdit
74 names []string
75 lit = curLit.Node().(*ast.CompositeLit)
76 compLitType = info.TypeOf(lit)
77 )
78
79
80
81 var checkLit func(lit *ast.CompositeLit)
82 checkLit = func(lit *ast.CompositeLit) {
83 for i, elt := range lit.Elts {
84
85 if kv, ok := elt.(*ast.KeyValueExpr); ok {
86 if innerLit := isEmbeddedFieldLit(info, compLitType, kv); innerLit != nil {
87
88
89
90
91
92 closingPos := innerLit.Elts[len(innerLit.Elts)-1].End()
93 file := astutil.EnclosingFile(curLit)
94
95 if !analyzerutil.FileUsesGoVersion(pass, file, versions.Go1_27) {
96 return
97 }
98
99 if !moreiters.Empty(astutil.Comments(file, kv.Pos(), innerLit.Lbrace+1)) ||
100 !moreiters.Empty(astutil.Comments(file, closingPos, innerLit.Rbrace+1)) {
101 continue
102 }
103
104
105
106
107 startPos := kv.Pos()
108 endPos := innerLit.Lbrace + 1
109
110
111
112
113
114
115
116
117
118
119 tokFile := pass.Fset.File(kv.Pos())
120 lineOf := func(pos token.Pos) int {
121 return tokFile.PositionFor(pos, false).Line
122 }
123 curLine := lineOf(kv.Pos())
124 var prevLine int
125 if i == 0 {
126
127 prevLine = lineOf(lit.Lbrace)
128 } else {
129 prevLine = lineOf(lit.Elts[i-1].End())
130 }
131
132
133
134
135
136 if prevLine < curLine && curLine < tokFile.LineCount() &&
137 lineOf(innerLit.Elts[0].Pos()) > curLine {
138 lineStart := tokFile.LineStart(curLine)
139 nextLineStart := tokFile.LineStart(curLine + 1)
140
141 if moreiters.Empty(astutil.Comments(file, lineStart, nextLineStart)) {
142 startPos = tokFile.LineStart(curLine)
143 endPos = nextLineStart
144 }
145 }
146
147 edits = append(edits, []analysis.TextEdit{
148
149
150 {
151
152 Pos: startPos,
153 End: endPos,
154 },
155 {
156
157
158
159 Pos: closingPos,
160 End: innerLit.Rbrace + 1,
161 },
162 }...)
163 names = append(names, kv.Key.(*ast.Ident).Name)
164 checkLit(innerLit)
165 }
166 }
167 }
168 }
169 checkLit(lit)
170 if len(edits) > 0 {
171 pass.Report(analysis.Diagnostic{
172 Pos: curLit.Node().Pos(),
173 End: curLit.Node().End(),
174 Message: "embedded field type can be removed from struct literal",
175 SuggestedFixes: []analysis.SuggestedFix{
176 {
177 Message: fmt.Sprintf("Remove embedded field type%s %s", cond(len(names) == 1, "", "s"), strings.Join(names, ", ")),
178 TextEdits: edits,
179 },
180 },
181 })
182 return true
183 }
184 return false
185 }
186
187
188
189
190
191 func embedlitCombine(pass *analysis.Pass, index *typeindex.Index, info *types.Info, curLit inspector.Cursor) error {
192 compLit := curLit.Node().(*ast.CompositeLit)
193 if !moreiters.Every(slices.Values(compLit.Elts), func(e ast.Expr) bool {
194 return is[*ast.KeyValueExpr](e)
195 }) {
196
197
198 return nil
199 }
200 var (
201
202 lhs *ast.Ident
203
204
205
206 curStmt inspector.Cursor
207 )
208 switch curLit.ParentEdgeKind() {
209 case edge.AssignStmt_Rhs:
210 assign := curLit.Parent().Node().(*ast.AssignStmt)
211
212
213 if len(assign.Lhs) != 1 {
214 return nil
215 }
216 if id, ok := assign.Lhs[0].(*ast.Ident); ok {
217 lhs = id
218 curStmt = curLit.Parent()
219 }
220 case edge.ValueSpec_Values:
221 spec := curLit.Parent().Node().(*ast.ValueSpec)
222
223 if len(spec.Names) != 1 {
224 return nil
225 }
226 lhs = spec.Names[0]
227 if decl, ok := moreiters.First(curLit.Enclosing((*ast.DeclStmt)(nil))); ok {
228 if gdecl, ok := decl.Node().(*ast.DeclStmt).Decl.(*ast.GenDecl); ok && len(gdecl.Specs) == 1 {
229 curStmt = decl
230 }
231 }
232 default:
233 return nil
234 }
235
236 if lhs == nil || !curStmt.Valid() {
237 return nil
238 }
239
240 var (
241 compLitType = info.TypeOf(compLit)
242 tObj = info.ObjectOf(lhs)
243
244
245 firstStmt, lastStmt inspector.Cursor
246 hasEmbeddedSelection bool
247 )
248 if compLitType == nil {
249 return nil
250 }
251
252
253
254
255
256 var fieldPaths [][]int
257 for _, elt := range compLit.Elts {
258 k, ok := elt.(*ast.KeyValueExpr).Key.(*ast.Ident)
259 if !ok {
260 return nil
261 }
262 _, idx, _ := types.LookupFieldOrMethod(compLitType, true, pass.Pkg, k.Name)
263 if len(idx) == 0 {
264 return nil
265 }
266 fieldPaths = append(fieldPaths, idx)
267 }
268
269 stmtloop:
270 for {
271 var ok bool
272 curStmt, ok = curStmt.NextSibling()
273 if !ok {
274 break
275 }
276
277
278 assign, ok := curStmt.Node().(*ast.AssignStmt)
279 if !ok || len(assign.Lhs) != 1 || !(assign.Tok == token.ASSIGN || assign.Tok == token.DEFINE) {
280
281 break
282 }
283 expr := assign.Lhs[0]
284 sel, ok := expr.(*ast.SelectorExpr)
285 if !ok {
286 break
287 }
288
289 selXId, ok := sel.X.(*ast.Ident)
290 if !ok {
291
292 break
293 }
294 obj := info.ObjectOf(selXId)
295 if obj != tObj {
296 break
297 }
298 fieldObj, assignIdx, indirect := types.LookupFieldOrMethod(compLitType, true, pass.Pkg, sel.Sel.Name)
299 fieldVar, ok := fieldObj.(*types.Var)
300 if !ok || len(assignIdx) == 0 || indirect {
301 break
302 }
303
304
305 if slices.ContainsFunc(fieldPaths, func(index []int) bool { return pathConflicts(index, assignIdx) }) {
306 break
307 }
308
309
310
311 if fieldVar.Embedded() || len(assignIdx) > 1 {
312 hasEmbeddedSelection = true
313 }
314
315 rhsCur := curStmt.ChildAt(edge.AssignStmt_Rhs, 0)
316 if uses(index, rhsCur, tObj) {
317 break
318 }
319 for c := range rhsCur.Preorder((*ast.Ident)(nil)) {
320 id := c.Node().(*ast.Ident)
321
322
323 if info.ObjectOf(id) == tObj {
324 break stmtloop
325 }
326
327
328
329
330 }
331
332
333
334 fieldPaths = append(fieldPaths, assignIdx)
335 if !firstStmt.Valid() {
336 firstStmt = curStmt
337 }
338 lastStmt = curStmt
339 }
340
341 if !firstStmt.Valid() || !hasEmbeddedSelection {
342
343 return nil
344 }
345
346 file := astutil.EnclosingFile(curLit)
347
348 if !analyzerutil.FileUsesGoVersion(pass, file, versions.Go1_27) {
349 return nil
350 }
351
352
353
354 tokFile := pass.Fset.File(compLit.Rbrace)
355 filename := tokFile.Name()
356 src, err := pass.ReadFile(filename)
357 if err != nil {
358 return err
359 }
360
361 hasTrailingComma := false
362 if len(compLit.Elts) > 0 {
363 lastElt := compLit.Elts[len(compLit.Elts)-1]
364 lastEltOffset := tokFile.Offset(lastElt.End())
365 rbraceOffset := tokFile.Offset(compLit.Rbrace)
366 span := bytes.Clone(src[lastEltOffset:rbraceOffset])
367
368
369 for co := range astutil.Comments(file, lastElt.End(), compLit.Rbrace) {
370 start := max(tokFile.Offset(co.Pos())-lastEltOffset, 0)
371 end := min(tokFile.Offset(co.End())-lastEltOffset, len(span))
372 if start < end {
373 clear(span[start:end])
374 }
375 }
376 hasTrailingComma = bytes.Contains(span, []byte(","))
377 }
378 var edits []analysis.TextEdit
379
380
381
382
383
384
385
386
387
388
389
390 if len(compLit.Elts) > 0 && !hasTrailingComma {
391 edits = append(edits, analysis.TextEdit{
392 Pos: compLit.Rbrace,
393 End: compLit.Rbrace + 1,
394 NewText: []byte(","),
395 })
396 } else {
397 edits = append(edits, analysis.TextEdit{
398 Pos: compLit.Rbrace,
399 End: compLit.Rbrace + 1,
400 })
401 }
402
403
404
405
406
407 curStmt = firstStmt
408 var prevStmt inspector.Cursor
409 for {
410 assign := curStmt.Node().(*ast.AssignStmt)
411 expr := assign.Lhs[0]
412 sel := expr.(*ast.SelectorExpr)
413
414 edits = append(edits, analysis.TextEdit{
415 Pos: assign.Pos(),
416 End: sel.Sel.Pos(),
417 })
418
419 edits = append(edits, analysis.TextEdit{
420 Pos: expr.End(),
421 End: assign.TokPos + 1,
422 NewText: []byte(":"),
423 })
424
425
426 if prevStmt.Valid() {
427 edits = append(edits, analysis.TextEdit{
428 Pos: prevStmt.Node().End(),
429 NewText: []byte(","),
430 })
431 }
432
433
434 if curStmt == lastStmt {
435 edits = append(edits, analysis.TextEdit{
436 Pos: assign.End(),
437 NewText: []byte("}"),
438 })
439 break
440 }
441 prevStmt = curStmt
442 curStmt, _ = curStmt.NextSibling()
443 }
444
445 pass.Report(analysis.Diagnostic{
446 Pos: curLit.Node().Pos(),
447 End: curLit.Node().End(),
448 Message: "embedded field assignment can be moved to struct literal",
449 SuggestedFixes: []analysis.SuggestedFix{
450 {
451 Message: "Move embedded field assignment to struct literal",
452 TextEdits: edits,
453 },
454 },
455 })
456 return nil
457 }
458
459
460
461
462
463 func isEmbeddedFieldLit(info *types.Info, topLevelType types.Type, kv *ast.KeyValueExpr) *ast.CompositeLit {
464 obj := keyedField(info, kv)
465 if obj == nil || !obj.Embedded() {
466 return nil
467 }
468 lit, ok := kv.Value.(*ast.CompositeLit)
469 if !ok || len(lit.Elts) == 0 {
470
471 return nil
472 }
473
474
475
476 for _, elt := range lit.Elts {
477 kv, ok := elt.(*ast.KeyValueExpr)
478 if !ok {
479 return nil
480 }
481 obj := keyedField(info, kv)
482 if obj == nil {
483 return nil
484 }
485 k := kv.Key.(*ast.Ident)
486
487
488
489
490
491
492
493 parentObj, _, _ := types.LookupFieldOrMethod(topLevelType, true, obj.Pkg(), k.Name)
494 if parentObj != obj {
495 return nil
496 }
497 }
498 return lit
499 }
500
501
502
503 func keyedField(info *types.Info, kv *ast.KeyValueExpr) *types.Var {
504 k, ok := kv.Key.(*ast.Ident)
505 if !ok {
506 return nil
507 }
508 obj, ok := info.ObjectOf(k).(*types.Var)
509 if !ok || !obj.IsField() {
510 return nil
511 }
512 return obj
513 }
514
515
516
517 func pathConflicts(p1, p2 []int) bool {
518 return (len(p1) >= len(p2) && slices.Equal(p1[:len(p2)], p2)) ||
519 (len(p2) >= len(p1) && slices.Equal(p2[:len(p1)], p1))
520 }
521
View as plain text