chore(relay): generate reconnect token size and version predicate from the protocol spec
The 32-byte reconnect token and its 43-character base64url shape were
pinned by hand in both server/main.go and the Dart peer service, and the
supported-version check was spelled out three times in Go. The spec now
carries reconnectTokenBytes and the generator emits the token size,
encoded length, and validator for Dart plus the size and
supportedRelayProtocolVersion for Go, failing before writing either
target when the key is missing. The last handwritten error code in lib/
('not_in_room') uses the generated constant, and releaseSession's
in-flight join is a FutureCoalescer.
This commit is contained in:
@@ -32,8 +32,25 @@ def validated_id_pattern(spec: dict) -> str:
|
||||
return pattern
|
||||
|
||||
|
||||
def validated_reconnect_token_bytes(spec: dict) -> int:
|
||||
try:
|
||||
size = spec["reconnectTokenBytes"]
|
||||
except KeyError:
|
||||
raise ValueError("reconnectTokenBytes is required") from None
|
||||
if isinstance(size, bool) or not isinstance(size, int) or size <= 0:
|
||||
raise ValueError("reconnectTokenBytes must be a positive integer")
|
||||
return size
|
||||
|
||||
|
||||
def reconnect_token_length(token_bytes: int) -> int:
|
||||
"""Length of the unpadded base64url encoding of ``token_bytes`` bytes."""
|
||||
return (token_bytes * 4 + 2) // 3
|
||||
|
||||
|
||||
def dart_source(spec: dict) -> str:
|
||||
id_pattern = validated_id_pattern(spec)
|
||||
token_bytes = validated_reconnect_token_bytes(spec)
|
||||
token_length = reconnect_token_length(token_bytes)
|
||||
lines = [
|
||||
"// Generated by scripts/codegen/generate_relay_protocol.py. Do not edit.",
|
||||
"",
|
||||
@@ -56,14 +73,20 @@ def dart_source(spec: dict) -> str:
|
||||
lines.append(f" static const int {name} = {value};")
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
f" static const int reconnectTokenBytes = {token_bytes};",
|
||||
f" static const int reconnectTokenLength = {token_length};",
|
||||
"",
|
||||
f" static final RegExp _idPattern = RegExp(r{id_pattern!r});",
|
||||
f" static final RegExp _reconnectTokenPattern = RegExp(r'^[A-Za-z0-9_-]{{{token_length}}}$');",
|
||||
"",
|
||||
" static bool isValidSessionId(String value) =>",
|
||||
" value.isNotEmpty && value.length <= maxSessionIdLength && _idPattern.hasMatch(value);",
|
||||
"",
|
||||
" static bool isValidPeerId(String value) =>",
|
||||
" value.isNotEmpty && value.length <= maxPeerIdLength && _idPattern.hasMatch(value);",
|
||||
"",
|
||||
" static bool isValidReconnectToken(String value) => _reconnectTokenPattern.hasMatch(value);",
|
||||
"}",
|
||||
"",
|
||||
]
|
||||
@@ -73,6 +96,7 @@ def dart_source(spec: dict) -> str:
|
||||
|
||||
def go_source(spec: dict) -> str:
|
||||
validated_id_pattern(spec)
|
||||
token_bytes = validated_reconnect_token_bytes(spec)
|
||||
lines = [
|
||||
"// Code generated by scripts/codegen/generate_relay_protocol.py. DO NOT EDIT.",
|
||||
"",
|
||||
@@ -112,6 +136,7 @@ def go_source(spec: dict) -> str:
|
||||
limit_constants = [
|
||||
(go_limit_names[name], str(value)) for name, value in spec["limits"].items()
|
||||
]
|
||||
limit_constants.append(("reconnectTokenSize", str(token_bytes)))
|
||||
limit_name_width = max(len(name) for name, _ in limit_constants)
|
||||
lines.extend(
|
||||
f"\t{name:<{limit_name_width}} = {value}"
|
||||
@@ -121,6 +146,10 @@ def go_source(spec: dict) -> str:
|
||||
[
|
||||
")",
|
||||
"",
|
||||
"func supportedRelayProtocolVersion(version int) bool {",
|
||||
"\treturn version == legacyRelayProtocolVersion || version == relayProtocolVersion",
|
||||
"}",
|
||||
"",
|
||||
"func validRelayID(value string, maxLength int) bool {",
|
||||
"\tif len(value) == 0 || len(value) > maxLength {",
|
||||
"\t\treturn false",
|
||||
|
||||
@@ -86,6 +86,65 @@ class RelayProtocolGeneratorTest(unittest.TestCase):
|
||||
with self.assertRaisesRegex(ValueError, "idPattern must be a string"):
|
||||
generator.validated_id_pattern(spec)
|
||||
|
||||
def test_reconnect_token_length_is_unpadded_base64url_length(self) -> None:
|
||||
for token_bytes, expected in ((1, 2), (2, 3), (3, 4), (32, 43), (33, 44)):
|
||||
with self.subTest(token_bytes=token_bytes):
|
||||
self.assertEqual(generator.reconnect_token_length(token_bytes), expected)
|
||||
|
||||
def test_reconnect_token_bytes_render_derived_symbols(self) -> None:
|
||||
spec = copy.deepcopy(self.spec)
|
||||
spec["reconnectTokenBytes"] = 24
|
||||
|
||||
dart_output = generator.dart_source(spec)
|
||||
go_output = generator.go_source(spec)
|
||||
|
||||
self.assertIn("static const int reconnectTokenBytes = 24;", dart_output)
|
||||
self.assertIn("static const int reconnectTokenLength = 32;", dart_output)
|
||||
self.assertIn("RegExp(r'^[A-Za-z0-9_-]{32}$')", dart_output)
|
||||
self.assertIn("static bool isValidReconnectToken(String value)", dart_output)
|
||||
self.assertIn("reconnectTokenSize = 24", go_output)
|
||||
self.assertIn("func supportedRelayProtocolVersion(version int) bool", go_output)
|
||||
|
||||
def test_missing_reconnect_token_bytes_fails_before_writing_either_target(self) -> None:
|
||||
spec = copy.deepcopy(self.spec)
|
||||
del spec["reconnectTokenBytes"]
|
||||
|
||||
for renderer in (generator.dart_source, generator.go_source):
|
||||
with self.subTest(renderer=renderer.__name__):
|
||||
with self.assertRaisesRegex(ValueError, "reconnectTokenBytes is required"):
|
||||
renderer(spec)
|
||||
|
||||
with tempfile.TemporaryDirectory() as temporary_directory:
|
||||
root = Path(temporary_directory)
|
||||
spec_path = root / "relay_protocol.json"
|
||||
dart_path = root / "relay_protocol.g.dart"
|
||||
go_path = root / "relay_protocol_gen.go"
|
||||
spec_path.write_text(json.dumps(spec), encoding="utf-8")
|
||||
dart_path.write_text("dart sentinel\n", encoding="utf-8")
|
||||
go_path.write_text("go sentinel\n", encoding="utf-8")
|
||||
|
||||
with (
|
||||
mock.patch.object(generator, "SPEC_PATH", spec_path),
|
||||
mock.patch.object(generator, "DART_PATH", dart_path),
|
||||
mock.patch.object(generator, "GO_PATH", go_path),
|
||||
):
|
||||
with self.assertRaisesRegex(ValueError, "reconnectTokenBytes"):
|
||||
generator.main()
|
||||
|
||||
self.assertEqual(dart_path.read_text(encoding="utf-8"), "dart sentinel\n")
|
||||
self.assertEqual(go_path.read_text(encoding="utf-8"), "go sentinel\n")
|
||||
|
||||
def test_invalid_reconnect_token_bytes_is_rejected(self) -> None:
|
||||
for value in (None, "32", 32.0, True, 0, -1):
|
||||
with self.subTest(value=value):
|
||||
spec = copy.deepcopy(self.spec)
|
||||
spec["reconnectTokenBytes"] = value
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "reconnectTokenBytes must be a positive integer"
|
||||
):
|
||||
generator.validated_reconnect_token_bytes(spec)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user