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",
}
PASSTHROUGH = {
"DescribedStreamsGroup": "returned verbatim by AdminClient::describe_streams_groups",
}
MESSAGES = "src/protocol/messages"
def structs() -> dict[str, tuple[str, str]]:
found: dict[str, tuple[str, str]] = {}
for path in sorted((ROOT / MESSAGES).glob("*.rs")):
source = path.read_text()
for match in re.finditer(r"pub struct (\w+) \{(.*?)\n\}", source, re.S):
found[match.group(1)] = (path.name, match.group(2))
return found
def contained(roots: list[str], defs: dict[str, tuple[str, str]]) -> set[str]:
seen: set[str] = set()
stack = list(roots)
while stack:
name = stack.pop()
if name in seen or name not in defs:
continue
seen.add(name)
for referenced in re.findall(r"\b([A-Z]\w+)\b", defs[name][1]):
if referenced in defs and referenced not in seen:
stack.append(referenced)
return seen
def response_fields() -> dict[str, set[str]]:
defs = structs()
reachable = contained([name for name in defs if "Response" in name], defs)
passthrough = contained([name for name in PASSTHROUGH if name in defs], defs)
found: dict[str, set[str]] = {}
for struct in sorted(reachable - passthrough):
file_name, body = defs[struct]
for field in re.findall(r"^\s+pub (\w+):", body, re.M):
found.setdefault(field, set()).add(f"{file_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))
stale += sorted(name for name in PASSTHROUGH if name not in structs())
if stale:
print("Protocol reachability check FAILED\n", file=sys.stderr)
for field in stale:
print(
f" - `{field}` is named by ALLOW or PASSTHROUGH but is no longer "
"a response field or struct.\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, "
f"{len(PASSTHROUGH)} type(s) returned verbatim)"
)
return 0
if __name__ == "__main__":
sys.exit(main())