返回 CodeWhale
check-blocking-calls-budget.py
根目录 / scripts / check-blocking-calls-budget.py
1 #!/usr/bin/env python3
2 """Ratchet for blocking calls that could land on Tokio workers (#6149).
3
4 Codewhale's convention: async code must not run blocking operations inline.
5 `std::fs`/`thread::sleep` (and friends) are fine inside `spawn_blocking`,
6 on dedicated `std::thread`s, and in synchronous entry points — but every
7 unprotected call site is one careless caller away from parking a runtime
8 worker. This check counts the sites that are NOT already inside a blocking
9 scope (`spawn_blocking`, `spawn_blocking_supervised`, `std::thread::spawn`,
10 `thread::Builder`) or test code, and fails if any file exceeds its recorded
11 budget in `check-blocking-calls-budget.json`.
12
13 Fix the call site — wrap the work in `spawn_blocking` (the established
14 pattern, ~80 sites) or switch to `tokio::fs`/`tokio::time` — or, if the site
15 is genuinely only reachable from synchronous code, acknowledge the debt by
16 raising the file's budget.
17
18 Run `python3 scripts/check-blocking-calls-budget.py --update` to regenerate
19 the budget after removing sites or after an intentional addition.
20 """
21
22 from __future__ import annotations
23
24 import json
25 import re
26 import sys
27 import tomllib
28 from pathlib import Path
29
30 ROOT = Path(__file__).resolve().parents[1]
31 CRATES = ROOT / "crates"
32 BUDGET_PATH = Path(__file__).with_suffix(".json")
33
34 PATTERNS = {
35 "thread_sleep": re.compile(r"\bthread::sleep\s*\("),
36 "std_fs": re.compile(
37 r"\bstd::fs::(?:read|read_to_string|write|create_dir|create_dir_all|"
38 r"remove_file|remove_dir|remove_dir_all|copy|rename|metadata|"
39 r"symlink_metadata|read_dir|canonicalize|exists|set_permissions|"
40 r"hard_link|soft_link|symlink|(?:File|OpenOptions|DirBuilder)(?=\s*::))\b"
41 ),
42 # Method form (`path.canonicalize()`) resolves the path on the calling
43 # thread exactly like `std::fs::canonicalize` (#6522 review).
44 "path_canonicalize": re.compile(r"\.canonicalize\s*\(\s*\)"),
45 }
46
47 ATTR_RE = re.compile(r"#\s*\[([^\]]*)\]")
48 FN_RE = re.compile(
49 r"\b(?:pub(?:\([^)]*\))?\s+)?(?:unsafe\s+)?(?:extern\s+\"[^\"]*\"\s+)?"
50 r"(async\s+)?fn\s+([A-Za-z_][\w]*)"
51 )
52 MOD_RE = re.compile(r"\bmod\s+([A-Za-z_][\w]*)")
53
54 TOKEN_RE = re.compile(
55 r"#\s*\[[^\]]*\]"
56 r"|\bmod\s+\w+"
57 r"|\bimpl\b"
58 r"|\basync\s+move\s*\{"
59 r"|\basync\s*\{"
60 r"|\b(?:pub(?:\([^)]*\))?\s+)?(?:unsafe\s+)?(?:async\s+)?fn\s+\w+"
61 r"|spawn_blocking(?:_supervised)?"
62 r"|thread::spawn"
63 r"|thread::Builder::new"
64 r"|[{}]"
65 )
66
67
68 def strip_comments_and_strings(text: str, *, module_literals: bool = False) -> str:
69 """Blank out comments and string/char literal contents, keeping newlines."""
70 out = list(text)
71 i, n = 0, len(text)
72 line_comment = block_comment = in_str = in_char = in_raw = False
73 block_depth = 0
74 raw_hashes = 0
75 while i < n:
76 c = text[i]
77 if line_comment:
78 if c == "\n":
79 line_comment = False
80 else:
81 out[i] = " "
82 i += 1
83 continue
84 if block_comment:
85 if text[i : i + 2] == "/*":
86 block_depth += 1
87 out[i] = out[i + 1] = " "
88 i += 2
89 continue
90 if text[i : i + 2] == "*/":
91 block_depth -= 1
92 out[i] = out[i + 1] = " "
93 i += 2
94 if block_depth == 0:
95 block_comment = False
96 continue
97 if c != "\n":
98 out[i] = " "
99 i += 1
100 continue
101 if in_str:
102 if c == "\\":
103 out[i] = out[i + 1] = " "
104 i += 2
105 continue
106 if c == '"':
107 in_str = False
108 elif c != "\n":
109 out[i] = " "
110 i += 1
111 continue
112 if in_char:
113 if c == "\\":
114 out[i] = out[i + 1] = " "
115 i += 2
116 continue
117 if c == "'":
118 in_char = False
119 elif c != "\n":
120 out[i] = " "
121 i += 1
122 continue
123 if in_raw:
124 if c == '"' and text[i + 1 : i + 1 + raw_hashes] == "#" * raw_hashes:
125 for j in range(1 + raw_hashes):
126 out[i + j] = " "
127 i += 1 + raw_hashes
128 in_raw = False
129 continue
130 if c != "\n":
131 out[i] = " "
132 i += 1
133 continue
134 if text[i : i + 2] == "//":
135 line_comment = True
136 out[i] = out[i + 1] = " "
137 i += 2
138 continue
139 if text[i : i + 2] == "/*":
140 block_comment = True
141 block_depth = 1
142 out[i] = out[i + 1] = " "
143 i += 2
144 continue
145 if c == "r":
146 # Production counting keeps its baseline masking unchanged. The
147 # module-edge reader additionally recognises zero-hash raw paths.
148 m = re.match(r'r(#*)"' if module_literals else r'r(#+)"', text[i:])
149 if m:
150 raw_hashes = len(m.group(1))
151 in_raw = True
152 for j in range(2 + raw_hashes):
153 out[i + j] = " "
154 i += 2 + raw_hashes
155 continue
156 if c == '"':
157 in_str = True
158 out[i] = " "
159 i += 1
160 continue
161 if c == "'" and re.match(r"'(?:\\.|[^'\\])'", text[i:]):
162 in_char = True
163 out[i] = " "
164 i += 1
165 continue
166 i += 1
167 return "".join(out)
168
169
170 def file_counts(path: Path) -> dict[str, int]:
171 """Count unprotected blocking-call sites in one Rust source file."""
172 text = path.read_text(encoding="utf-8", errors="replace")
173 code = strip_comments_and_strings(text)
174 counts = {name: 0 for name in PATTERNS}
175 # Scope stack: entries are dicts {kind, open_depth} where kind is
176 # 'test', 'blocking', 'fn', 'mod', or 'impl'. A hit counts only when the
177 # innermost enclosing scope is neither test code nor a blocking pool /
178 # dedicated-thread closure.
179 stack: list[dict] = []
180 pending_attr_test = False
181 pending_blocking = False
182 depth = 0
183 for line in code.split("\n"):
184 # Interleave pattern hits and scope tokens in column order so a
185 # one-liner like `fn f() { thread::sleep(..) }` sees the fn scope.
186 events: list[tuple[int, str, object]] = []
187 for name, pat in PATTERNS.items():
188 for m in pat.finditer(line):
189 events.append((m.start(), "hit", name))
190 for m in TOKEN_RE.finditer(line):
191 events.append((m.start(), "tok", m.group(0)))
192 events.sort(key=lambda e: e[0])
193 for _col, kind, payload in events:
194 if kind == "hit":
195 if not any(s["kind"] in ("test", "blocking") for s in stack):
196 counts[payload] += 1 # type: ignore[index]
197 continue
198 tok = payload # type: ignore[assignment]
199 if tok.startswith("#"):
200 inner = tok[tok.index("[") + 1 : -1]
201 if "test" in inner:
202 pending_attr_test = True
203 continue
204 if tok == "{":
205 depth += 1
206 if pending_blocking:
207 stack.append({"kind": "blocking", "open": depth})
208 elif stack and stack[-1]["open"] is None:
209 stack[-1]["open"] = depth
210 pending_blocking = False
211 continue
212 if tok == "}":
213 while stack and stack[-1]["open"] == depth:
214 stack.pop()
215 depth -= 1
216 continue
217 if "spawn_blocking" in tok or tok in ("thread::spawn", "thread::Builder::new"):
218 pending_blocking = True
219 continue
220 if tok.startswith("async") and tok.endswith("{"):
221 stack.append({"kind": "fn", "open": depth + 1})
222 depth += 1
223 pending_attr_test = False
224 pending_blocking = False
225 continue
226 fm = FN_RE.match(tok)
227 if fm:
228 kind = "test" if pending_attr_test else "fn"
229 stack.append({"kind": kind, "open": None})
230 pending_attr_test = False
231 pending_blocking = False
232 continue
233 mm = MOD_RE.match(tok)
234 if mm:
235 name = mm.group(1)
236 kind = "test" if (pending_attr_test or name.startswith("test")) else "mod"
237 stack.append({"kind": kind, "open": None})
238 pending_attr_test = False
239 pending_blocking = False
240 continue
241 if tok == "impl":
242 stack.append({"kind": "impl", "open": None})
243 pending_attr_test = False
244 pending_blocking = False
245 continue
246 return {k: v for k, v in counts.items() if v}
247
248
249 def _rust_path_literal(text: str) -> tuple[str, int] | None:
250 """Only a literal include/path: never evaluate concat or a Rust expression."""
251 normal = re.match(r'"(?:[^"\\]|\\.)*"', text)
252 if normal:
253 try:
254 value = json.loads(normal.group())
255 except ValueError:
256 return None
257 return value, normal.end()
258 raw = re.match(r'r(#+|)"', text)
259 if raw:
260 end = text.find('"' + raw.group(1), raw.end())
261 if end >= 0:
262 return text[raw.end():end], end + 1 + len(raw.group(1))
263 return None
264
265
266 def _test_module_edges(path: Path, text: str, *, crate_root: bool = False) -> list[tuple[Path, bool]]:
267 """Literal item edges and their actual lexical cfg(test) module scope."""
268 code = strip_comments_and_strings(text, module_literals=True)
269 closes: dict[int, int] = {}
270 braces: list[int] = []
271 for position, char in enumerate(code):
272 if char == "{":
273 braces.append(position)
274 elif char == "}" and braces:
275 closes[braces.pop()] = position
276 modules = list(re.finditer(
277 r"(?P<attrs>(?:#\s*\[[^\]]*\]\s*)*)"
278 r"(?:pub(?:\([^)]*\))?\s+)?mod\s+(?P<name>[A-Za-z_]\w*)\s*(?P<end>[{;])",
279 code,
280 ))
281 inline = [entry for entry in modules if entry.group("end") == "{" and entry.end() - 1 in closes]
282 def test_attribute(entry: re.Match[str]) -> bool:
283 return any(re.fullmatch(r"#\s*\[\s*cfg\s*\(\s*test\s*\)\s*\]", attr.group())
284 for attr in re.finditer(r"#\s*\[[^\]]*\]", entry.group("attrs")))
285 def ancestors(position: int) -> list[re.Match[str]]:
286 return [entry for entry in inline if entry.end() <= position < closes[entry.end() - 1]]
287 def under_test(position: int) -> bool:
288 return any(test_attribute(entry) for entry in ancestors(position))
289 base = path.parent if crate_root or path.name in ("mod.rs", "lib.rs", "main.rs") else path.parent / path.stem
290 edges: list[tuple[Path, bool]] = []
291 for entry in modules:
292 if entry.group("end") != ";":
293 continue
294 parents = ancestors(entry.start())
295 # Path attributes on an inline module can change its directory owner;
296 # refuse to infer such a layout instead of exempting a guessed file.
297 if any(re.search(r"#\s*\[\s*path\s*=", parent.group("attrs")) for parent in parents):
298 continue
299 directory = base.joinpath(*(parent.group("name") for parent in parents))
300 attrs = text[entry.start("attrs"):entry.end("attrs")]
301 configured = re.search(r"#\s*\[\s*path\s*=\s*", attrs)
302 if configured:
303 literal = _rust_path_literal(attrs[configured.end():])
304 if not literal:
305 continue
306 directory = directory if parents else path.parent
307 candidates = [directory / literal[0]]
308 else:
309 name = entry.group("name")
310 candidates = [directory / f"{name}.rs", directory / name / "mod.rs"]
311 found = [candidate.resolve() for candidate in candidates if candidate.is_file()]
312 if len(found) == 1:
313 edges.append((found[0], test_attribute(entry) or under_test(entry.start())))
314 for entry in re.finditer(r"(?<![\w:])include\s*!\s*\(\s*", code):
315 # String contents are masked, so whitespace in `code` also covers the
316 # path. Recover only the original argument immediately after `(`.
317 opening = code.index("(", entry.start(), entry.end())
318 argument = text[opening + 1:].lstrip()
319 literal = _rust_path_literal(argument)
320 if not literal or not re.match(r"\s*,?\s*\)", argument[literal[1]:]):
321 continue
322 target = (path.parent / literal[0]).resolve()
323 if target.is_file():
324 edges.append((target, under_test(entry.start())))
325 return edges
326
327
328 def _cargo_crate_roots(scan_root: Path) -> tuple[set[Path], set[Path]]:
329 """Classify actual Cargo entry points; production always wins a conflict."""
330 roots: set[Path] = set()
331 test_roots: set[Path] = set()
332 # Preserve conventional roots even in hermetic scopes without a manifest.
333 for source in (scan_root, *scan_root.rglob("src")):
334 if not source.is_dir():
335 continue
336 candidates = [source / "lib.rs", source / "main.rs"]
337 candidates.extend((source / "bin").glob("*.rs"))
338 candidates.extend((source / "bin").glob("*/main.rs"))
339 roots.update(path.resolve() for path in candidates if path.is_file())
340 for manifest in scan_root.rglob("Cargo.toml"):
341 try:
342 config = tomllib.loads(manifest.read_text(encoding="utf-8"))
343 except (OSError, ValueError):
344 # A malformed/unreadable manifest must not exempt an unknown root.
345 roots.update(path.resolve() for path in manifest.parent.rglob("*.rs"))
346 continue
347 if not isinstance(config.get("package"), dict):
348 continue
349 targets = [config.get("lib", {}), *config.get("bin", [])]
350 for target in targets:
351 if isinstance(target, dict) and isinstance(target.get("path"), str):
352 path = manifest.parent / target["path"]
353 if path.is_file():
354 roots.add(path.resolve())
355 declared_names = set()
356 for target in config.get("test", []):
357 if not isinstance(target, dict):
358 continue
359 name, path = target.get("name"), target.get("path")
360 if isinstance(name, str):
361 declared_names.add(name)
362 if isinstance(path, str):
363 candidates = [manifest.parent / path]
364 elif isinstance(name, str):
365 candidates = [manifest.parent / "tests" / f"{name}.rs",
366 manifest.parent / "tests" / name / "main.rs"]
367 else:
368 continue
369 found = [candidate.resolve() for candidate in candidates if candidate.is_file()]
370 if len(found) == 1:
371 test_roots.add(found[0])
372 else:
373 roots.update(found) # Invalid/ambiguous target cannot exempt code.
374 if config["package"].get("autotests", True):
375 automatic: dict[str, list[Path]] = {}
376 directory = manifest.parent / "tests"
377 for candidate in (*directory.glob("*.rs"), *directory.glob("*/main.rs")):
378 name = candidate.parent.name if candidate.name == "main.rs" else candidate.stem
379 if name not in declared_names and candidate.is_file():
380 automatic.setdefault(name, []).append(candidate.resolve())
381 for found in automatic.values():
382 if len(found) == 1:
383 test_roots.add(found[0])
384 else:
385 roots.update(found)
386 return roots, test_roots
387
388
389 def module_file_scopes(scan_root: Path | None = None) -> tuple[set[Path], set[Path]]:
390 """Test-only files and production conflicts from the same literal graph.
391
392 Follow literal includes recursively from actual cfg(test) inline or
393 external modules. A file also included by production remains counted,
394 as do its unguarded descendants. Comments, strings and dynamic include
395 expressions are never an exemption, nor is a `tests`-looking file name.
396 """
397 scan_root = (CRATES if scan_root is None else scan_root).resolve()
398 cargo_production, cargo_tests = _cargo_crate_roots(scan_root)
399 graph: dict[Path, list[tuple[Path, bool]]] = {}
400 for path in scan_root.rglob("*.rs"):
401 try:
402 graph[path.resolve()] = _test_module_edges(
403 path, path.read_text(encoding="utf-8", errors="ignore"),
404 crate_root=path.resolve() in cargo_production | cargo_tests,
405 )
406 except OSError:
407 continue
408 test_reachable = {target for edges in graph.values() for target, test in edges if test}
409 test_reachable.update(cargo_tests & set(graph))
410 pending = list(test_reachable)
411 while pending:
412 for target, _test in graph.get(pending.pop(), []):
413 if target not in test_reachable:
414 test_reachable.add(target)
415 pending.append(target)
416 crate_roots = cargo_production & set(graph)
417 production = (set(graph) - test_reachable) | crate_roots
418 pending = list(production)
419 while pending:
420 for target, test in graph.get(pending.pop(), []):
421 if not test and target not in production:
422 production.add(target)
423 pending.append(target)
424 # A caller with legacy test-path heuristics must not exempt a known
425 # production/test conflict or a real Cargo entry point by its filename.
426 return test_reachable - production, (test_reachable & production) | crate_roots
427
428
429 def cfg_test_module_files(scan_root: Path | None = None) -> set[Path]:
430 """Reuse the literal graph with an explicit scope; default remains CRATES."""
431 return module_file_scopes(scan_root)[0]
432
433
434 def collect_current() -> dict[str, dict[str, int]]:
435 budget: dict[str, dict[str, int]] = {}
436 test_only = cfg_test_module_files()
437 for path in sorted(CRATES.rglob("*.rs")):
438 if path.resolve() in test_only:
439 continue
440 try:
441 counts = file_counts(path)
442 except OSError:
443 continue
444 if counts:
445 rel = str(path.relative_to(ROOT))
446 budget[rel] = counts
447 return budget
448
449
450 def main() -> int:
451 update = "--update" in sys.argv
452 current = collect_current()
453 if update:
454 BUDGET_PATH.write_text(
455 json.dumps(current, indent=2, sort_keys=True) + "\n",
456 encoding="utf-8",
457 )
458 total = sum(sum(v.values()) for v in current.values())
459 print(f"wrote {BUDGET_PATH.name}: {total} sites across {len(current)} files")
460 return 0
461
462 if not BUDGET_PATH.exists():
463 print(f"missing {BUDGET_PATH.name}; run with --update to create it", file=sys.stderr)
464 return 2
465 budget = json.loads(BUDGET_PATH.read_text(encoding="utf-8"))
466
467 failures: list[str] = []
468 savings: list[str] = []
469 for path, counts in sorted(current.items()):
470 allowed = budget.get(path, {})
471 for name, count in counts.items():
472 limit = allowed.get(name, 0)
473 if count > limit:
474 failures.append(
475 f"{path}: {name} sites {count} > budget {limit}"
476 )
477 elif count < limit:
478 savings.append(
479 f"{path}: {name} sites {count} < budget {limit} — tighten with --update"
480 )
481 for path, counts in sorted(budget.items()):
482 if path not in current:
483 savings.append(f"{path}: file clean — tighten with --update")
484 else:
485 for name in counts:
486 if name not in current[path]:
487 savings.append(
488 f"{path}: {name} sites 0 < budget {counts[name]} — tighten with --update"
489 )
490
491 for line in savings:
492 print(line)
493 if failures:
494 print(
495 "\nBlocking-call budget exceeded — new `thread::sleep`/`std::fs` call "
496 "sites appeared outside spawn_blocking/dedicated-thread/test scopes:",
497 file=sys.stderr,
498 )
499 for line in failures:
500 print(f" {line}", file=sys.stderr)
501 print(
502 "Move the work into `tokio::task::spawn_blocking` (or use tokio::fs "
503 "/ tokio::time). If the site can only run on synchronous code, land "
504 "the raised budget in this PR and say why in the PR description:\n"
505 " python3 scripts/check-blocking-calls-budget.py --update\n"
506 "See #6149.",
507 file=sys.stderr,
508 )
509 return 1
510 total = sum(sum(v.values()) for v in current.values())
511 print(f"blocking-call budget: {total} sites across {len(current)} files, within budget")
512 return 0
513
514
515 if __name__ == "__main__":
516 raise SystemExit(main())
517
517 lines PYTHON