from __future__ import annotations
import re
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
ALLOW = {
"allow_replication_factor_change": "echo of the AlterPartitionReassignments request field",
"endpoint_type": "krafka routes to the controller by node id, never by endpoint type",
}
MESSAGES = "src/protocol/messages"
def response_fields() -> dict[str, set[str]]:
found: dict[str, set[str]] = {}
for path in sorted((ROOT / MESSAGES).glob("*.rs")):
source = path.read_text()
for match in re.finditer(r"pub struct (\w*Response\w*) \{(.*?)\n\}", source, re.S):
struct, body = match.group(1), match.group(2)
for field in re.findall(r"^\s+pub (\w+):", body, re.M):
found.setdefault(field, set()).add(f"{path.name}::{struct}")
return found
def client_source() -> str:
parts = []
for path in sorted((ROOT / "src").rglob("*.rs")):
rel = path.relative_to(ROOT)
if "protocol" in rel.parts or "testing" in rel.parts:
continue
source = path.read_text()
cut = source.find("\n#[cfg(test)]")
if cut > 0:
source = source[:cut]
parts.append(source)
return "\n".join(parts)
def main() -> int:
fields = response_fields()
client = client_source()
unread = [
(field, sites)
for field, sites in sorted(fields.items())
if field not in ALLOW and not re.search(r"\b" + re.escape(field) + r"\b", client)
]
stale = sorted(set(ALLOW) - set(fields))
if stale:
print("Protocol reachability check FAILED\n", file=sys.stderr)
for field in stale:
print(
f" - ALLOW names `{field}`, which is no longer a response field.\n"
" Remove the entry from xtask/protocol_reachability.py.\n",
file=sys.stderr,
)
return 1
if unread:
print("Protocol reachability check FAILED\n", file=sys.stderr)
for field, sites in unread:
where = ", ".join(sorted(sites))
print(
f" - `{field}` is decoded by {where} and read by no client code.\n"
" Either use it, or add it to ALLOW in "
"xtask/protocol_reachability.py with the\n reason it is "
"decode-only. A field the broker sends and the client throws away\n"
" is information the application cannot get any other way.\n",
file=sys.stderr,
)
return 1
print(
f"✓ Protocol reachability: {len(fields)} response fields, "
f"every one read by client code ({len(ALLOW)} documented decode-only)"
)
return 0
if __name__ == "__main__":
sys.exit(main())