From ab477c971d4e4fff89736004f3502b408037111c Mon Sep 17 00:00:00 2001 From: edde746 <86283021+edde746@users.noreply.github.com> Date: Thu, 3 Sep 2026 08:20:35 +0200 Subject: [PATCH] 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. --- .../providers/watch_together_provider.dart | 2 +- .../services/relay_protocol.g.dart | 6 ++ .../services/watch_together_peer_service.dart | 19 ++---- relay_protocol.json | 3 +- scripts/codegen/generate_relay_protocol.py | 29 +++++++++ .../codegen/test_generate_relay_protocol.py | 59 +++++++++++++++++++ server/main.go | 7 +-- server/main_test.go | 14 +++++ server/relay_protocol_gen.go | 5 ++ 9 files changed, 124 insertions(+), 20 deletions(-) diff --git a/lib/watch_together/providers/watch_together_provider.dart b/lib/watch_together/providers/watch_together_provider.dart index cfe3e1fa6..04264aad0 100644 --- a/lib/watch_together/providers/watch_together_provider.dart +++ b/lib/watch_together/providers/watch_together_provider.dart @@ -691,7 +691,7 @@ class WatchTogetherProvider with ChangeNotifier { _errorSubscription = peerService.onError.listen((error) { if (_disposed || !identical(_peerService, peerService)) return; final hostPeerId = _session?.hostPeerId; - if (error.serverCode == 'not_in_room' && + if (error.serverCode == RelayProtocol.notInRoomCode && !isHost && hostPeerId != null && !peerService.connectedPeers.contains(hostPeerId)) { diff --git a/lib/watch_together/services/relay_protocol.g.dart b/lib/watch_together/services/relay_protocol.g.dart index 5537c21ac..63c4df3bf 100644 --- a/lib/watch_together/services/relay_protocol.g.dart +++ b/lib/watch_together/services/relay_protocol.g.dart @@ -39,11 +39,17 @@ abstract final class RelayProtocol { static const int maxSessionIdLength = 64; static const int maxPeerIdLength = 128; + static const int reconnectTokenBytes = 32; + static const int reconnectTokenLength = 43; + static final RegExp _idPattern = RegExp(r'^[A-Za-z0-9_-]+$'); + static final RegExp _reconnectTokenPattern = RegExp(r'^[A-Za-z0-9_-]{43}$'); 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); } diff --git a/lib/watch_together/services/watch_together_peer_service.dart b/lib/watch_together/services/watch_together_peer_service.dart index e45002261..1fc32950c 100644 --- a/lib/watch_together/services/watch_together_peer_service.dart +++ b/lib/watch_together/services/watch_together_peer_service.dart @@ -8,6 +8,7 @@ import 'package:web_socket_channel/web_socket_channel.dart'; import '../../i18n/strings.g.dart'; import '../../services/base_peer_service.dart'; +import '../../services/trackers/future_coalescer.dart'; import '../../utils/app_logger.dart'; import '../models/sync_message.dart'; import 'relay_protocol.g.dart'; @@ -92,7 +93,7 @@ class WatchTogetherPeerService with KeepaliveMixin { bool _disposed = false; bool _initialSetupInProgress = false; bool _teardownInProgress = false; - Future? _releaseFuture; + final FutureCoalescer _release = FutureCoalescer(); /// Called after a successful reconnection so the provider can re-announce join. void Function()? onReconnected; @@ -159,7 +160,7 @@ class WatchTogetherPeerService with KeepaliveMixin { /// retried without relying on server-returned state. static String _mintReconnectToken() { final random = Random.secure(); - final bytes = List.generate(32, (_) => random.nextInt(256), growable: false); + final bytes = List.generate(RelayProtocol.reconnectTokenBytes, (_) => random.nextInt(256), growable: false); return base64Url.encode(bytes).replaceAll('=', ''); } @@ -259,8 +260,6 @@ class WatchTogetherPeerService with KeepaliveMixin { ); } - static final RegExp _reconnectTokenPattern = RegExp(r'^[A-Za-z0-9_-]{43}$'); - PeerError _invalidSetupResponse(String type) { appLogger.w('WatchTogether: Relay returned an invalid $type response'); return PeerError(type: PeerErrorType.serverError, message: t.watchTogether.errors.invalidRelayResponse); @@ -278,7 +277,7 @@ class WatchTogetherPeerService with KeepaliveMixin { (_isHost && hostPeerId != _myPeerId) || reconnectToken is! String || reconnectToken != _reconnectToken || - !_reconnectTokenPattern.hasMatch(reconnectToken) || + !RelayProtocol.isValidReconnectToken(reconnectToken) || protocolVersion != _relayProtocolVersion) { throw _invalidSetupResponse(type); } @@ -809,15 +808,7 @@ class WatchTogetherPeerService with KeepaliveMixin { /// guests release their reserved reconnect identity. If transport was lost, /// authenticate a fresh connection first so an intentional exit is not /// mistaken for a transient disconnect. - Future releaseSession() { - final active = _releaseFuture; - if (active != null) return active; - final operation = _releaseSession(); - _releaseFuture = operation; - return operation.whenComplete(() { - if (identical(_releaseFuture, operation)) _releaseFuture = null; - }); - } + Future releaseSession() => _release.run(_releaseSession); Future _releaseSession() async { if (_sessionId == null || _myPeerId == null || _reconnectToken == null) return; diff --git a/relay_protocol.json b/relay_protocol.json index 3b66cebc9..7eed08af9 100644 --- a/relay_protocol.json +++ b/relay_protocol.json @@ -42,5 +42,6 @@ "maxSessionIdLength": 64, "maxPeerIdLength": 128 }, - "idPattern": "^[A-Za-z0-9_-]+$" + "idPattern": "^[A-Za-z0-9_-]+$", + "reconnectTokenBytes": 32 } diff --git a/scripts/codegen/generate_relay_protocol.py b/scripts/codegen/generate_relay_protocol.py index e748425d4..a52959d7a 100755 --- a/scripts/codegen/generate_relay_protocol.py +++ b/scripts/codegen/generate_relay_protocol.py @@ -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", diff --git a/scripts/codegen/test_generate_relay_protocol.py b/scripts/codegen/test_generate_relay_protocol.py index 7a07ceaba..75ecee11d 100644 --- a/scripts/codegen/test_generate_relay_protocol.py +++ b/scripts/codegen/test_generate_relay_protocol.py @@ -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() diff --git a/server/main.go b/server/main.go index 9f87dbfdb..269b04cdd 100644 --- a/server/main.go +++ b/server/main.go @@ -75,7 +75,6 @@ const ( maxRetainedRooms = 2000 connRateBurst = 5 connRateSustained = 1 - reconnectTokenSize = 32 snapshotFormatVersion = 4 snapshotDebounce = 100 * time.Millisecond snapshotFlushTimeout = 5 * time.Second @@ -1251,7 +1250,7 @@ func (s *Server) loadSnapshot(path string) (bool, error) { skipped++ continue } - if r.ProtocolVersion != legacyRelayProtocolVersion && r.ProtocolVersion != relayProtocolVersion { + if !supportedRelayProtocolVersion(r.ProtocolVersion) { skipped++ continue } @@ -1784,7 +1783,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid sessionId or peerId"}) continue } - if msg.ProtocolVersion != legacyRelayProtocolVersion && msg.ProtocolVersion != relayProtocolVersion { + if !supportedRelayProtocolVersion(msg.ProtocolVersion) { client.sendJSON(serverMsg{ Type: relayTypeError, Code: relayErrorProtocolMismatch, @@ -1930,7 +1929,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid sessionId or peerId"}) continue } - if msg.ProtocolVersion != legacyRelayProtocolVersion && msg.ProtocolVersion != relayProtocolVersion { + if !supportedRelayProtocolVersion(msg.ProtocolVersion) { client.sendJSON(serverMsg{ Type: relayTypeError, Code: relayErrorProtocolMismatch, diff --git a/server/main_test.go b/server/main_test.go index c917511ff..53dd7d366 100644 --- a/server/main_test.go +++ b/server/main_test.go @@ -33,6 +33,7 @@ func TestGeneratedRelayProtocolVersionsMatchSpec(t *testing.T) { var spec struct { ProtocolVersion int `json:"protocolVersion"` LegacyProtocolVersion int `json:"legacyProtocolVersion"` + ReconnectTokenBytes int `json:"reconnectTokenBytes"` } if err := json.Unmarshal(data, &spec); err != nil { t.Fatalf("decode relay protocol spec: %v", err) @@ -47,6 +48,19 @@ func TestGeneratedRelayProtocolVersionsMatchSpec(t *testing.T) { spec.LegacyProtocolVersion, ) } + if reconnectTokenSize != spec.ReconnectTokenBytes { + t.Fatalf("generated reconnectTokenSize=%d, spec=%d", reconnectTokenSize, spec.ReconnectTokenBytes) + } + for _, version := range []int{legacyRelayProtocolVersion, relayProtocolVersion} { + if !supportedRelayProtocolVersion(version) { + t.Fatalf("supportedRelayProtocolVersion(%d)=false", version) + } + } + for _, version := range []int{relayProtocolVersion + 1, -1} { + if supportedRelayProtocolVersion(version) { + t.Fatalf("supportedRelayProtocolVersion(%d)=true", version) + } + } } // newTestServer builds a goroutine-free, network-free test server. diff --git a/server/relay_protocol_gen.go b/server/relay_protocol_gen.go index 144ca1ed7..01d81f18a 100644 --- a/server/relay_protocol_gen.go +++ b/server/relay_protocol_gen.go @@ -40,8 +40,13 @@ const ( maxMessageSize = 65536 maxSessionIDLength = 64 maxPeerIDLength = 128 + reconnectTokenSize = 32 ) +func supportedRelayProtocolVersion(version int) bool { + return version == legacyRelayProtocolVersion || version == relayProtocolVersion +} + func validRelayID(value string, maxLength int) bool { if len(value) == 0 || len(value) > maxLength { return false