Something went wrong. Try again.
atproto Thingiverse but good
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650#!/usr/bin/env python3"""Reject runtime SQLx query calls in production Rust sources.
The production policy applies to handwritten Rust under ``src/``. Explicitlyowned test files and test-only module bodies are outside this check becausetheir dynamic SQL is fixture/setup code rather than application SQL. Files suchas ``contest.rs`` remain production source; a substring in a filename is nottest ownership."""
from __future__ import annotations
import argparseimport reimport sysfrom pathlib import Pathfrom tempfile import TemporaryDirectory
# Comments are blanked before this expression runs, so ``::`` may be separated# from either identifier by whitespace or comments. The optional turbofish is# deliberately accepted only before a call parenthesis; ``query!`` therefore# remains a compile-time macro and is not reported.RUNTIME_CALL = re.compile( r"\bsqlx\s*::\s*query(?:_as|_scalar)?\s*" r"(?:::\s*<[^;{}\n]*>\s*)?\(")CFG_ATTRIBUTE_START = re.compile(r"#\s*\[\s*cfg\s*\(")
def _skip_rust_token(source: str, index: int) -> int | None: """Return the end of a comment/string/char token at ``index``.""" if source.startswith("//", index): newline = source.find("\n", index + 2) return len(source) if newline == -1 else newline
if source.startswith("/*", index): depth = 1 cursor = index + 2 while cursor < len(source) and depth: if source.startswith("/*", cursor): depth += 1 cursor += 2 elif source.startswith("*/", cursor): depth -= 1 cursor += 2 else: cursor += 1 return cursor
raw_prefix_end: int | None = None if source.startswith("br", index): raw_prefix_end = index + 2 elif source.startswith("r", index): raw_prefix_end = index + 1 if raw_prefix_end is not None: cursor = raw_prefix_end while cursor < len(source) and source[cursor] == "#": cursor += 1 if cursor < len(source) and source[cursor] == '"': hashes = cursor - raw_prefix_end closing = '"' + ("#" * hashes) end = source.find(closing, cursor + 1) return len(source) if end == -1 else end + len(closing)
quote_index: int | None = None if source[index] == '"': quote_index = index elif index + 1 < len(source) and source[index] in "bc" and source[index + 1] == '"': quote_index = index + 1 if quote_index is not None: cursor = quote_index + 1 escaped = False while cursor < len(source): char = source[cursor] if escaped: escaped = False elif char == "\\": escaped = True elif char == '"': return cursor + 1 cursor += 1 return len(source)
if source[index] == "'": cursor = index + 1 if cursor >= len(source): return None if source[cursor] == "\\": cursor += 2 if source[index + 1 : index + 3] == "\\u" and cursor < len(source) and source[cursor] == "{": closing = source.find("}", cursor + 1) cursor = len(source) if closing == -1 else closing + 1 else: cursor += 1 if cursor < len(source) and source[cursor] == "'": return cursor + 1
return None
def _is_code_position(source: str, target: int) -> bool: """Return whether ``target`` is outside comments, strings, and chars.""" index = 0 while index < target: end = _skip_rust_token(source, index) if end is not None: if index <= target < end: return False index = end else: index += 1 return True
def _find_identifier(source: str, start: int, wanted: str) -> int | None: index = start while index < len(source): end = _skip_rust_token(source, index) if end is not None: index = end continue if source.startswith(wanted, index): before = source[index - 1] if index else " " after_index = index + len(wanted) after = source[after_index] if after_index < len(source) else " " if not (before.isalnum() or before == "_") and not ( after.isalnum() or after == "_" ): return index index += 1 return None
def _find_next_code_char(source: str, start: int, wanted: str) -> int | None: index = start while index < len(source): end = _skip_rust_token(source, index) if end is not None: index = end elif source[index] == wanted: return index else: index += 1 return None
def _find_matching_brace(source: str, body_start: int) -> int | None: depth = 0 index = body_start while index < len(source): end = _skip_rust_token(source, index) if end is not None: index = end continue if source[index] == "{": depth += 1 elif source[index] == "}": depth -= 1 if depth == 0: return index + 1 index += 1 return None
def _matching_delimiter(source: str, opening: int, left: str, right: str) -> int | None: """Find a matching delimiter while ignoring comments and literals.""" depth = 0 index = opening while index < len(source): end = _skip_rust_token(source, index) if end is not None: index = end continue if source[index] == left: depth += 1 elif source[index] == right: depth -= 1 if depth == 0: return index index += 1 return None
def _find_item_terminator(source: str, start: int) -> int | None: """Return the end of a const/static item, ignoring nested delimiters.""" pairs = {"(": ")", "[": "]", "{": "}"} closing = {right: left for left, right in pairs.items()} stack: list[str] = [] index = start while index < len(source): end = _skip_rust_token(source, index) if end is not None: index = end continue char = source[index] if char in pairs: stack.append(char) elif char in closing: if stack and stack[-1] == closing[char]: stack.pop() elif char == ";" and not stack: return index + 1 index += 1 return None
def _cfg_truth_without_test(expression: str) -> bool | None: """Evaluate a cfg expression with ``test`` disabled when provable.
``None`` means that an unknown predicate may be either true or false. This three-valued evaluation keeps negation conservative: only expressions proven false are masked from the production SQL scan. """ expression = expression.strip() if expression == "test": return False if expression.startswith("all(") and expression.endswith(")"): values = [ _cfg_truth_without_test(part) for part in _split_cfg_arguments(expression[4:-1]) ] if any(value is False for value in values): return False if all(value is True for value in values): return True return None if expression.startswith("any(") and expression.endswith(")"): values = [ _cfg_truth_without_test(part) for part in _split_cfg_arguments(expression[4:-1]) ] if any(value is True for value in values): return True if all(value is False for value in values): return False return None if expression.startswith("not(") and expression.endswith(")"): value = _cfg_truth_without_test(expression[4:-1]) return None if value is None else not value return None
def _cfg_can_be_true_without_test(expression: str) -> bool: """Return whether a cfg expression can match with ``test`` disabled.""" return _cfg_truth_without_test(expression) is not False
def _split_cfg_arguments(expression: str) -> list[str]: parts: list[str] = [] start = 0 depth = 0 index = 0 while index < len(expression): end = _skip_rust_token(expression, index) if end is not None: index = end continue char = expression[index] if char == "(": depth += 1 elif char == ")": depth -= 1 elif char == "," and depth == 0: parts.append(expression[start:index].strip()) start = index + 1 index += 1 parts.append(expression[start:].strip()) return [part for part in parts if part]
def _skip_space(source: str, index: int) -> int: while index < len(source) and source[index].isspace(): index += 1 return index
def _identifier_at(source: str, index: int, identifier: str) -> int | None: if not source.startswith(identifier, index): return None before = source[index - 1] if index else " " after_index = index + len(identifier) after = source[after_index] if after_index < len(source) else " " if before.isalnum() or before == "_" or after.isalnum() or after == "_": return None return after_index
def _attributed_item_body(code: str, attribute_start: int, attribute_end: int) -> tuple[int, int] | None: """Return ``(attribute_start, body_end)`` for a cfg-attributed item body.""" index = _skip_space(code, attribute_end) # Rust permits a stack of attributes before the item. Consume only complete # attributes immediately adjacent to this item; never search ahead for an # unrelated declaration. while index < len(code) and code.startswith("#[", index): closing = _matching_delimiter(code, index + 1, "[", "]") if closing is None: return None index = _skip_space(code, closing + 1)
if code.startswith("pub", index): after_pub = _identifier_at(code, index, "pub") if after_pub is None: return None index = _skip_space(code, after_pub) if index < len(code) and code[index] == "(": visibility_end = _matching_delimiter(code, index, "(", ")") if visibility_end is None: return None index = _skip_space(code, visibility_end + 1)
after_mod = _identifier_at(code, index, "mod") if after_mod is not None: index = _skip_space(code, after_mod) name_end = index while name_end < len(code) and (code[name_end].isalnum() or code[name_end] == "_"): name_end += 1 if name_end == index: return None index = _skip_space(code, name_end) if index >= len(code) or code[index] != "{": # External modules have no body in this file. Their separate file # is checked by source_paths according to its explicit ownership. return None body_end = _find_matching_brace(code, index) return None if body_end is None else (attribute_start, body_end)
for item_kind in ("impl", "trait"): after_item_kind = _identifier_at(code, index, item_kind) if after_item_kind is None: continue body_open = _find_next_code_char(code, after_item_kind, "{") if body_open is None: return None body_end = _find_matching_brace(code, body_open) return None if body_end is None else (attribute_start, body_end)
after_const = _identifier_at(code, index, "const") if after_const is not None: after_const_name = _skip_space(code, after_const) # `const fn` is handled by the function path below; other const items # end at their top-level semicolon, including block initializers. if _identifier_at(code, after_const_name, "fn") is None: body_end = _find_item_terminator(code, after_const) return None if body_end is None else (attribute_start, body_end)
after_static = _identifier_at(code, index, "static") if after_static is not None: body_end = _find_item_terminator(code, after_static) return None if body_end is None else (attribute_start, body_end)
# Test helpers are commonly free functions rather than nested modules. # Consume the ordinary qualifiers before `fn`, then find the function body # without crossing a declaration terminator. while True: qualifier = None for candidate in ("const", "async", "unsafe"): end = _identifier_at(code, index, candidate) if end is not None: qualifier = end break if qualifier is None: break index = _skip_space(code, qualifier) after_fn = _identifier_at(code, index, "fn") if after_fn is None: return None body_open = _find_next_code_char(code, after_fn, "{") if body_open is None: return None terminator = _find_next_code_char(code, after_fn, ";") if terminator is not None and terminator < body_open: return None body_end = _find_matching_brace(code, body_open) return None if body_end is None else (attribute_start, body_end)
def mask_test_modules(source: str) -> str: """Blank test-only cfg-attributed module and function bodies.""" code = mask_non_code_tokens(source) result = list(source) for match in CFG_ATTRIBUTE_START.finditer(code): opening = code.find("(", match.start(), match.end()) if opening == -1: continue closing = _matching_delimiter(code, opening, "(", ")") if closing is None: continue bracket = _skip_space(code, closing + 1) if bracket >= len(code) or code[bracket] != "]": continue expression = source[opening + 1 : closing] if _cfg_can_be_true_without_test(expression): continue item = _attributed_item_body(code, match.start(), bracket + 1) if item is None: continue start, body_end = item for index in range(start, body_end): if result[index] != "\n": result[index] = " " return "".join(result)
def mask_non_code_tokens(source: str) -> str: """Blank comments and literals so regex checks only inspect Rust code.""" result = list(source) index = 0 while index < len(source): end = _skip_rust_token(source, index) if end is None: index += 1 continue for token_index in range(index, end): if result[token_index] != "\n": result[token_index] = " " index = end return "".join(result)
def production_source(path: Path) -> str: source = path.read_text(encoding="utf-8") return mask_non_code_tokens(mask_test_modules(source))
def _is_explicit_test_owned(path: Path, root: Path) -> bool: """Recognize conventional test ownership without substring heuristics.""" relative = path.relative_to(root / "src") parts = relative.parts if any(part in {"tests", "test"} for part in parts[:-1]): return True return path.stem in {"tests", "test"} or path.name.endswith("_tests.rs")
PATH_MODULE = re.compile( r"#\s*\[\s*path\s*=\s*([\"'])([^\"']+)\1\s*\]")FILE_CFG_TEST = re.compile(r"(?m)^\s*#!\s*\[\s*cfg\s*\(\s*test\s*\)\s*\]")CFG_TEST_ATTRIBUTE = re.compile(r"#\s*\[\s*cfg\s*\(\s*test\s*\)\s*\]\s*$")
def _test_owned_path_modules(root: Path) -> set[Path]: """Find sibling files imported only through test-gated path modules.""" owned: set[Path] = set() src = root / "src" for parent in src.rglob("*.rs"): source = parent.read_text(encoding="utf-8") file_is_test = FILE_CFG_TEST.search(source) is not None for match in PATH_MODULE.finditer(source): preceding = source[: match.start()] if not file_is_test and CFG_TEST_ATTRIBUTE.search(preceding) is None: continue target = (parent.parent / match.group(2)).resolve() if target.is_file() and target.is_relative_to(src.resolve()): owned.add(target) return owned
def source_paths(root: Path) -> list[Path]: src = root / "src" if not src.is_dir(): return [] test_owned_modules = _test_owned_path_modules(root) return sorted( path for path in src.rglob("*.rs") if not _is_explicit_test_owned(path, root) and path.resolve() not in test_owned_modules )
def violations(root: Path) -> list[tuple[Path, int]]: findings: list[tuple[Path, int]] = [] for path in source_paths(root): source = production_source(path) for match in RUNTIME_CALL.finditer(source): findings.append((path, source.count("\n", 0, match.start()) + 1)) return findings
def run_self_test() -> None: with TemporaryDirectory() as directory: root = Path(directory) (root / "src" / "tests").mkdir(parents=True) (root / "migrations").mkdir() (root / "src" / "production.rs").write_text( "pub fn seed() {\n" " sqlx /* path */ :: query /* call */ (\"SELECT 1\");\n" " sqlx::query_as :: <_, Row>(\"SELECT 2\");\n" " sqlx::query_scalar!(\"SELECT 3\");\n" "}\n", encoding="utf-8", ) (root / "src" / "nested.rs").write_text( "#[cfg(test)]\nmod tests {\n" " /* outer { /* nested } */ still test */\n" " let sql = r#\"fixture { cfg(test) }\"#;\n" " let lifetime: &'a str = 'x';\n" " sqlx::query(\"fixture\");\n" "}\n\n" "#[cfg(test)]\nfn test_function_does_not_mask_following_module() {\n" " sqlx::query(\"test-only function fixture\");\n" "}\n" "mod production_after_test_function {\n" " sqlx::query_scalar(\"SELECT 4\");\n" "}\n\n" "#[cfg(test)]\nmod external_tests;\n\n" "pub fn checked() { sqlx::query_scalar(\"SELECT 5\"); }\n", encoding="utf-8", ) (root / "src" / "cfg_forms.rs").write_text( "#[cfg(test)]\n" "impl TestImpl {\n" " fn query(&self) {\n" " sqlx::query(\"masked impl\");\n" " }\n" "}\n" "#[cfg(test)]\n" "trait TestTrait {\n" " fn query(&self) {\n" " sqlx::query(\"masked trait\");\n" " }\n" "}\n" "#[cfg(test)]\n" "const TEST_CONST: &str = {\n" " sqlx::query(\"masked const\");\n" " \"fixture\"\n" "};\n" "#[cfg(test)]\n" "static TEST_STATIC: &str = {\n" " sqlx::query(\"masked static\");\n" " \"fixture\"\n" "};\n" "#[cfg(all(test, feature = \"fixtures\"))]\n" "mod all_tests { sqlx::query(\"masked all\"); }\n" "#[cfg(any(test, feature = \"fixtures\"))]\n" "mod mixed_tests { sqlx::query(\"reported any\"); }\n" "#[cfg(not(any(test, feature = \"fixtures\")))]\n" "mod negated_mixed { sqlx::query(\"reported negated any\"); }\n" "#[cfg(not(feature = \"fixtures\"))]\n" "mod negated_feature { sqlx::query(\"reported negated feature\"); }\n", encoding="utf-8", ) (root / "src" / "oauth.rs").write_text( "#[cfg(test)]\nmod tests {\n" " let raw = br##\"{ nested \\\"quotes\\\" }\"##;\n" " sqlx::query_as(\"fixture\");\n" "}\n", encoding="utf-8", ) (root / "src" / "comments.rs").write_text( "// #[cfg(test)] mod ignored { sqlx::query(\"comment\"); }\n" "pub const TEXT: &str = r#\"#[cfg(test)] mod ignored { }\"#;\n", encoding="utf-8", ) (root / "src" / "fixture_tests.rs").write_text( "sqlx::query_as(\"fixture file\");\n", encoding="utf-8" ) (root / "src" / "test_dispatch.rs").write_text( "#![cfg(test)]\n" "#[path = \"test_support.rs\"]\n" "mod test_support;\n", encoding="utf-8", ) (root / "src" / "test_support.rs").write_text( "sqlx::query(\"test fixture module\");\n", encoding="utf-8" ) (root / "src" / "production_path.rs").write_text( "#[path = \"production_support.rs\"]\n" "mod production_support;\n", encoding="utf-8", ) (root / "src" / "production_support.rs").write_text( "sqlx::query(\"production path module\");\n", encoding="utf-8" ) (root / "src" / "contest.rs").write_text( "pub fn production_contest() { sqlx::query(\"contest\"); }\n", encoding="utf-8", ) (root / "migrations" / "0001.sql").write_text( "sqlx::query(\"migration text\");\n", encoding="utf-8" )
found = violations(root) expected = [ (root / "src" / "cfg_forms.rs", 26), (root / "src" / "cfg_forms.rs", 28), (root / "src" / "cfg_forms.rs", 30), (root / "src" / "contest.rs", 1), (root / "src" / "nested.rs", 14), (root / "src" / "nested.rs", 20), (root / "src" / "production.rs", 2), (root / "src" / "production.rs", 3), (root / "src" / "production_support.rs", 1), ] if found != expected: raise AssertionError(f"checker self-test mismatch: {found!r}")
(root / "src" / "production.rs").write_text( "pub fn seed() { sqlx::query!(\"SELECT 1\"); }\n", encoding="utf-8" ) (root / "src" / "contest.rs").write_text( "pub fn production_contest() { sqlx::query!(\"contest\"); }\n", encoding="utf-8", ) (root / "src" / "nested.rs").write_text( "#[cfg(test)]\nmod tests { sqlx::query(\"masked\"); }\n", encoding="utf-8", ) (root / "src" / "cfg_forms.rs").write_text( "#[cfg(all(test, feature = \"fixtures\"))]\n" "mod all_tests { sqlx::query(\"masked\"); }\n", encoding="utf-8", ) (root / "src" / "production_support.rs").write_text( "sqlx::query!(\"production path\");\n", encoding="utf-8" ) found = violations(root) if found: raise AssertionError(f"allowed macro self-test mismatch: {found!r}") print("check-production-sql self-test: PASS")
def main() -> int: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--self-test", action="store_true") parser.add_argument( "root", nargs="?", type=Path, default=Path(__file__).resolve().parent.parent, help="repository root (defaults to the directory containing this script)", ) args = parser.parse_args() if args.self_test: run_self_test() return 0
root = args.root.resolve() findings = violations(root) if findings: for path, line in findings: print(f"{path.relative_to(root)}:{line}: runtime sqlx query in production source") return 1 print("check-production-sql: PASS (no runtime SQLx queries in production src)") return 0
if __name__ == "__main__": sys.exit(main())