1
2
3
4
5 package lostcancel
6
7 import (
8 _ "embed"
9 "fmt"
10 "go/ast"
11 "go/types"
12
13 "golang.org/x/tools/go/analysis"
14 "golang.org/x/tools/go/analysis/passes/ctrlflow"
15 "golang.org/x/tools/go/analysis/passes/inspect"
16 "golang.org/x/tools/go/ast/inspector"
17 "golang.org/x/tools/go/cfg"
18 "golang.org/x/tools/internal/analysis/analyzerutil"
19 "golang.org/x/tools/internal/typesinternal"
20 )
21
22
23 var doc string
24
25 var Analyzer = &analysis.Analyzer{
26 Name: "lostcancel",
27 Doc: analyzerutil.MustExtractDoc(doc, "lostcancel"),
28 URL: "https://pkg.go.dev/golang.org/x/tools/go/analysis/passes/lostcancel",
29 Run: run,
30 Requires: []*analysis.Analyzer{
31 inspect.Analyzer,
32 ctrlflow.Analyzer,
33 },
34 }
35
36 const debug = false
37
38 var contextPackage = "context"
39
40
41
42
43
44
45
46
47
48
49
50 func run(pass *analysis.Pass) (any, error) {
51
52 if !typesinternal.Imports(pass.Pkg, contextPackage) {
53 return nil, nil
54 }
55
56
57 inspect := pass.ResultOf[inspect.Analyzer].(*inspector.Inspector)
58 nodeTypes := []ast.Node{
59 (*ast.FuncLit)(nil),
60 (*ast.FuncDecl)(nil),
61 }
62 inspect.Preorder(nodeTypes, func(n ast.Node) {
63 runFunc(pass, n)
64 })
65 return nil, nil
66 }
67
68 func runFunc(pass *analysis.Pass, node ast.Node) {
69
70 var funcScope *types.Scope
71 switch v := node.(type) {
72 case *ast.FuncLit:
73 funcScope = pass.TypesInfo.Scopes[v.Type]
74 case *ast.FuncDecl:
75 funcScope = pass.TypesInfo.Scopes[v.Type]
76 }
77
78
79 cancelvars := make(map[*types.Var]ast.Node)
80
81
82
83
84
85
86 ast.PreorderStack(node, nil, func(n ast.Node, stack []ast.Node) bool {
87 if _, ok := n.(*ast.FuncLit); ok && len(stack) > 0 {
88 return false
89 }
90
91
92
93
94
95
96
97 if !isContextWithCancel(pass.TypesInfo, n) || !isCall(stack[len(stack)-1]) {
98 return true
99 }
100 var id *ast.Ident
101 stmt := stack[len(stack)-2]
102 switch stmt := stmt.(type) {
103 case *ast.ValueSpec:
104 if len(stmt.Names) > 1 {
105 id = stmt.Names[1]
106 }
107 case *ast.AssignStmt:
108 if len(stmt.Lhs) > 1 {
109 id, _ = stmt.Lhs[1].(*ast.Ident)
110 }
111 }
112 if id != nil {
113 if id.Name == "_" {
114 pass.ReportRangef(id,
115 "the cancel function returned by context.%s should be called, not discarded, to avoid a context leak",
116 n.(*ast.SelectorExpr).Sel.Name)
117 } else if v, ok := pass.TypesInfo.Uses[id].(*types.Var); ok {
118
119
120 if funcScope.Contains(v.Pos()) {
121 cancelvars[v] = stmt
122 }
123 } else if v, ok := pass.TypesInfo.Defs[id].(*types.Var); ok {
124 cancelvars[v] = stmt
125 }
126 }
127 return true
128 })
129
130 if len(cancelvars) == 0 {
131 return
132 }
133
134
135 cfgs := pass.ResultOf[ctrlflow.Analyzer].(*ctrlflow.CFGs)
136 var g *cfg.CFG
137 var sig *types.Signature
138 switch node := node.(type) {
139 case *ast.FuncDecl:
140 sig, _ = pass.TypesInfo.Defs[node.Name].Type().(*types.Signature)
141 if node.Name.Name == "main" && sig.Recv() == nil && pass.Pkg.Name() == "main" {
142
143
144 return
145 }
146 g = cfgs.FuncDecl(node)
147
148 case *ast.FuncLit:
149 sig, _ = pass.TypesInfo.Types[node.Type].Type.(*types.Signature)
150 g = cfgs.FuncLit(node)
151 }
152 if sig == nil {
153 return
154 }
155
156
157 if debug {
158 fmt.Println(g.Format(pass.Fset))
159 }
160
161
162
163
164 for v, stmt := range cancelvars {
165 if ret := lostCancelPath(pass, g, v, stmt, sig); ret != nil {
166 lineno := pass.Fset.Position(stmt.Pos()).Line
167 pass.ReportRangef(stmt, "the %s function is not used on all paths (possible context leak)", v.Name())
168
169 pos, end := ret.Pos(), ret.End()
170
171
172 if pass.Fset.File(pos) != pass.Fset.File(end) {
173 end = pos
174 }
175 pass.Report(analysis.Diagnostic{
176 Pos: pos,
177 End: end,
178 Message: fmt.Sprintf("this return statement may be reached without using the %s var defined on line %d", v.Name(), lineno),
179 })
180 }
181 }
182 }
183
184 func isCall(n ast.Node) bool { _, ok := n.(*ast.CallExpr); return ok }
185
186
187
188 func isContextWithCancel(info *types.Info, n ast.Node) bool {
189 sel, ok := n.(*ast.SelectorExpr)
190 if !ok {
191 return false
192 }
193 switch sel.Sel.Name {
194 case "WithCancel", "WithCancelCause",
195 "WithTimeout", "WithTimeoutCause",
196 "WithDeadline", "WithDeadlineCause":
197 default:
198 return false
199 }
200 if x, ok := sel.X.(*ast.Ident); ok {
201 if pkgname, ok := info.Uses[x].(*types.PkgName); ok {
202 return pkgname.Imported().Path() == contextPackage
203 }
204
205
206 return x.Name == "context"
207 }
208 return false
209 }
210
211
212
213
214
215 func lostCancelPath(pass *analysis.Pass, g *cfg.CFG, v *types.Var, stmt ast.Node, sig *types.Signature) *ast.ReturnStmt {
216 vIsNamedResult := sig != nil && tupleContains(sig.Results(), v)
217
218
219 uses := func(pass *analysis.Pass, v *types.Var, stmts []ast.Node) bool {
220 found := false
221 for _, stmt := range stmts {
222 ast.Inspect(stmt, func(n ast.Node) bool {
223 switch n := n.(type) {
224 case *ast.Ident:
225 if pass.TypesInfo.Uses[n] == v {
226 found = true
227 }
228 case *ast.ReturnStmt:
229
230
231 if n.Results == nil && vIsNamedResult {
232 found = true
233 }
234 }
235 return !found
236 })
237 }
238 return found
239 }
240
241
242 memo := make(map[*cfg.Block]bool)
243 blockUses := func(pass *analysis.Pass, v *types.Var, b *cfg.Block) bool {
244 res, ok := memo[b]
245 if !ok {
246 res = uses(pass, v, b.Nodes)
247 memo[b] = res
248 }
249 return res
250 }
251
252
253
254 var defblock *cfg.Block
255 var rest []ast.Node
256 outer:
257 for _, b := range g.Blocks {
258 for i, n := range b.Nodes {
259 if n == stmt {
260 defblock = b
261 rest = b.Nodes[i+1:]
262 break outer
263 }
264 }
265 }
266 if defblock == nil {
267 panic("internal error: can't find defining block for cancel var")
268 }
269
270
271 if uses(pass, v, rest) {
272 return nil
273 }
274
275
276 if ret := defblock.Return(); ret != nil {
277 return ret
278 }
279
280
281
282 seen := make(map[*cfg.Block]bool)
283 var search func(blocks []*cfg.Block) *ast.ReturnStmt
284 search = func(blocks []*cfg.Block) *ast.ReturnStmt {
285 for _, b := range blocks {
286 if seen[b] {
287 continue
288 }
289 seen[b] = true
290
291
292 if blockUses(pass, v, b) {
293 continue
294 }
295
296
297 if ret := b.Return(); ret != nil {
298 if debug {
299 fmt.Printf("found path to return in block %s\n", b)
300 }
301 return ret
302 }
303
304
305 if ret := search(b.Succs); ret != nil {
306 if debug {
307 fmt.Printf(" from block %s\n", b)
308 }
309 return ret
310 }
311 }
312 return nil
313 }
314 return search(defblock.Succs)
315 }
316
317 func tupleContains(tuple *types.Tuple, v *types.Var) bool {
318 for v0 := range tuple.Variables() {
319 if v0 == v {
320 return true
321 }
322 }
323 return false
324 }
325
View as plain text