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