From 43a8fe020d6a7812b022409b465ad4638975ffac Mon Sep 17 00:00:00 2001 From: edde746 <86283021+edde746@users.noreply.github.com> Date: Thu, 23 Jul 2026 02:18:03 +0200 Subject: [PATCH] fix(relay): secure reconnect and room ownership --- .github/workflows/ci.yml | 1 + lib/i18n/bg.i18n.json | 1 + lib/i18n/da.i18n.json | 1 + lib/i18n/de.i18n.json | 1 + lib/i18n/en.i18n.json | 1 + lib/i18n/es.i18n.json | 1 + lib/i18n/fr.i18n.json | 1 + lib/i18n/it.i18n.json | 1 + lib/i18n/ja.i18n.json | 1 + lib/i18n/ko.i18n.json | 1 + lib/i18n/nb.i18n.json | 1 + lib/i18n/nl.i18n.json | 1 + lib/i18n/pl.i18n.json | 1 + lib/i18n/pt.i18n.json | 1 + lib/i18n/ru.i18n.json | 1 + lib/i18n/strings.g.dart | 2 +- lib/i18n/strings_bg.g.dart | 6 +- lib/i18n/strings_da.g.dart | 6 +- lib/i18n/strings_de.g.dart | 6 +- lib/i18n/strings_en.g.dart | 8 +- lib/i18n/strings_es.g.dart | 6 +- lib/i18n/strings_fr.g.dart | 6 +- lib/i18n/strings_it.g.dart | 6 +- lib/i18n/strings_ja.g.dart | 6 +- lib/i18n/strings_ko.g.dart | 6 +- lib/i18n/strings_nb.g.dart | 6 +- lib/i18n/strings_nl.g.dart | 6 +- lib/i18n/strings_pl.g.dart | 6 +- lib/i18n/strings_pt.g.dart | 6 +- lib/i18n/strings_ru.g.dart | 6 +- lib/i18n/strings_sv.g.dart | 6 +- lib/i18n/strings_zh.g.dart | 6 +- lib/i18n/sv.i18n.json | 1 + lib/i18n/zh.i18n.json | 1 + lib/providers/companion_remote_provider.dart | 451 +- .../mobile_remote_screen.dart | 11 +- lib/screens/settings/logs_screen.dart | 43 +- lib/screens/settings/settings_screen.dart | 28 +- .../companion_remote_peer_service.dart | 1040 +++- lib/services/settings_service.dart | 37 +- lib/services/trackers/oauth_proxy_client.dart | 4 +- lib/watch_together/models/sync_message.dart | 4 +- lib/watch_together/primitives.dart | 2 - .../providers/watch_together_provider.dart | 277 +- .../screens/watch_together_screen.dart | 72 +- .../services/host_playback_coordinator.dart | 8 +- .../services/recent_rooms_service.dart | 82 +- .../services/recent_rooms_service.g.dart | 2 + .../services/relay_protocol.g.dart | 9 + .../services/watch_together_controller.dart | 14 +- .../services/watch_together_peer_service.dart | 485 +- .../watch_together_relay_endpoint.dart | 85 + .../widgets/watch_together_overlay.dart | 6 +- relay_protocol.json | 14 +- scripts/ci_checks.sh | 1 + scripts/generate_relay_protocol.py | 40 +- scripts/test_generate_relay_protocol.py | 72 + server/client_ip.go | 149 + server/docker-compose.yml | 1 + server/main.go | 1673 +++++-- server/main_test.go | 4288 ++++++++++++++++- server/oauth.go | 260 +- server/oauth_test.go | 554 ++- server/rate_limit.go | 123 +- server/relay_protocol_gen.go | 47 +- test/models/json_model_round_trip_test.dart | 25 +- .../profile_session_screen_test.dart | 28 +- .../companion_remote_provider_test.dart | 586 ++- test/screens/settings/logs_screen_test.dart | 124 + .../settings/settings_screen_test.dart | 33 + .../companion_remote_peer_service_test.dart | 744 ++- .../settings_export_service_test.dart | 8 +- test/services/settings_service_test.dart | 80 + test/test_helpers/watch_together_fakes.dart | 6 +- .../host_playback_coordinator_test.dart | 130 +- test/watch_together/playback_state_test.dart | 7 +- test/watch_together/primitives_test.dart | 6 - .../recent_rooms_service_test.dart | 116 + .../watch_together_controller_test.dart | 144 +- .../watch_together_overlay_test.dart | 22 +- .../watch_together_peer_service_test.dart | 1007 +++- .../watch_together_provider_test.dart | 658 +++ 82 files changed, 12341 insertions(+), 1382 deletions(-) create mode 100644 lib/watch_together/services/watch_together_relay_endpoint.dart create mode 100644 scripts/test_generate_relay_protocol.py create mode 100644 server/client_ip.go create mode 100644 test/watch_together/recent_rooms_service_test.dart diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d162aa51..f88aaaf0 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -58,6 +58,7 @@ jobs: python3 scripts/check_workflow_action_pins.py python3 scripts/test_check_workflow_action_pins.py python3 scripts/test_check_codegen.py + python3 scripts/test_generate_relay_protocol.py python3 scripts/test_format_native.py python3 scripts/test_run_maestro.py python3 scripts/check_update_packages_workflow.py diff --git a/lib/i18n/bg.i18n.json b/lib/i18n/bg.i18n.json index 17ddfe3b..66d4f49d 100644 --- a/lib/i18n/bg.i18n.json +++ b/lib/i18n/bg.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Релей сървър за гледане заедно", "watchTogetherRelayDescription": "Задай собствен релей сървър. Всички трябва да използват същия сървър.", "watchTogetherRelayHint": "https://my-relay.example.com", + "watchTogetherRelayInvalid": "", "crashReporting": "Докладване на сривове", "crashReportingDescription": "Изпращай доклади за сривове, за да помогнеш за подобряване на приложението", "debugLogging": "Логове за отстраняване на грешки", diff --git a/lib/i18n/da.i18n.json b/lib/i18n/da.i18n.json index ce698ea4..3cfdc961 100644 --- a/lib/i18n/da.i18n.json +++ b/lib/i18n/da.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Watch Together-relay", "watchTogetherRelayDescription": "Angiv en brugerdefineret relay. Alle skal bruge samme server.", "watchTogetherRelayHint": "https://min-relay.eksempel.dk", + "watchTogetherRelayInvalid": "", "crashReporting": "Fejlrapportering", "crashReportingDescription": "Send fejlrapporter for at hjælpe med at forbedre appen", "debugLogging": "Fejlfindingslogning", diff --git a/lib/i18n/de.i18n.json b/lib/i18n/de.i18n.json index 23f85a7a..9636482c 100644 --- a/lib/i18n/de.i18n.json +++ b/lib/i18n/de.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Gemeinsam Schauen Relay", "watchTogetherRelayDescription": "Eigenes Relay festlegen. Alle müssen denselben Server verwenden.", "watchTogetherRelayHint": "https://mein-relay.beispiel.de", + "watchTogetherRelayInvalid": "", "crashReporting": "Absturzberichte", "crashReportingDescription": "Absturzberichte senden, um die App zu verbessern", "debugLogging": "Debug-Protokollierung", diff --git a/lib/i18n/en.i18n.json b/lib/i18n/en.i18n.json index bb3d6c4d..52b3a29d 100644 --- a/lib/i18n/en.i18n.json +++ b/lib/i18n/en.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Watch Together Relay", "watchTogetherRelayDescription": "Set a custom relay. Everyone must use the same server.", "watchTogetherRelayHint": "https://my-relay.example.com", + "watchTogetherRelayInvalid": "Enter a valid HTTP or HTTPS relay base URL.", "crashReporting": "Crash Reporting", "crashReportingDescription": "Send crash reports to help improve the app", "debugLogging": "Debug Logging", diff --git a/lib/i18n/es.i18n.json b/lib/i18n/es.i18n.json index 6666d652..2cf4344c 100644 --- a/lib/i18n/es.i18n.json +++ b/lib/i18n/es.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Relay de Ver Juntos", "watchTogetherRelayDescription": "Define un relay personalizado. Todos deben usar el mismo servidor.", "watchTogetherRelayHint": "https://mi-relay.ejemplo.com", + "watchTogetherRelayInvalid": "", "crashReporting": "Informes de Errores", "crashReportingDescription": "Enviar informes de errores para mejorar la aplicación", "debugLogging": "Registro de Depuración", diff --git a/lib/i18n/fr.i18n.json b/lib/i18n/fr.i18n.json index ecacf40c..ba1b77f9 100644 --- a/lib/i18n/fr.i18n.json +++ b/lib/i18n/fr.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Relais Regarder Ensemble", "watchTogetherRelayDescription": "Définir un relay personnalisé. Tout le monde doit utiliser le même serveur.", "watchTogetherRelayHint": "https://mon-relais.exemple.fr", + "watchTogetherRelayInvalid": "", "crashReporting": "Rapports de plantage", "crashReportingDescription": "Envoyer des rapports de plantage pour améliorer l'application", "debugLogging": "Journalisation de débogage", diff --git a/lib/i18n/it.i18n.json b/lib/i18n/it.i18n.json index 63f21978..6eba0dbf 100644 --- a/lib/i18n/it.i18n.json +++ b/lib/i18n/it.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Relay Guarda Insieme", "watchTogetherRelayDescription": "Imposta un relay personalizzato. Tutti devono usare lo stesso server.", "watchTogetherRelayHint": "https://mio-relay.esempio.it", + "watchTogetherRelayInvalid": "", "crashReporting": "Segnalazione errori", "crashReportingDescription": "Invia segnalazioni di errori per migliorare l'app", "debugLogging": "Log di debug", diff --git a/lib/i18n/ja.i18n.json b/lib/i18n/ja.i18n.json index 64858c91..b63f6cd8 100644 --- a/lib/i18n/ja.i18n.json +++ b/lib/i18n/ja.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "一緒に視聴リレーサーバー", "watchTogetherRelayDescription": "カスタムリレーを設定します。全員が同じサーバーを使う必要があります。", "watchTogetherRelayHint": "https://my-relay.example.com", + "watchTogetherRelayInvalid": "", "crashReporting": "クラッシュレポート", "crashReportingDescription": "アプリの改善に役立つクラッシュレポートを送信", "debugLogging": "デバッグログ", diff --git a/lib/i18n/ko.i18n.json b/lib/i18n/ko.i18n.json index 25413102..0e955d95 100644 --- a/lib/i18n/ko.i18n.json +++ b/lib/i18n/ko.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "함께 보기 릴레이", "watchTogetherRelayDescription": "사용자 지정 릴레이를 설정합니다. 모두 같은 서버를 사용해야 합니다.", "watchTogetherRelayHint": "https://my-relay.example.com", + "watchTogetherRelayInvalid": "", "crashReporting": "충돌 보고", "crashReportingDescription": "앱 개선을 위해 충돌 보고서 전송", "debugLogging": "디버그 로깅", diff --git a/lib/i18n/nb.i18n.json b/lib/i18n/nb.i18n.json index 146ff8ff..d7660f76 100644 --- a/lib/i18n/nb.i18n.json +++ b/lib/i18n/nb.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Se Sammen-relay", "watchTogetherRelayDescription": "Angi en egendefinert relay. Alle må bruke samme server.", "watchTogetherRelayHint": "https://min-relay.eksempel.no", + "watchTogetherRelayInvalid": "", "crashReporting": "Krasjrapportering", "crashReportingDescription": "Send krasjrapporter for å hjelpe med å forbedre appen", "debugLogging": "Feilsøkingslogging", diff --git a/lib/i18n/nl.i18n.json b/lib/i18n/nl.i18n.json index f378c411..7e369c3d 100644 --- a/lib/i18n/nl.i18n.json +++ b/lib/i18n/nl.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Samen Kijken Relay", "watchTogetherRelayDescription": "Stel een aangepaste relay in. Iedereen moet dezelfde server gebruiken.", "watchTogetherRelayHint": "https://mijn-relay.voorbeeld.nl", + "watchTogetherRelayInvalid": "", "crashReporting": "Crashrapportage", "crashReportingDescription": "Crashrapporten verzenden om de app te verbeteren", "debugLogging": "Debug logging", diff --git a/lib/i18n/pl.i18n.json b/lib/i18n/pl.i18n.json index 4f802842..72038214 100644 --- a/lib/i18n/pl.i18n.json +++ b/lib/i18n/pl.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Relay Oglądaj Razem", "watchTogetherRelayDescription": "Ustaw własny relay. Wszyscy muszą używać tego samego serwera.", "watchTogetherRelayHint": "https://moj-relay.przyklad.pl", + "watchTogetherRelayInvalid": "", "crashReporting": "Raportowanie błędów", "crashReportingDescription": "Wysyłaj raporty o błędach, aby pomóc ulepszyć aplikację", "debugLogging": "Logowanie debugowania", diff --git a/lib/i18n/pt.i18n.json b/lib/i18n/pt.i18n.json index be126208..2cd061e3 100644 --- a/lib/i18n/pt.i18n.json +++ b/lib/i18n/pt.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Relay do Assistir Juntos", "watchTogetherRelayDescription": "Defina um relay personalizado. Todos devem usar o mesmo servidor.", "watchTogetherRelayHint": "https://meu-relay.exemplo.com.br", + "watchTogetherRelayInvalid": "", "crashReporting": "Relatório de Erros", "crashReportingDescription": "Enviar relatórios de erros para ajudar a melhorar o app", "debugLogging": "Log de Depuração", diff --git a/lib/i18n/ru.i18n.json b/lib/i18n/ru.i18n.json index 01bcdd26..2929356f 100644 --- a/lib/i18n/ru.i18n.json +++ b/lib/i18n/ru.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Relay совместного просмотра", "watchTogetherRelayDescription": "Задайте свой relay. Все должны использовать один сервер.", "watchTogetherRelayHint": "https://my-relay.example.com", + "watchTogetherRelayInvalid": "", "crashReporting": "Отчёты об ошибках", "crashReportingDescription": "Отправлять отчёты об ошибках для улучшения приложения", "debugLogging": "Журнал отладки", diff --git a/lib/i18n/strings.g.dart b/lib/i18n/strings.g.dart index f5372e74..901990ae 100644 --- a/lib/i18n/strings.g.dart +++ b/lib/i18n/strings.g.dart @@ -4,7 +4,7 @@ /// To regenerate, run: `dart run slang` /// /// Locales: 16 -/// Strings: 23011 (1438 per locale) +/// Strings: 23027 (1439 per locale) // coverage:ignore-file // ignore_for_file: type=lint, unused_import diff --git a/lib/i18n/strings_bg.g.dart b/lib/i18n/strings_bg.g.dart index 183f5d2b..acd97fdc 100644 --- a/lib/i18n/strings_bg.g.dart +++ b/lib/i18n/strings_bg.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsBg extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Релей сървър за гледане заедно'; @override String get watchTogetherRelayDescription => 'Задай собствен релей сървър. Всички трябва да използват същия сървър.'; @override String get watchTogetherRelayHint => 'https://my-relay.example.com'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Докладване на сривове'; @override String get crashReportingDescription => 'Изпращай доклади за сривове, за да помогнеш за подобряване на приложението'; @override String get debugLogging => 'Логове за отстраняване на грешки'; @@ -2303,6 +2304,7 @@ extension on TranslationsBg { 'settings.watchTogetherRelay' => 'Релей сървър за гледане заедно', 'settings.watchTogetherRelayDescription' => 'Задай собствен релей сървър. Всички трябва да използват същия сървър.', 'settings.watchTogetherRelayHint' => 'https://my-relay.example.com', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Докладване на сривове', 'settings.crashReportingDescription' => 'Изпращай доклади за сривове, за да помогнеш за подобряване на приложението', 'settings.debugLogging' => 'Логове за отстраняване на грешки', @@ -2644,9 +2646,9 @@ extension on TranslationsBg { 'messages.markedAsUnwatchedOffline' => 'Маркирано като негледано (ще се синхронизира, когато сте онлайн)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Автоматично премахнато: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('bg'))(n, one: 'Автоматично премахнато ${n} гледано изтегляне', other: 'Автоматично премахнати ${n} гледани изтегляния', ), - 'messages.removedFromContinueWatching' => 'Премахнато от продължаване на гледането', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Премахнато от продължаване на гледането', 'messages.errorLoading' => ({required Object error}) => 'Грешка: ${error}', 'messages.streamInterrupted' => 'Потокът прекъсна. Натиснете „Пусни“ или превъртете, за да опитате отново.', 'messages.liveStreamInterrupted' => 'Потокът на живо прекъсна. Натиснете „Пусни“, за да опитате отново.', @@ -3158,9 +3160,9 @@ extension on TranslationsBg { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} превъртя', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} буферира', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} е с по-стара версия на приложението — синхронизирането не е налично', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Продължаване без ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Продължаване без ${name}', 'watchTogether.waitingForParticipants' => 'Изчакване другите да заредят...', 'watchTogether.waitingForName' => ({required Object name}) => 'Изчакване на ${name}...', 'watchTogether.recentRooms' => 'Скорошни стаи', diff --git a/lib/i18n/strings_da.g.dart b/lib/i18n/strings_da.g.dart index 0a99230d..811d3063 100644 --- a/lib/i18n/strings_da.g.dart +++ b/lib/i18n/strings_da.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsDa extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Watch Together-relay'; @override String get watchTogetherRelayDescription => 'Angiv en brugerdefineret relay. Alle skal bruge samme server.'; @override String get watchTogetherRelayHint => 'https://min-relay.eksempel.dk'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Fejlrapportering'; @override String get crashReportingDescription => 'Send fejlrapporter for at hjælpe med at forbedre appen'; @override String get debugLogging => 'Fejlfindingslogning'; @@ -2303,6 +2304,7 @@ extension on TranslationsDa { 'settings.watchTogetherRelay' => 'Watch Together-relay', 'settings.watchTogetherRelayDescription' => 'Angiv en brugerdefineret relay. Alle skal bruge samme server.', 'settings.watchTogetherRelayHint' => 'https://min-relay.eksempel.dk', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Fejlrapportering', 'settings.crashReportingDescription' => 'Send fejlrapporter for at hjælpe med at forbedre appen', 'settings.debugLogging' => 'Fejlfindingslogning', @@ -2644,9 +2646,9 @@ extension on TranslationsDa { 'messages.markedAsUnwatchedOffline' => 'Markeret som uset (synkroniseres online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Automatisk fjernet: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('da'))(n, one: 'Fjernede automatisk ${n} set download', other: 'Fjernede automatisk ${n} sete downloads', ), - 'messages.removedFromContinueWatching' => 'Fjernet fra Fortsæt med at se', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Fjernet fra Fortsæt med at se', 'messages.errorLoading' => ({required Object error}) => 'Fejl: ${error}', 'messages.streamInterrupted' => 'Streamen blev afbrudt. Tryk på afspil, eller spol for at prøve igen.', 'messages.liveStreamInterrupted' => 'Livestreamen blev afbrudt. Tryk på afspil for at prøve igen.', @@ -3158,9 +3160,9 @@ extension on TranslationsDa { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} spoled', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} bufferer', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} bruger en ældre appversion — synkronisering er ikke tilgængelig', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Fortsætter uden ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Fortsætter uden ${name}', 'watchTogether.waitingForParticipants' => 'Venter på at andre indlæser...', 'watchTogether.waitingForName' => ({required Object name}) => 'Venter på ${name}...', 'watchTogether.recentRooms' => 'Seneste rum', diff --git a/lib/i18n/strings_de.g.dart b/lib/i18n/strings_de.g.dart index 9c0fe6aa..0672887d 100644 --- a/lib/i18n/strings_de.g.dart +++ b/lib/i18n/strings_de.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsDe extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Gemeinsam Schauen Relay'; @override String get watchTogetherRelayDescription => 'Eigenes Relay festlegen. Alle müssen denselben Server verwenden.'; @override String get watchTogetherRelayHint => 'https://mein-relay.beispiel.de'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Absturzberichte'; @override String get crashReportingDescription => 'Absturzberichte senden, um die App zu verbessern'; @override String get debugLogging => 'Debug-Protokollierung'; @@ -2303,6 +2304,7 @@ extension on TranslationsDe { 'settings.watchTogetherRelay' => 'Gemeinsam Schauen Relay', 'settings.watchTogetherRelayDescription' => 'Eigenes Relay festlegen. Alle müssen denselben Server verwenden.', 'settings.watchTogetherRelayHint' => 'https://mein-relay.beispiel.de', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Absturzberichte', 'settings.crashReportingDescription' => 'Absturzberichte senden, um die App zu verbessern', 'settings.debugLogging' => 'Debug-Protokollierung', @@ -2644,9 +2646,9 @@ extension on TranslationsDe { 'messages.markedAsUnwatchedOffline' => 'Als ungesehen markiert (wird synchronisiert, wenn online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Automatisch entfernt: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('de'))(n, one: 'Automatisch entfernt: ${n} angesehener Download', other: 'Automatisch entfernt: ${n} angesehene Downloads', ), - 'messages.removedFromContinueWatching' => 'Aus ‚Weiterschauen\' entfernt', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Aus ‚Weiterschauen\' entfernt', 'messages.errorLoading' => ({required Object error}) => 'Fehler: ${error}', 'messages.streamInterrupted' => 'Der Stream wurde unterbrochen. Drücke auf Wiedergabe oder spule, um es erneut zu versuchen.', 'messages.liveStreamInterrupted' => 'Der Livestream wurde unterbrochen. Drücke auf Wiedergabe, um es erneut zu versuchen.', @@ -3158,9 +3160,9 @@ extension on TranslationsDe { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} hat gespult', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} puffert', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} verwendet eine ältere Appversion — Synchronisierung nicht verfügbar', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Fortfahren ohne ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Fortfahren ohne ${name}', 'watchTogether.waitingForParticipants' => 'Warte auf andere zum Laden...', 'watchTogether.waitingForName' => ({required Object name}) => 'Warten auf ${name}...', 'watchTogether.recentRooms' => 'Letzte Räume', diff --git a/lib/i18n/strings_en.g.dart b/lib/i18n/strings_en.g.dart index bf917ab7..68d61429 100644 --- a/lib/i18n/strings_en.g.dart +++ b/lib/i18n/strings_en.g.dart @@ -653,6 +653,9 @@ class TranslationsSettingsEn { /// en: 'https://my-relay.example.com' String get watchTogetherRelayHint => 'https://my-relay.example.com'; + /// en: 'Enter a valid HTTP or HTTPS relay base URL.' + String get watchTogetherRelayInvalid => 'Enter a valid HTTP or HTTPS relay base URL.'; + /// en: 'Crash Reporting' String get crashReporting => 'Crash Reporting'; @@ -5180,6 +5183,7 @@ extension on Translations { 'settings.watchTogetherRelay' => 'Watch Together Relay', 'settings.watchTogetherRelayDescription' => 'Set a custom relay. Everyone must use the same server.', 'settings.watchTogetherRelayHint' => 'https://my-relay.example.com', + 'settings.watchTogetherRelayInvalid' => 'Enter a valid HTTP or HTTPS relay base URL.', 'settings.crashReporting' => 'Crash Reporting', 'settings.crashReportingDescription' => 'Send crash reports to help improve the app', 'settings.debugLogging' => 'Debug Logging', @@ -5522,9 +5526,9 @@ extension on Translations { 'messages.markedAsUnwatchedOffline' => 'Marked as unwatched (will sync when online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Auto-removed: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('en'))(n, one: 'Auto-removed ${n} watched download', other: 'Auto-removed ${n} watched downloads', ), - 'messages.removedFromContinueWatching' => 'Removed from Continue Watching', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Removed from Continue Watching', 'messages.errorLoading' => ({required Object error}) => 'Error: ${error}', 'messages.streamInterrupted' => 'The stream was interrupted. Press play or seek to retry.', 'messages.liveStreamInterrupted' => 'The live stream was interrupted. Press play to retry.', @@ -6037,9 +6041,9 @@ extension on Translations { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} seeked', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} is buffering', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} is on an older app version — sync unavailable', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Resuming without ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Resuming without ${name}', 'watchTogether.waitingForParticipants' => 'Waiting for others to load...', 'watchTogether.waitingForName' => ({required Object name}) => 'Waiting for ${name}...', 'watchTogether.recentRooms' => 'Recent Rooms', diff --git a/lib/i18n/strings_es.g.dart b/lib/i18n/strings_es.g.dart index 6187f2f0..9571b1d4 100644 --- a/lib/i18n/strings_es.g.dart +++ b/lib/i18n/strings_es.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsEs extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Relay de Ver Juntos'; @override String get watchTogetherRelayDescription => 'Define un relay personalizado. Todos deben usar el mismo servidor.'; @override String get watchTogetherRelayHint => 'https://mi-relay.ejemplo.com'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Informes de Errores'; @override String get crashReportingDescription => 'Enviar informes de errores para mejorar la aplicación'; @override String get debugLogging => 'Registro de Depuración'; @@ -2303,6 +2304,7 @@ extension on TranslationsEs { 'settings.watchTogetherRelay' => 'Relay de Ver Juntos', 'settings.watchTogetherRelayDescription' => 'Define un relay personalizado. Todos deben usar el mismo servidor.', 'settings.watchTogetherRelayHint' => 'https://mi-relay.ejemplo.com', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Informes de Errores', 'settings.crashReportingDescription' => 'Enviar informes de errores para mejorar la aplicación', 'settings.debugLogging' => 'Registro de Depuración', @@ -2644,9 +2646,9 @@ extension on TranslationsEs { 'messages.markedAsUnwatchedOffline' => 'Marcado como no visto (se sincronizará al estar en línea)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Eliminado automáticamente: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('es'))(n, one: 'Se eliminó automáticamente ${n} descarga vista', other: 'Se eliminaron automáticamente ${n} descargas vistas', ), - 'messages.removedFromContinueWatching' => 'Eliminado de Seguir Viendo', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Eliminado de Seguir Viendo', 'messages.errorLoading' => ({required Object error}) => 'Error: ${error}', 'messages.streamInterrupted' => 'La reproducción se interrumpió. Pulsa reproducir o avanza para volver a intentarlo.', 'messages.liveStreamInterrupted' => 'La transmisión en vivo se interrumpió. Pulsa reproducir para volver a intentarlo.', @@ -3158,9 +3160,9 @@ extension on TranslationsEs { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} avanzó', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} está cargando', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} usa una versión anterior de la app — sincronización no disponible', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Reanudando sin ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Reanudando sin ${name}', 'watchTogether.waitingForParticipants' => 'Esperando a que otros carguen...', 'watchTogether.waitingForName' => ({required Object name}) => 'Esperando a ${name}...', 'watchTogether.recentRooms' => 'Salas recientes', diff --git a/lib/i18n/strings_fr.g.dart b/lib/i18n/strings_fr.g.dart index 215f55ab..286fed9e 100644 --- a/lib/i18n/strings_fr.g.dart +++ b/lib/i18n/strings_fr.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsFr extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Relais Regarder Ensemble'; @override String get watchTogetherRelayDescription => 'Définir un relay personnalisé. Tout le monde doit utiliser le même serveur.'; @override String get watchTogetherRelayHint => 'https://mon-relais.exemple.fr'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Rapports de plantage'; @override String get crashReportingDescription => 'Envoyer des rapports de plantage pour améliorer l\'application'; @override String get debugLogging => 'Journalisation de débogage'; @@ -2303,6 +2304,7 @@ extension on TranslationsFr { 'settings.watchTogetherRelay' => 'Relais Regarder Ensemble', 'settings.watchTogetherRelayDescription' => 'Définir un relay personnalisé. Tout le monde doit utiliser le même serveur.', 'settings.watchTogetherRelayHint' => 'https://mon-relais.exemple.fr', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Rapports de plantage', 'settings.crashReportingDescription' => 'Envoyer des rapports de plantage pour améliorer l\'application', 'settings.debugLogging' => 'Journalisation de débogage', @@ -2644,9 +2646,9 @@ extension on TranslationsFr { 'messages.markedAsUnwatchedOffline' => 'Marqué comme non vu (sera synchronisé lorsque vous serez en ligne)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Supprimé automatiquement : ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('fr'))(n, one: '${n} téléchargement vu supprimé automatiquement', other: '${n} téléchargements vus supprimés automatiquement', ), - 'messages.removedFromContinueWatching' => 'Supprimer de "Continuer à regarder"', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Supprimer de "Continuer à regarder"', 'messages.errorLoading' => ({required Object error}) => 'Erreur: ${error}', 'messages.streamInterrupted' => 'La lecture a été interrompue. Appuyez sur Lecture ou avancez pour réessayer.', 'messages.liveStreamInterrupted' => 'Le direct a été interrompu. Appuyez sur Lecture pour réessayer.', @@ -3158,9 +3160,9 @@ extension on TranslationsFr { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} a avancé', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} met en mémoire tampon', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} utilise une ancienne version de l’app — synchronisation indisponible', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Reprise sans ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Reprise sans ${name}', 'watchTogether.waitingForParticipants' => 'En attente du chargement des autres...', 'watchTogether.waitingForName' => ({required Object name}) => 'En attente de ${name}...', 'watchTogether.recentRooms' => 'Salons récents', diff --git a/lib/i18n/strings_it.g.dart b/lib/i18n/strings_it.g.dart index 24d0d268..f6cbcd27 100644 --- a/lib/i18n/strings_it.g.dart +++ b/lib/i18n/strings_it.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsIt extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Relay Guarda Insieme'; @override String get watchTogetherRelayDescription => 'Imposta un relay personalizzato. Tutti devono usare lo stesso server.'; @override String get watchTogetherRelayHint => 'https://mio-relay.esempio.it'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Segnalazione errori'; @override String get crashReportingDescription => 'Invia segnalazioni di errori per migliorare l\'app'; @override String get debugLogging => 'Log di debug'; @@ -2303,6 +2304,7 @@ extension on TranslationsIt { 'settings.watchTogetherRelay' => 'Relay Guarda Insieme', 'settings.watchTogetherRelayDescription' => 'Imposta un relay personalizzato. Tutti devono usare lo stesso server.', 'settings.watchTogetherRelayHint' => 'https://mio-relay.esempio.it', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Segnalazione errori', 'settings.crashReportingDescription' => 'Invia segnalazioni di errori per migliorare l\'app', 'settings.debugLogging' => 'Log di debug', @@ -2644,9 +2646,9 @@ extension on TranslationsIt { 'messages.markedAsUnwatchedOffline' => 'Segnato come non visto (sincronizzato online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Rimosso automaticamente: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('it'))(n, one: 'Rimosso automaticamente ${n} download già visto', other: 'Rimossi automaticamente ${n} download già visti', ), - 'messages.removedFromContinueWatching' => 'Rimosso da Continua a guardare', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Rimosso da Continua a guardare', 'messages.errorLoading' => ({required Object error}) => 'Errore: ${error}', 'messages.streamInterrupted' => 'La riproduzione si è interrotta. Premi Riproduci o scorri per riprovare.', 'messages.liveStreamInterrupted' => 'La diretta si è interrotta. Premi Riproduci per riprovare.', @@ -3158,9 +3160,9 @@ extension on TranslationsIt { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} ha cercato', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} sta caricando', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} usa una versione precedente dell\'app — sincronizzazione non disponibile', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Ripresa senza ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Ripresa senza ${name}', 'watchTogether.waitingForParticipants' => 'In attesa che gli altri carichino...', 'watchTogether.waitingForName' => ({required Object name}) => 'In attesa di ${name}...', 'watchTogether.recentRooms' => 'Stanze recenti', diff --git a/lib/i18n/strings_ja.g.dart b/lib/i18n/strings_ja.g.dart index 0b89634e..b48ce00e 100644 --- a/lib/i18n/strings_ja.g.dart +++ b/lib/i18n/strings_ja.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsJa extends TranslationsSettingsEn { @override String get watchTogetherRelay => '一緒に視聴リレーサーバー'; @override String get watchTogetherRelayDescription => 'カスタムリレーを設定します。全員が同じサーバーを使う必要があります。'; @override String get watchTogetherRelayHint => 'https://my-relay.example.com'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'クラッシュレポート'; @override String get crashReportingDescription => 'アプリの改善に役立つクラッシュレポートを送信'; @override String get debugLogging => 'デバッグログ'; @@ -2300,6 +2301,7 @@ extension on TranslationsJa { 'settings.watchTogetherRelay' => '一緒に視聴リレーサーバー', 'settings.watchTogetherRelayDescription' => 'カスタムリレーを設定します。全員が同じサーバーを使う必要があります。', 'settings.watchTogetherRelayHint' => 'https://my-relay.example.com', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'クラッシュレポート', 'settings.crashReportingDescription' => 'アプリの改善に役立つクラッシュレポートを送信', 'settings.debugLogging' => 'デバッグログ', @@ -2641,9 +2643,9 @@ extension on TranslationsJa { 'messages.markedAsUnwatchedOffline' => '未視聴にしました(オンライン時に同期)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => '自動削除: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('ja'))(n, other: '視聴済みダウンロードを${n}件自動削除しました', ), - 'messages.removedFromContinueWatching' => '視聴中から削除しました', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => '視聴中から削除しました', 'messages.errorLoading' => ({required Object error}) => 'エラー: ${error}', 'messages.streamInterrupted' => 'ストリームが中断されました。再生を押すかシークして再試行してください。', 'messages.liveStreamInterrupted' => 'ライブストリームが中断されました。再生を押して再試行してください。', @@ -3155,9 +3157,9 @@ extension on TranslationsJa { 'watchTogether.participantSeeked' => ({required Object name}) => '${name}がシークしました', 'watchTogether.participantBuffering' => ({required Object name}) => '${name}がバッファリング中', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} は古いバージョンのアプリを使用しているため、同期できません', - 'watchTogether.resumingWithout' => ({required Object name}) => '${name} なしで再開', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => '${name} なしで再開', 'watchTogether.waitingForParticipants' => '他の参加者の読み込みを待っています...', 'watchTogether.waitingForName' => ({required Object name}) => '${name}を待っています...', 'watchTogether.recentRooms' => '最近のルーム', diff --git a/lib/i18n/strings_ko.g.dart b/lib/i18n/strings_ko.g.dart index f59cf7c2..1ad98e9f 100644 --- a/lib/i18n/strings_ko.g.dart +++ b/lib/i18n/strings_ko.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsKo extends TranslationsSettingsEn { @override String get watchTogetherRelay => '함께 보기 릴레이'; @override String get watchTogetherRelayDescription => '사용자 지정 릴레이를 설정합니다. 모두 같은 서버를 사용해야 합니다.'; @override String get watchTogetherRelayHint => 'https://my-relay.example.com'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => '충돌 보고'; @override String get crashReportingDescription => '앱 개선을 위해 충돌 보고서 전송'; @override String get debugLogging => '디버그 로깅'; @@ -2300,6 +2301,7 @@ extension on TranslationsKo { 'settings.watchTogetherRelay' => '함께 보기 릴레이', 'settings.watchTogetherRelayDescription' => '사용자 지정 릴레이를 설정합니다. 모두 같은 서버를 사용해야 합니다.', 'settings.watchTogetherRelayHint' => 'https://my-relay.example.com', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => '충돌 보고', 'settings.crashReportingDescription' => '앱 개선을 위해 충돌 보고서 전송', 'settings.debugLogging' => '디버그 로깅', @@ -2641,9 +2643,9 @@ extension on TranslationsKo { 'messages.markedAsUnwatchedOffline' => '미시청으로 표시됨 (연결 시 동기화됨)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => '자동 삭제됨: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('ko'))(n, other: '시청한 다운로드 ${n}개를 자동 삭제했습니다', ), - 'messages.removedFromContinueWatching' => '계속 시청 목록에서 제거됨', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => '계속 시청 목록에서 제거됨', 'messages.errorLoading' => ({required Object error}) => '오류: ${error}', 'messages.streamInterrupted' => '스트림이 중단되었습니다. 재생을 누르거나 탐색하여 다시 시도하세요.', 'messages.liveStreamInterrupted' => '라이브 스트림이 중단되었습니다. 재생을 눌러 다시 시도하세요.', @@ -3155,9 +3157,9 @@ extension on TranslationsKo { 'watchTogether.participantSeeked' => ({required Object name}) => '${name}님이 탐색했습니다', 'watchTogether.participantBuffering' => ({required Object name}) => '${name}님이 버퍼링 중입니다', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name}님이 이전 버전의 앱을 사용 중입니다 — 동기화를 사용할 수 없습니다', - 'watchTogether.resumingWithout' => ({required Object name}) => '${name}님 없이 재생을 재개합니다', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => '${name}님 없이 재생을 재개합니다', 'watchTogether.waitingForParticipants' => '다른 참가자의 로딩을 기다리는 중...', 'watchTogether.waitingForName' => ({required Object name}) => '${name}님을 기다리는 중...', 'watchTogether.recentRooms' => '최근 방', diff --git a/lib/i18n/strings_nb.g.dart b/lib/i18n/strings_nb.g.dart index 80c89788..2681dce7 100644 --- a/lib/i18n/strings_nb.g.dart +++ b/lib/i18n/strings_nb.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsNb extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Se Sammen-relay'; @override String get watchTogetherRelayDescription => 'Angi en egendefinert relay. Alle må bruke samme server.'; @override String get watchTogetherRelayHint => 'https://min-relay.eksempel.no'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Krasjrapportering'; @override String get crashReportingDescription => 'Send krasjrapporter for å hjelpe med å forbedre appen'; @override String get debugLogging => 'Feilsøkingslogging'; @@ -2303,6 +2304,7 @@ extension on TranslationsNb { 'settings.watchTogetherRelay' => 'Se Sammen-relay', 'settings.watchTogetherRelayDescription' => 'Angi en egendefinert relay. Alle må bruke samme server.', 'settings.watchTogetherRelayHint' => 'https://min-relay.eksempel.no', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Krasjrapportering', 'settings.crashReportingDescription' => 'Send krasjrapporter for å hjelpe med å forbedre appen', 'settings.debugLogging' => 'Feilsøkingslogging', @@ -2644,9 +2646,9 @@ extension on TranslationsNb { 'messages.markedAsUnwatchedOffline' => 'Merket som usett (synkroniseres når tilkoblet)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Automatisk fjernet: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('nb'))(n, one: 'Fjernet automatisk ${n} sett nedlasting', other: 'Fjernet automatisk ${n} sette nedlastinger', ), - 'messages.removedFromContinueWatching' => 'Fjernet fra Fortsett å se', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Fjernet fra Fortsett å se', 'messages.errorLoading' => ({required Object error}) => 'Feil: ${error}', 'messages.streamInterrupted' => 'Avspillingen ble avbrutt. Trykk på Spill av eller spol for å prøve på nytt.', 'messages.liveStreamInterrupted' => 'Direktesendingen ble avbrutt. Trykk på Spill av for å prøve på nytt.', @@ -3158,9 +3160,9 @@ extension on TranslationsNb { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} spolet', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} buffrer', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} bruker en eldre appversjon — synkronisering er ikke tilgjengelig', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Fortsetter uten ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Fortsetter uten ${name}', 'watchTogether.waitingForParticipants' => 'Venter på at andre laster inn...', 'watchTogether.waitingForName' => ({required Object name}) => 'Venter på ${name}...', 'watchTogether.recentRooms' => 'Nylige rom', diff --git a/lib/i18n/strings_nl.g.dart b/lib/i18n/strings_nl.g.dart index 5e0977b7..694031eb 100644 --- a/lib/i18n/strings_nl.g.dart +++ b/lib/i18n/strings_nl.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsNl extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Samen Kijken Relay'; @override String get watchTogetherRelayDescription => 'Stel een aangepaste relay in. Iedereen moet dezelfde server gebruiken.'; @override String get watchTogetherRelayHint => 'https://mijn-relay.voorbeeld.nl'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Crashrapportage'; @override String get crashReportingDescription => 'Crashrapporten verzenden om de app te verbeteren'; @override String get debugLogging => 'Debug logging'; @@ -2303,6 +2304,7 @@ extension on TranslationsNl { 'settings.watchTogetherRelay' => 'Samen Kijken Relay', 'settings.watchTogetherRelayDescription' => 'Stel een aangepaste relay in. Iedereen moet dezelfde server gebruiken.', 'settings.watchTogetherRelayHint' => 'https://mijn-relay.voorbeeld.nl', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Crashrapportage', 'settings.crashReportingDescription' => 'Crashrapporten verzenden om de app te verbeteren', 'settings.debugLogging' => 'Debug logging', @@ -2644,9 +2646,9 @@ extension on TranslationsNl { 'messages.markedAsUnwatchedOffline' => 'Gemarkeerd als ongekeken (sync wanneer online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Automatisch verwijderd: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('nl'))(n, one: 'Automatisch ${n} bekeken download verwijderd', other: 'Automatisch ${n} bekeken downloads verwijderd', ), - 'messages.removedFromContinueWatching' => 'Verwijderd uit Doorgaan met kijken', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Verwijderd uit Doorgaan met kijken', 'messages.errorLoading' => ({required Object error}) => 'Fout: ${error}', 'messages.streamInterrupted' => 'De stream is onderbroken. Druk op afspelen of spoel om het opnieuw te proberen.', 'messages.liveStreamInterrupted' => 'De livestream is onderbroken. Druk op afspelen om het opnieuw te proberen.', @@ -3158,9 +3160,9 @@ extension on TranslationsNl { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} heeft gespoeld', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} is aan het bufferen', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} gebruikt een oudere appversie — synchronisatie niet beschikbaar', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Hervatten zonder ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Hervatten zonder ${name}', 'watchTogether.waitingForParticipants' => 'Wachten tot anderen geladen zijn...', 'watchTogether.waitingForName' => ({required Object name}) => 'Wachten op ${name}...', 'watchTogether.recentRooms' => 'Recente kamers', diff --git a/lib/i18n/strings_pl.g.dart b/lib/i18n/strings_pl.g.dart index d2a2cd87..5a69c6c3 100644 --- a/lib/i18n/strings_pl.g.dart +++ b/lib/i18n/strings_pl.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsPl extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Relay Oglądaj Razem'; @override String get watchTogetherRelayDescription => 'Ustaw własny relay. Wszyscy muszą używać tego samego serwera.'; @override String get watchTogetherRelayHint => 'https://moj-relay.przyklad.pl'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Raportowanie błędów'; @override String get crashReportingDescription => 'Wysyłaj raporty o błędach, aby pomóc ulepszyć aplikację'; @override String get debugLogging => 'Logowanie debugowania'; @@ -2309,6 +2310,7 @@ extension on TranslationsPl { 'settings.watchTogetherRelay' => 'Relay Oglądaj Razem', 'settings.watchTogetherRelayDescription' => 'Ustaw własny relay. Wszyscy muszą używać tego samego serwera.', 'settings.watchTogetherRelayHint' => 'https://moj-relay.przyklad.pl', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Raportowanie błędów', 'settings.crashReportingDescription' => 'Wysyłaj raporty o błędach, aby pomóc ulepszyć aplikację', 'settings.debugLogging' => 'Logowanie debugowania', @@ -2650,9 +2652,9 @@ extension on TranslationsPl { 'messages.markedAsUnwatchedOffline' => 'Oznaczono jako nieobejrzane (zsynchronizuje się po połączeniu)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Automatycznie usunięto: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('pl'))(n, one: 'Automatycznie usunięto ${n} obejrzane pobranie', few: 'Automatycznie usunięto ${n} obejrzane pobrania', many: 'Automatycznie usunięto ${n} obejrzanych pobrań', other: 'Automatycznie usunięto ${n} obejrzanego pobrania', ), - 'messages.removedFromContinueWatching' => 'Usunięto z kontynuowania oglądania', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Usunięto z kontynuowania oglądania', 'messages.errorLoading' => ({required Object error}) => 'Błąd: ${error}', 'messages.streamInterrupted' => 'Strumień został przerwany. Naciśnij odtwarzanie lub przewiń, aby spróbować ponownie.', 'messages.liveStreamInterrupted' => 'Transmisja na żywo została przerwana. Naciśnij odtwarzanie, aby spróbować ponownie.', @@ -3164,9 +3166,9 @@ extension on TranslationsPl { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} przewinął', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} buforuje', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} używa starszej wersji aplikacji — synchronizacja jest niedostępna', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Wznawianie bez ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Wznawianie bez ${name}', 'watchTogether.waitingForParticipants' => 'Oczekiwanie na załadowanie u innych...', 'watchTogether.waitingForName' => ({required Object name}) => 'Oczekiwanie na ${name}...', 'watchTogether.recentRooms' => 'Ostatnie pokoje', diff --git a/lib/i18n/strings_pt.g.dart b/lib/i18n/strings_pt.g.dart index dd868417..9c6c030d 100644 --- a/lib/i18n/strings_pt.g.dart +++ b/lib/i18n/strings_pt.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsPt extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Relay do Assistir Juntos'; @override String get watchTogetherRelayDescription => 'Defina um relay personalizado. Todos devem usar o mesmo servidor.'; @override String get watchTogetherRelayHint => 'https://meu-relay.exemplo.com.br'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Relatório de Erros'; @override String get crashReportingDescription => 'Enviar relatórios de erros para ajudar a melhorar o app'; @override String get debugLogging => 'Log de Depuração'; @@ -2303,6 +2304,7 @@ extension on TranslationsPt { 'settings.watchTogetherRelay' => 'Relay do Assistir Juntos', 'settings.watchTogetherRelayDescription' => 'Defina um relay personalizado. Todos devem usar o mesmo servidor.', 'settings.watchTogetherRelayHint' => 'https://meu-relay.exemplo.com.br', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Relatório de Erros', 'settings.crashReportingDescription' => 'Enviar relatórios de erros para ajudar a melhorar o app', 'settings.debugLogging' => 'Log de Depuração', @@ -2644,9 +2646,9 @@ extension on TranslationsPt { 'messages.markedAsUnwatchedOffline' => 'Marcado como não assistido (será sincronizado quando online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Removido automaticamente: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('pt'))(n, one: 'Removido automaticamente ${n} download assistido', other: 'Removidos automaticamente ${n} downloads assistidos', ), - 'messages.removedFromContinueWatching' => 'Removido de Continuar Assistindo', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Removido de Continuar Assistindo', 'messages.errorLoading' => ({required Object error}) => 'Erro: ${error}', 'messages.streamInterrupted' => 'A transmissão foi interrompida. Toque em reproduzir ou avance para tentar novamente.', 'messages.liveStreamInterrupted' => 'A transmissão ao vivo foi interrompida. Toque em reproduzir para tentar novamente.', @@ -3158,9 +3160,9 @@ extension on TranslationsPt { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} avançou', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} está carregando', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} está em uma versão mais antiga do aplicativo — sincronização indisponível', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Retomando sem ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Retomando sem ${name}', 'watchTogether.waitingForParticipants' => 'Aguardando outros carregarem...', 'watchTogether.waitingForName' => ({required Object name}) => 'Aguardando ${name}...', 'watchTogether.recentRooms' => 'Salas recentes', diff --git a/lib/i18n/strings_ru.g.dart b/lib/i18n/strings_ru.g.dart index f4b222be..b783415b 100644 --- a/lib/i18n/strings_ru.g.dart +++ b/lib/i18n/strings_ru.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsRu extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Relay совместного просмотра'; @override String get watchTogetherRelayDescription => 'Задайте свой relay. Все должны использовать один сервер.'; @override String get watchTogetherRelayHint => 'https://my-relay.example.com'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Отчёты об ошибках'; @override String get crashReportingDescription => 'Отправлять отчёты об ошибках для улучшения приложения'; @override String get debugLogging => 'Журнал отладки'; @@ -2309,6 +2310,7 @@ extension on TranslationsRu { 'settings.watchTogetherRelay' => 'Relay совместного просмотра', 'settings.watchTogetherRelayDescription' => 'Задайте свой relay. Все должны использовать один сервер.', 'settings.watchTogetherRelayHint' => 'https://my-relay.example.com', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Отчёты об ошибках', 'settings.crashReportingDescription' => 'Отправлять отчёты об ошибках для улучшения приложения', 'settings.debugLogging' => 'Журнал отладки', @@ -2650,9 +2652,9 @@ extension on TranslationsRu { 'messages.markedAsUnwatchedOffline' => 'Отмечено как непросмотренное (синхронизируется при подключении)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Автоудалено: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('ru'))(n, one: 'Автоматически удалена ${n} просмотренная загрузка', few: 'Автоматически удалены ${n} просмотренные загрузки', many: 'Автоматически удалено ${n} просмотренных загрузок', other: 'Автоматически удалено ${n} просмотренной загрузки', ), - 'messages.removedFromContinueWatching' => 'Удалено из «Продолжить просмотр»', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Удалено из «Продолжить просмотр»', 'messages.errorLoading' => ({required Object error}) => 'Ошибка: ${error}', 'messages.streamInterrupted' => 'Поток прервался. Нажмите «Воспроизвести» или перемотайте, чтобы повторить попытку.', 'messages.liveStreamInterrupted' => 'Прямая трансляция прервалась. Нажмите «Воспроизвести», чтобы повторить попытку.', @@ -3164,9 +3166,9 @@ extension on TranslationsRu { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} перемотал', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} буферизует', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} использует старую версию приложения — синхронизация недоступна', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Возобновление без ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Возобновление без ${name}', 'watchTogether.waitingForParticipants' => 'Ожидание загрузки у других...', 'watchTogether.waitingForName' => ({required Object name}) => 'Ожидание ${name}...', 'watchTogether.recentRooms' => 'Недавние комнаты', diff --git a/lib/i18n/strings_sv.g.dart b/lib/i18n/strings_sv.g.dart index d221d51d..ccab5580 100644 --- a/lib/i18n/strings_sv.g.dart +++ b/lib/i18n/strings_sv.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsSv extends TranslationsSettingsEn { @override String get watchTogetherRelay => 'Titta Tillsammans-relay'; @override String get watchTogetherRelayDescription => 'Ange en anpassad relay. Alla måste använda samma server.'; @override String get watchTogetherRelayHint => 'https://min-relay.exempel.se'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => 'Kraschrapportering'; @override String get crashReportingDescription => 'Skicka kraschrapporter för att förbättra appen'; @override String get debugLogging => 'Felsökningsloggning'; @@ -2303,6 +2304,7 @@ extension on TranslationsSv { 'settings.watchTogetherRelay' => 'Titta Tillsammans-relay', 'settings.watchTogetherRelayDescription' => 'Ange en anpassad relay. Alla måste använda samma server.', 'settings.watchTogetherRelayHint' => 'https://min-relay.exempel.se', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => 'Kraschrapportering', 'settings.crashReportingDescription' => 'Skicka kraschrapporter för att förbättra appen', 'settings.debugLogging' => 'Felsökningsloggning', @@ -2644,9 +2646,9 @@ extension on TranslationsSv { 'messages.markedAsUnwatchedOffline' => 'Markerad som osedd (synkroniseras när online)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => 'Automatiskt borttagen: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('sv'))(n, one: 'Tog automatiskt bort ${n} sedd nedladdning', other: 'Tog automatiskt bort ${n} sedda nedladdningar', ), - 'messages.removedFromContinueWatching' => 'Borttagen från Fortsätt titta', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => 'Borttagen från Fortsätt titta', 'messages.errorLoading' => ({required Object error}) => 'Fel: ${error}', 'messages.streamInterrupted' => 'Uppspelningen avbröts. Tryck på play eller spola för att försöka igen.', 'messages.liveStreamInterrupted' => 'Livestreamen avbröts. Tryck på play för att försöka igen.', @@ -3158,9 +3160,9 @@ extension on TranslationsSv { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} spolade', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} buffrar', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} använder en äldre appversion — synkronisering är inte tillgänglig', - 'watchTogether.resumingWithout' => ({required Object name}) => 'Återupptar utan ${name}', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => 'Återupptar utan ${name}', 'watchTogether.waitingForParticipants' => 'Väntar på att andra laddar...', 'watchTogether.waitingForName' => ({required Object name}) => 'Väntar på ${name}...', 'watchTogether.recentRooms' => 'Senaste rum', diff --git a/lib/i18n/strings_zh.g.dart b/lib/i18n/strings_zh.g.dart index a23f5b4e..b7a3b098 100644 --- a/lib/i18n/strings_zh.g.dart +++ b/lib/i18n/strings_zh.g.dart @@ -311,6 +311,7 @@ class _TranslationsSettingsZh extends TranslationsSettingsEn { @override String get watchTogetherRelay => '一起看中继服务器'; @override String get watchTogetherRelayDescription => '设置自定义中继。所有人必须使用同一服务器。'; @override String get watchTogetherRelayHint => 'https://my-relay.example.com'; + @override String get watchTogetherRelayInvalid => ''; @override String get crashReporting => '崩溃报告'; @override String get crashReportingDescription => '发送崩溃报告以帮助改进应用'; @override String get debugLogging => '调试日志'; @@ -2300,6 +2301,7 @@ extension on TranslationsZh { 'settings.watchTogetherRelay' => '一起看中继服务器', 'settings.watchTogetherRelayDescription' => '设置自定义中继。所有人必须使用同一服务器。', 'settings.watchTogetherRelayHint' => 'https://my-relay.example.com', + 'settings.watchTogetherRelayInvalid' => '', 'settings.crashReporting' => '崩溃报告', 'settings.crashReportingDescription' => '发送崩溃报告以帮助改进应用', 'settings.debugLogging' => '调试日志', @@ -2641,9 +2643,9 @@ extension on TranslationsZh { 'messages.markedAsUnwatchedOffline' => '已标记为未观看 (将在联网时同步)', 'messages.autoRemovedWatchedDownload' => ({required Object title}) => '已自动移除: ${title}', 'messages.autoRemovedWatchedDownloads' => ({required num n}) => (_root.$meta.cardinalResolver ?? PluralResolvers.cardinal('zh'))(n, other: '已自动移除 ${n} 个看过的下载', ), - 'messages.removedFromContinueWatching' => '已从继续观看中移除', _ => null, } ?? switch (path) { + 'messages.removedFromContinueWatching' => '已从继续观看中移除', 'messages.errorLoading' => ({required Object error}) => '错误: ${error}', 'messages.streamInterrupted' => '视频流已中断。按播放键或拖动进度条重试。', 'messages.liveStreamInterrupted' => '直播流已中断。按播放键重试。', @@ -3155,9 +3157,9 @@ extension on TranslationsZh { 'watchTogether.participantSeeked' => ({required Object name}) => '${name} 跳转了', 'watchTogether.participantBuffering' => ({required Object name}) => '${name} 正在缓冲', 'watchTogether.participantNeedsUpdate' => ({required Object name}) => '${name} 正在使用较旧版本的应用,无法同步', - 'watchTogether.resumingWithout' => ({required Object name}) => '不等待 ${name},继续播放', _ => null, } ?? switch (path) { + 'watchTogether.resumingWithout' => ({required Object name}) => '不等待 ${name},继续播放', 'watchTogether.waitingForParticipants' => '等待其他人加载...', 'watchTogether.waitingForName' => ({required Object name}) => '正在等待 ${name}...', 'watchTogether.recentRooms' => '最近的房间', diff --git a/lib/i18n/sv.i18n.json b/lib/i18n/sv.i18n.json index a5ca4d97..dd20a8db 100644 --- a/lib/i18n/sv.i18n.json +++ b/lib/i18n/sv.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "Titta Tillsammans-relay", "watchTogetherRelayDescription": "Ange en anpassad relay. Alla måste använda samma server.", "watchTogetherRelayHint": "https://min-relay.exempel.se", + "watchTogetherRelayInvalid": "", "crashReporting": "Kraschrapportering", "crashReportingDescription": "Skicka kraschrapporter för att förbättra appen", "debugLogging": "Felsökningsloggning", diff --git a/lib/i18n/zh.i18n.json b/lib/i18n/zh.i18n.json index d305f097..8e8da55e 100644 --- a/lib/i18n/zh.i18n.json +++ b/lib/i18n/zh.i18n.json @@ -180,6 +180,7 @@ "watchTogetherRelay": "一起看中继服务器", "watchTogetherRelayDescription": "设置自定义中继。所有人必须使用同一服务器。", "watchTogetherRelayHint": "https://my-relay.example.com", + "watchTogetherRelayInvalid": "", "crashReporting": "崩溃报告", "crashReportingDescription": "发送崩溃报告以帮助改进应用", "debugLogging": "调试日志", diff --git a/lib/providers/companion_remote_provider.dart b/lib/providers/companion_remote_provider.dart index 5a40648d..652fe112 100644 --- a/lib/providers/companion_remote_provider.dart +++ b/lib/providers/companion_remote_provider.dart @@ -25,6 +25,8 @@ export '../services/companion_remote/lan_discovery_service.dart' show Discovered typedef CommandReceivedCallback = void Function(RemoteCommand command); typedef PlexHomeResolver = Future Function(String connectionId); +typedef CompanionRemotePeerServiceFactory = CompanionRemotePeerService Function(); +typedef LanDiscoveryServiceFactory = LanDiscoveryService Function(); String _localizedRemoteError(Object error, String Function(String details) fallback) { if (error is RemotePeerError) return error.message; @@ -32,8 +34,24 @@ String _localizedRemoteError(Object error, String Function(String details) fallb } class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin { + CompanionRemoteProvider() : this._(CompanionRemotePeerService.new, LanDiscoveryService.new); + + @visibleForTesting + CompanionRemoteProvider.forTesting({ + required CompanionRemotePeerServiceFactory peerServiceFactory, + LanDiscoveryServiceFactory discoveryServiceFactory = LanDiscoveryService.new, + }) : this._(peerServiceFactory, discoveryServiceFactory); + + CompanionRemoteProvider._(this._peerServiceFactory, this._discoveryServiceFactory) { + _initializeDeviceInfo(); + } + + final CompanionRemotePeerServiceFactory _peerServiceFactory; + final LanDiscoveryServiceFactory _discoveryServiceFactory; RemoteSession? _session; CompanionRemotePeerService? _peerService; + CompanionRemotePeerService? _pendingRemotePeer; + final Expando> _peerDisposals = Expando>('companion remote peer disposal'); LanDiscoveryService? _discoveryService; String _deviceName = t.companionRemote.unknownDevice; String _platform = 'unknown'; @@ -44,7 +62,11 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin Timer? _reconnectTimer; int _reconnectAttempts = 0; Future? _activeReconnect; - bool _intentionalDisconnect = false; + int? _activeReconnectGeneration; + int _remoteGeneration = 0; + final Expando _intentionalDisconnectGeneration = Expando( + 'companion remote intentional disconnect generation', + ); // Reconnection context (only hostAddresses and hostClientId are connection-specific) List? _lastHostAddresses; @@ -90,10 +112,6 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin bool get isPlayerActive => _isPlayerActive; bool get isHostServerRunning => _peerService?.isServerRunning ?? false; - CompanionRemoteProvider() { - _initializeDeviceInfo(); - } - Future _initializeDeviceInfo() async { final identity = await DeviceIdentityService.resolve(); _deviceName = identity.deviceName ?? t.companionRemote.unknownDevice; @@ -420,16 +438,110 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin _cryptoProfileId = null; } + void _markIntentionalDisconnect(CompanionRemotePeerService? peer, int generation) { + if (peer != null) { + _intentionalDisconnectGeneration[peer] = generation; + } + } + + void _clearIntentionalDisconnect(CompanionRemotePeerService? peer, int generation) { + if (peer != null && _intentionalDisconnectGeneration[peer] == generation) { + _intentionalDisconnectGeneration[peer] = null; + } + } + + bool _isIntentionalDisconnect(CompanionRemotePeerService peer, int generation) { + return _intentionalDisconnectGeneration[peer] == generation; + } + + bool _ownsPeer(CompanionRemotePeerService peer, int generation) { + return !isDisposed && + generation == _remoteGeneration && + (identical(_peerService, peer) || identical(_pendingRemotePeer, peer)); + } + + ({CompanionRemotePeerService? current, CompanionRemotePeerService? pending}) _invalidateRemoteLifecycle() { + _remoteGeneration++; + _reconnectTimer?.cancel(); + _reconnectTimer = null; + + final pending = _pendingRemotePeer; + final current = _session?.isRemote == true ? _peerService : null; + _pendingRemotePeer = null; + if (identical(_peerService, current)) { + _peerService = null; + } + if (current != null || pending != null) { + _cleanupSubscriptions(); + } + return (current: current, pending: pending); + } + + Future _disposePeerOnce(CompanionRemotePeerService peer) { + final existing = _peerDisposals[peer]; + if (existing != null) return existing; + + final disposal = () async { + try { + await peer.dispose(); + } catch (error, stackTrace) { + appLogger.d('CompanionRemote: Peer cleanup ignored', error: error, stackTrace: stackTrace); + } + }(); + _peerDisposals[peer] = disposal; + return disposal; + } + + Future _disposeDetachedPeers( + ({CompanionRemotePeerService? current, CompanionRemotePeerService? pending}) peers, + ) async { + final current = peers.current; + final pending = peers.pending; + if (current != null) { + await _disposePeerOnce(current); + } + if (pending != null && !identical(pending, current)) { + await _disposePeerOnce(pending); + } + } + + _RemoteConnectRequest _beginRemoteConnectRequest() { + final wasHost = isHost || isHostServerRunning; + final peers = _invalidateRemoteLifecycle(); + _reconnectAttempts = 0; + _session = null; + _isPlayerActive = false; + safeNotifyListeners(); + return _RemoteConnectRequest( + generation: _remoteGeneration, + wasHost: wasHost, + current: peers.current, + pending: peers.pending, + ); + } + + Future _prepareRemoteConnect(_RemoteConnectRequest request) async { + await _disposeDetachedPeers((current: request.current, pending: request.pending)); + if (request.wasHost) { + await _serializeLifecycle(_stopHostServerLocked); + } else { + stopDiscovery(); + } + return !isDisposed && request.generation == _remoteGeneration; + } + /// Fully tear down network/session state and forget derived crypto material. /// Used by logout so an app-level provider surviving route replacement does /// not keep broadcasting with the previous Plex Home identity. Future resetForLogout() { + final detachedPeers = _invalidateRemoteLifecycle(); + _reconnectAttempts = 0; + _lastHostAddresses = null; + _lastHostClientId = null; + _lastAuthContextId = null; + return _serializeLifecycle(() async { - _reconnectTimer?.cancel(); - _reconnectAttempts = 0; - _lastHostAddresses = null; - _lastHostClientId = null; - _lastAuthContextId = null; + await _disposeDetachedPeers(detachedPeers); await _stopHostServerLocked(); stopDiscovery(); _clearCryptoContext(); @@ -450,6 +562,12 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin @visibleForTesting List get debugCryptoConnectionIds => _authContexts.map((context) => context.connectionId).toList(); + @visibleForTesting + bool get debugIsDiscoveryBroadcasting => _discoveryService?.isBroadcasting ?? false; + + @visibleForTesting + bool get debugIsDiscoveryListening => _discoveryService?.isListening ?? false; + Future startHostServer() => _serializeLifecycle(_startHostServerLocked); Future _startHostServerLocked() async { @@ -461,8 +579,8 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin appLogger.d('CompanionRemote: Starting host server'); - _peerService ??= CompanionRemotePeerService(); - _setupPeerServiceListeners(); + final peer = _peerService ??= _peerServiceFactory(); + _setupPeerServiceListeners(peer, _remoteGeneration); try { final contexts = List.unmodifiable(_authContexts); @@ -476,7 +594,7 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin safeNotifyListeners(); // Start LAN discovery broadcasting - _discoveryService ??= LanDiscoveryService(); + _discoveryService ??= _discoveryServiceFactory(); final localIps = result.addresses.map((a) => a.split(':').first).toList(); await _discoveryService!.startBroadcastingForContexts( contexts: contexts, @@ -503,19 +621,32 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin Future stopHostServer() => _serializeLifecycle(_stopHostServerLocked); Future _stopHostServerLocked() async { - _intentionalDisconnect = true; - await _discoveryService?.stopBroadcasting(); + final stopGeneration = _remoteGeneration; + final peer = _peerService; + _markIntentionalDisconnect(peer, stopGeneration); - if (_peerService != null) { - await _peerService!.disconnect(); - _peerService = null; + try { + await _discoveryService?.stopBroadcasting(); + _discoveryService?.stopListening(); + + if (identical(_peerService, peer)) { + _peerService = null; + _cleanupSubscriptions(); + } + if (peer != null) { + await _disposePeerOnce(peer); + } + + // A newer remote request may start while the stopped peer's asynchronous + // disposal is settling. Never let the older stop erase that replacement. + if (_remoteGeneration == stopGeneration) { + _session = null; + _isPlayerActive = false; + } + safeNotifyListeners(); + } finally { + _clearIntentionalDisconnect(peer, stopGeneration); } - _cleanupSubscriptions(); - - _session = null; - _isPlayerActive = false; - _intentionalDisconnect = false; - safeNotifyListeners(); } Stream>? discoverHosts() { @@ -524,7 +655,7 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin return null; } - _discoveryService ??= LanDiscoveryService(); + _discoveryService ??= _discoveryServiceFactory(); return _discoveryService!.startListeningForContexts(_authContexts); } @@ -543,26 +674,28 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin throw RemotePeerError(type: RemotePeerErrorType.authFailed, message: t.companionRemote.pairing.authFailed); } - await leaveSession(); + final request = _beginRemoteConnectRequest(); + if (!await _prepareRemoteConnect(request)) return; + final generation = request.generation; _lastHostAddresses = host.addresses; _lastHostClientId = host.clientId; _lastAuthContextId = authContext.id; appLogger.d('CompanionRemote: Connecting to ${host.name} at ${host.addresses}'); - _peerService = CompanionRemotePeerService(); - _setupPeerServiceListeners(); - + final candidate = _peerServiceFactory(); + _pendingRemotePeer = candidate; _session = RemoteSession( role: RemoteSessionRole.remote, status: RemoteSessionStatus.connecting, createdAt: DateTime.now(), ); + _setupPeerServiceListeners(candidate, generation); safeNotifyListeners(); try { - final winner = await _peerService!.joinSessionRacingWithContexts( + final winner = await candidate.joinSessionRacingWithContexts( _deviceName, _platform, host.addresses, @@ -570,18 +703,35 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin authContextId: authContext.id, expectedHostClientId: host.clientId, ); - _lastHostAddresses = [winner]; - _lastAuthContextId = _peerService!.selectedAuthContextId ?? authContext.id; - _lastHostClientId = _peerService!.selectedHostClientId ?? host.clientId; + if (!_ownsPeer(candidate, generation)) { + await _disposePeerOnce(candidate); + return; + } + _pendingRemotePeer = null; + _peerService = candidate; + _lastHostAddresses = [winner]; + _lastAuthContextId = candidate.selectedAuthContextId ?? authContext.id; + _lastHostClientId = candidate.selectedHostClientId ?? host.clientId; _session = _session?.copyWith(status: RemoteSessionStatus.connected); safeNotifyListeners(); appLogger.d('CompanionRemote: Connected to ${host.name} via $winner'); - } catch (e) { - appLogger.e('CompanionRemote: Failed to connect to host', error: e); + } catch (error, stackTrace) { + if (!_ownsPeer(candidate, generation)) { + await _disposePeerOnce(candidate); + return; + } + + _pendingRemotePeer = null; + _cleanupSubscriptions(); + await _disposePeerOnce(candidate); + appLogger.e('CompanionRemote: Failed to connect to host', error: error, stackTrace: stackTrace); _session = _session?.copyWith( status: RemoteSessionStatus.error, - errorMessage: _localizedRemoteError(e, (details) => t.companionRemote.pairing.failedToConnect(error: details)), + errorMessage: _localizedRemoteError( + error, + (details) => t.companionRemote.pairing.failedToConnect(error: details), + ), ); safeNotifyListeners(); rethrow; @@ -594,49 +744,69 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin throw RemotePeerError(type: RemotePeerErrorType.authFailed, message: t.companionRemote.pairing.cryptoInitFailed); } - await leaveSession(); + final request = _beginRemoteConnectRequest(); + if (!await _prepareRemoteConnect(request)) return; + final generation = request.generation; _lastHostAddresses = [hostAddress]; _lastHostClientId = ''; _lastAuthContextId = null; appLogger.d('CompanionRemote: Connecting to manual host $hostAddress'); - _peerService = CompanionRemotePeerService(); - _setupPeerServiceListeners(); - + final candidate = _peerServiceFactory(); + _pendingRemotePeer = candidate; _session = RemoteSession( role: RemoteSessionRole.remote, status: RemoteSessionStatus.connecting, createdAt: DateTime.now(), ); + _setupPeerServiceListeners(candidate, generation); safeNotifyListeners(); try { - await _peerService!.joinSessionWithContexts(_deviceName, _platform, hostAddress, _authContexts); - _lastAuthContextId = _peerService!.selectedAuthContextId; - _lastHostClientId = _peerService!.selectedHostClientId ?? ''; + await candidate.joinSessionWithContexts(_deviceName, _platform, hostAddress, _authContexts); + if (!_ownsPeer(candidate, generation)) { + await _disposePeerOnce(candidate); + return; + } + _pendingRemotePeer = null; + _peerService = candidate; + _lastAuthContextId = candidate.selectedAuthContextId; + _lastHostClientId = candidate.selectedHostClientId ?? ''; _session = _session?.copyWith(status: RemoteSessionStatus.connected); safeNotifyListeners(); - } catch (e) { - appLogger.e('CompanionRemote: Failed to connect to manual host', error: e); + } catch (error, stackTrace) { + if (!_ownsPeer(candidate, generation)) { + await _disposePeerOnce(candidate); + return; + } + + _pendingRemotePeer = null; + _cleanupSubscriptions(); + await _disposePeerOnce(candidate); + appLogger.e('CompanionRemote: Failed to connect to manual host', error: error, stackTrace: stackTrace); _session = _session?.copyWith( status: RemoteSessionStatus.error, - errorMessage: _localizedRemoteError(e, (details) => t.companionRemote.pairing.failedToConnect(error: details)), + errorMessage: _localizedRemoteError( + error, + (details) => t.companionRemote.pairing.failedToConnect(error: details), + ), ); safeNotifyListeners(); rethrow; } } - void _setupPeerServiceListeners() { - // A rebuild/reconnect can re-enter here with live subscriptions from the - // previous peer service still attached; drop them first so events don't - // fan out to a stale service. + void _setupPeerServiceListeners(CompanionRemotePeerService peer, int generation) { + // Only one peer owns the provider's listener set at a time. The identity + // and generation checks also reject events already queued when teardown + // synchronously cancels these subscriptions. _cleanupSubscriptions(); - _commandSubscription = _peerService!.onCommandReceived.listen( + _commandSubscription = peer.onCommandReceived.listen( (command) { + if (!_ownsPeer(peer, generation)) return; appLogger.d('CompanionRemote: Command received: ${command.type}'); if (command.type == RemoteCommandType.deviceInfo) { @@ -649,20 +819,24 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin onCommandReceived?.call(command); } }, - onError: (error) { + onError: (Object error) { + if (!_ownsPeer(peer, generation)) return; appLogger.e('CompanionRemote: Stream error', error: error); }, ); - _deviceConnectedSubscription = _peerService!.onDeviceConnected.listen((device) { + _deviceConnectedSubscription = peer.onDeviceConnected.listen((device) { + if (!_ownsPeer(peer, generation)) return; appLogger.d('CompanionRemote: Device connected: ${device.name}'); _session = _session?.copyWith(status: RemoteSessionStatus.connected, connectedDevice: device); safeNotifyListeners(); }); - _deviceDisconnectedSubscription = _peerService!.onDeviceDisconnected.listen((_) { - appLogger.d('CompanionRemote: Device disconnected (intentional: $_intentionalDisconnect)'); - if (_intentionalDisconnect) { + _deviceDisconnectedSubscription = peer.onDeviceDisconnected.listen((_) { + if (!_ownsPeer(peer, generation)) return; + final intentional = _isIntentionalDisconnect(peer, generation); + appLogger.d('CompanionRemote: Device disconnected (intentional: $intentional)'); + if (intentional) { _session = _session?.copyWith(status: RemoteSessionStatus.disconnected, connectedDevice: null); safeNotifyListeners(); } else if (isHost) { @@ -676,17 +850,19 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin } else { _session = _session?.copyWith(status: RemoteSessionStatus.reconnecting); safeNotifyListeners(); - _scheduleReconnect(); + _scheduleReconnect(generation); } }); - _errorSubscription = _peerService!.onError.listen((error) { + _errorSubscription = peer.onError.listen((error) { + if (!_ownsPeer(peer, generation)) return; appLogger.e('CompanionRemote: Error: ${error.message}'); _session = _session?.copyWith(status: RemoteSessionStatus.error, errorMessage: error.message); safeNotifyListeners(); }); - _statusSubscription = _peerService!.onConnectionStateChanged.listen((status) { + _statusSubscription = peer.onConnectionStateChanged.listen((status) { + if (!_ownsPeer(peer, generation)) return; appLogger.d('CompanionRemote: Status changed: $status'); _session = _session?.copyWith(status: status); safeNotifyListeners(); @@ -740,7 +916,8 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin _peerService!.sendCommand(RemoteCommand(type: type, data: data)); } - void _scheduleReconnect() { + void _scheduleReconnect(int generation) { + if (generation != _remoteGeneration || isDisposed) return; if (_reconnectAttempts >= _maxReconnectAttempts) { appLogger.w('CompanionRemote: Max reconnect attempts reached'); _session = _session?.copyWith( @@ -757,22 +934,35 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin appLogger.d('CompanionRemote: Reconnect attempt $_reconnectAttempts in ${delay.inSeconds}s'); _reconnectTimer?.cancel(); - _reconnectTimer = Timer(delay, _attemptReconnect); + _reconnectTimer = Timer(delay, () { + if (generation != _remoteGeneration || isDisposed) return; + unawaited(_attemptReconnect()); + }); } Future _attemptReconnect() { + final generation = _remoteGeneration; final active = _activeReconnect; - if (active != null) return active; + if (active != null && _activeReconnectGeneration == generation) { + return active; + } + late final Future attempt; - attempt = _runReconnectAttempt().whenComplete(() { - if (identical(_activeReconnect, attempt)) _activeReconnect = null; + attempt = _runReconnectAttempt(generation).whenComplete(() { + if (identical(_activeReconnect, attempt)) { + _activeReconnect = null; + _activeReconnectGeneration = null; + } }); _activeReconnect = attempt; + _activeReconnectGeneration = generation; return attempt; } - Future _runReconnectAttempt() async { - if (_lastHostAddresses == null || !isCryptoReady) { + Future _runReconnectAttempt(int generation) async { + if (generation != _remoteGeneration || isDisposed) return; + final hostAddresses = _lastHostAddresses; + if (hostAddresses == null || !isCryptoReady) { appLogger.w('CompanionRemote: No stored context for reconnect'); _session = _session?.copyWith( status: RemoteSessionStatus.error, @@ -782,86 +972,133 @@ class CompanionRemoteProvider with ChangeNotifier, DisposableChangeNotifierMixin return; } - try { - appLogger.d('CompanionRemote: Attempting reconnect...'); - _cleanupSubscriptions(); - try { - await _peerService?.disconnect(); - } finally { - _peerService = CompanionRemotePeerService(); - _setupPeerServiceListeners(); - } + appLogger.d('CompanionRemote: Attempting reconnect...'); + final oldPeer = _peerService; + _peerService = null; + _cleanupSubscriptions(); + if (oldPeer != null) { + await _disposePeerOnce(oldPeer); + } + if (generation != _remoteGeneration || isDisposed) return; - final authContextId = _authContextForId(_lastAuthContextId)?.id; - await _peerService!.joinSessionWithContexts( + final candidate = _peerServiceFactory(); + _pendingRemotePeer = candidate; + _setupPeerServiceListeners(candidate, generation); + final authContextId = _authContextForId(_lastAuthContextId)?.id; + final expectedHostClientId = _lastHostClientId ?? ''; + + try { + await candidate.joinSessionWithContexts( _deviceName, _platform, - _lastHostAddresses!.first, + hostAddresses.first, _authContexts, authContextId: authContextId, - expectedHostClientId: _lastHostClientId ?? '', + expectedHostClientId: expectedHostClientId, ); - _lastAuthContextId = _peerService!.selectedAuthContextId ?? authContextId; - _lastHostClientId = _peerService!.selectedHostClientId ?? _lastHostClientId; + if (!_ownsPeer(candidate, generation)) { + await _disposePeerOnce(candidate); + return; + } + _pendingRemotePeer = null; + _peerService = candidate; + _lastAuthContextId = candidate.selectedAuthContextId ?? authContextId; + _lastHostClientId = candidate.selectedHostClientId ?? _lastHostClientId; _session = _session?.copyWith(status: RemoteSessionStatus.connected, errorMessage: null); _reconnectAttempts = 0; safeNotifyListeners(); appLogger.d('CompanionRemote: Reconnected successfully'); - } catch (e) { - appLogger.e('CompanionRemote: Reconnect failed', error: e); - if (_session?.status == RemoteSessionStatus.reconnecting) { - _scheduleReconnect(); + } catch (error, stackTrace) { + if (!_ownsPeer(candidate, generation)) { + await _disposePeerOnce(candidate); + return; + } + + _pendingRemotePeer = null; + _cleanupSubscriptions(); + await _disposePeerOnce(candidate); + appLogger.e('CompanionRemote: Reconnect failed', error: error, stackTrace: stackTrace); + if (generation == _remoteGeneration && _session?.status == RemoteSessionStatus.reconnecting) { + _scheduleReconnect(generation); } } } - void retryReconnectNow() { + Future retryReconnectNow() { _reconnectTimer?.cancel(); + _reconnectTimer = null; _reconnectAttempts = 0; - _attemptReconnect(); + return _attemptReconnect(); } - void cancelReconnect() { - _reconnectTimer?.cancel(); + Future cancelReconnect() async { + final wasHost = isHost || isHostServerRunning; + final detachedPeers = _invalidateRemoteLifecycle(); _reconnectAttempts = 0; - _session = _session?.copyWith(status: RemoteSessionStatus.disconnected, connectedDevice: null); + _session = null; + _isPlayerActive = false; safeNotifyListeners(); + await _disposeDetachedPeers(detachedPeers); + if (wasHost) { + await _serializeLifecycle(_stopHostServerLocked); + } else { + stopDiscovery(); + } } Future leaveSession() async { - _intentionalDisconnect = true; - _reconnectTimer?.cancel(); + final leavingGeneration = _remoteGeneration; + _markIntentionalDisconnect(_peerService, leavingGeneration); + _markIntentionalDisconnect(_pendingRemotePeer, leavingGeneration); + final wasHost = isHost || isHostServerRunning; + final detachedPeers = _invalidateRemoteLifecycle(); _reconnectAttempts = 0; - - // Don't stop the host server when leaving — only stop discovery listening - if (_peerService != null && !isHost) { - appLogger.d('CompanionRemote: Leaving session'); - await _peerService!.disconnect(); - _peerService = null; - } - - _cleanupSubscriptions(); - - if (!isHost) { - _session = null; - } + _session = null; _isPlayerActive = false; - _intentionalDisconnect = false; safeNotifyListeners(); + + await _disposeDetachedPeers(detachedPeers); + if (wasHost) { + await _serializeLifecycle(_stopHostServerLocked); + } else { + stopDiscovery(); + } } @override void dispose() { + _remoteGeneration++; _reconnectTimer?.cancel(); + _reconnectTimer = null; + final current = _peerService; + final pending = _pendingRemotePeer; + _peerService = null; + _pendingRemotePeer = null; + _cleanupSubscriptions(); + _boundActiveProfile?.removeListener(_scheduleAuthContextRefresh); for (final sub in _profileServiceSubs) { sub.cancel(); } _profileServiceSubs.clear(); _discoveryService?.dispose(); - _peerService?.dispose(); + unawaited(_disposeDetachedPeers((current: current, pending: pending))); RemoteAuthService.instance.clearCache(); super.dispose(); } } + +class _RemoteConnectRequest { + const _RemoteConnectRequest({ + required this.generation, + required this.wasHost, + required this.current, + required this.pending, + }); + + final int generation; + final bool wasHost; + final CompanionRemotePeerService? current; + final CompanionRemotePeerService? pending; +} diff --git a/lib/screens/companion_remote/mobile_remote_screen.dart b/lib/screens/companion_remote/mobile_remote_screen.dart index 67e879d2..5e2098c1 100644 --- a/lib/screens/companion_remote/mobile_remote_screen.dart +++ b/lib/screens/companion_remote/mobile_remote_screen.dart @@ -82,10 +82,17 @@ class _MobileRemoteScreenState extends State { Row( mainAxisAlignment: .center, children: [ - OutlinedButton(onPressed: () => provider.cancelReconnect(), child: Text(t.common.cancel)), + OutlinedButton( + onPressed: () async { + await provider.cancelReconnect(); + }, + child: Text(t.common.cancel), + ), const SizedBox(width: 16), FilledButton( - onPressed: () => provider.retryReconnectNow(), + onPressed: () async { + await provider.retryReconnectNow(); + }, child: Text(t.companionRemote.remote.retryNow), ), ], diff --git a/lib/screens/settings/logs_screen.dart b/lib/screens/settings/logs_screen.dart index 6c50b5f4..9123f196 100644 --- a/lib/screens/settings/logs_screen.dart +++ b/lib/screens/settings/logs_screen.dart @@ -58,7 +58,9 @@ String constrainLogUploadPayload({required String header, required String logs, } class LogsScreen extends StatefulWidget { - const LogsScreen({super.key}); + const LogsScreen({super.key, this.httpClient}); + + final MediaServerHttpClient? httpClient; @override State createState() => _LogsScreenState(); @@ -69,6 +71,8 @@ class _LogsScreenState extends State with MountedSetStateMixin { String _deviceInfo = ''; final ScrollController _scrollController = ScrollController(); + MediaServerHttpClient get _httpClient => widget.httpClient ?? httpClient; + @override void initState() { super.initState(); @@ -182,7 +186,7 @@ class _LogsScreenState extends State with MountedSetStateMixin { showLoadingDialog(context); try { - final response = await httpClient.post( + final response = await _httpClient.post( 'https://ice.plezy.app/logs', body: logText, headers: {'Content-Type': 'text/plain'}, @@ -199,20 +203,29 @@ class _LogsScreenState extends State with MountedSetStateMixin { context: context, builder: (ctx) => AlertDialog( title: Text(t.messages.logsUploaded), - content: Row( + content: Column( + mainAxisSize: MainAxisSize.min, + crossAxisAlignment: CrossAxisAlignment.start, children: [ - Text('${t.messages.logId}: '), - SelectableText( - id, - style: const TextStyle(fontWeight: .bold, fontFamily: 'monospace', fontSize: 18), - ), - const SizedBox(width: 8), - IconButton( - icon: const AppIcon(Symbols.content_copy_rounded, size: 20), - onPressed: () { - Clipboard.setData(ClipboardData(text: id)); - showSuccessSnackBar(context, t.messages.logsCopied); - }, + Text('${t.messages.logId}:'), + const SizedBox(height: 8), + Row( + children: [ + Expanded( + child: SelectableText( + id, + style: const TextStyle(fontWeight: .bold, fontFamily: 'monospace', fontSize: 18), + ), + ), + const SizedBox(width: 8), + IconButton( + icon: const AppIcon(Symbols.content_copy_rounded, size: 20), + onPressed: () { + Clipboard.setData(ClipboardData(text: id)); + showSuccessSnackBar(context, t.messages.logsCopied); + }, + ), + ], ), ], ), diff --git a/lib/screens/settings/settings_screen.dart b/lib/screens/settings/settings_screen.dart index 44742b7b..2f5c634f 100644 --- a/lib/screens/settings/settings_screen.dart +++ b/lib/screens/settings/settings_screen.dart @@ -44,6 +44,7 @@ import '../../widgets/settings_builder.dart'; import '../../widgets/settings_section.dart'; import '../../profiles/active_profile_provider.dart'; import '../../profiles/profile.dart'; +import '../../watch_together/services/watch_together_relay_endpoint.dart'; import 'about_screen.dart'; import 'add_connection_screen.dart'; import 'appearance_settings_screen.dart'; @@ -801,6 +802,7 @@ class _RelayUrlDialog extends StatefulWidget { class _RelayUrlDialogState extends State<_RelayUrlDialog> { late final TextEditingController _controller; final _saveFocusNode = FocusNode(debugLabel: 'WatchTogetherRelaySave'); + bool _relayUrlInvalid = false; @override void initState() { @@ -824,8 +826,19 @@ class _RelayUrlDialogState extends State<_RelayUrlDialog> { } Future _save() async { - final trimmed = _controller.text.trim(); - await widget.settingsService.write(settings.SettingsService.customRelayUrl, trimmed.isEmpty ? null : trimmed); + final value = _controller.text; + if (value.trim().isEmpty) { + await widget.settingsService.write(settings.SettingsService.customRelayUrl, null); + if (mounted) Navigator.pop(context); + return; + } + + final endpoint = WatchTogetherRelayEndpoint.tryParseCustom(value); + if (endpoint == null) { + setState(() => _relayUrlInvalid = true); + return; + } + await widget.settingsService.write(settings.SettingsService.customRelayUrl, endpoint.canonicalBaseUrl); if (mounted) Navigator.pop(context); } @@ -835,9 +848,18 @@ class _RelayUrlDialogState extends State<_RelayUrlDialog> { title: Text(t.settings.watchTogetherRelay), content: FocusableTextField( controller: _controller, - decoration: InputDecoration(labelText: 'URL', hintText: t.settings.watchTogetherRelayHint), + decoration: InputDecoration( + labelText: 'URL', + hintText: t.settings.watchTogetherRelayHint, + errorText: _relayUrlInvalid ? t.settings.watchTogetherRelayInvalid : null, + ), autofocus: true, textInputAction: TextInputAction.done, + onChanged: (_) { + if (_relayUrlInvalid) { + setState(() => _relayUrlInvalid = false); + } + }, onEditingComplete: () => _saveFocusNode.requestFocus(), ), actions: [ diff --git a/lib/services/companion_remote/companion_remote_peer_service.dart b/lib/services/companion_remote/companion_remote_peer_service.dart index 986fa3a3..b989a2c5 100644 --- a/lib/services/companion_remote/companion_remote_peer_service.dart +++ b/lib/services/companion_remote/companion_remote_peer_service.dart @@ -23,15 +23,98 @@ export '../base_peer_service.dart' show PeerError, PeerErrorType; typedef RemotePeerErrorType = PeerErrorType; typedef RemotePeerError = PeerError; +typedef _SessionKeyDeriver = + Future> Function(List homeSecret, List hostNonce, List clientNonce); +typedef _RaceProbeConnection = ({Future Function() close, Future ready, Stream stream}); + +typedef _RaceProbeFactory = _RaceProbeConnection Function(Uri uri); + class CompanionRemotePeerService with KeepaliveMixin { + static const int _productionMaxTotalHostConnections = 100; + static const int _productionMaxHostConnectionsPerSource = 5; + static const int _productionMaxPreAuthMessageBytes = 65536; + static const Duration _productionAuthTimeout = Duration(seconds: 10); + static const int _productionMaxFailedAuthAttempts = 5; + static const Duration _productionAuthLockoutDuration = Duration(seconds: 30); + + CompanionRemotePeerService() + : this.forTesting( + maxTotalHostConnections: _productionMaxTotalHostConnections, + maxHostConnectionsPerSource: _productionMaxHostConnectionsPerSource, + maxPreAuthMessageBytes: _productionMaxPreAuthMessageBytes, + authTimeout: _productionAuthTimeout, + maxFailedAuthAttempts: _productionMaxFailedAuthAttempts, + authLockoutDuration: _productionAuthLockoutDuration, + ); + + CompanionRemotePeerService.forTesting({ + int maxTotalHostConnections = _productionMaxTotalHostConnections, + int maxHostConnectionsPerSource = _productionMaxHostConnectionsPerSource, + int maxPreAuthMessageBytes = _productionMaxPreAuthMessageBytes, + Duration authTimeout = _productionAuthTimeout, + int maxFailedAuthAttempts = _productionMaxFailedAuthAttempts, + Duration authLockoutDuration = _productionAuthLockoutDuration, + Future> Function(List homeSecret, List hostNonce, List clientNonce)? deriveSessionEncKey, + ({Future Function() close, Future ready, Stream stream}) Function(Uri uri)? raceProbeFactory, + this._afterHostUpgrade, + }) : assert(maxTotalHostConnections > 0), + assert(maxHostConnectionsPerSource > 0), + assert(maxPreAuthMessageBytes > 0), + assert(authTimeout > Duration.zero), + assert(maxFailedAuthAttempts > 0), + _maxTotalHostConnections = maxTotalHostConnections, + _maxHostConnectionsPerSource = maxHostConnectionsPerSource, + _maxPreAuthMessageBytes = maxPreAuthMessageBytes, + _authTimeout = authTimeout, + _maxFailedAuthAttempts = maxFailedAuthAttempts, + _authLockoutDuration = authLockoutDuration, + _deriveSessionEncKey = + deriveSessionEncKey ?? + ((homeSecret, hostNonce, clientNonce) { + return RemoteAuthService.instance.deriveSessionEncKey(homeSecret, hostNonce, clientNonce); + }), + _raceProbeFactory = raceProbeFactory ?? _openRaceProbe; + + static _RaceProbeConnection _openRaceProbe(Uri uri) { + final channel = IOWebSocketChannel.connect(uri, connectTimeout: const Duration(seconds: 5)); + return ( + close: () async { + await channel.sink.close(); + }, + ready: channel.ready, + stream: channel.stream, + ); + } + + final int _maxTotalHostConnections; + final int _maxHostConnectionsPerSource; + final int _maxPreAuthMessageBytes; + final Duration _authTimeout; + final int _maxFailedAuthAttempts; + final Duration _authLockoutDuration; + final _SessionKeyDeriver _deriveSessionEncKey; + final _RaceProbeFactory _raceProbeFactory; + final void Function()? _afterHostUpgrade; + // Server-side (host) fields HttpServer? _server; WebSocket? _clientSocket; + _HostAdmission? _currentHostAdmission; + final Set<_HostAdmission> _hostAdmissions = {}; + final Map _hostAdmissionsBySource = {}; + int _hostAdmissionCount = 0; + int _authenticationCommitGeneration = 0; + Future _hostAuthenticationCommitTail = Future.value(); + bool _acceptingHostConnections = false; + bool _isDisconnecting = false; + Future? _disconnectInProgress; + Future? _disposeInProgress; + bool _disposed = false; // Client-side (remote) fields IOWebSocketChannel? _channel; - StreamSubscription? _clientSocketSubscription; StreamSubscription? _channelSubscription; + int _remoteConnectionGeneration = 0; String? _myPeerId; String? _hostAddress; // Format: "ip:port" @@ -60,8 +143,6 @@ class CompanionRemotePeerService with KeepaliveMixin { // Auth rate limiting (per source IP) final Map _failedAuthAttempts = {}; final Map _authLockouts = {}; - static const int _maxFailedAuthAttempts = 5; - static const Duration _authLockoutDuration = Duration(seconds: 30); Stream get onCommandReceived => _commandReceivedController.stream; Stream get onDeviceConnected => _deviceConnectedController.stream; @@ -119,6 +200,10 @@ class CompanionRemotePeerService with KeepaliveMixin { String platform, List authContexts, ) async { + if (_disposed) { + throw StateError('CompanionRemotePeerService is disposed'); + } + if (authContexts.isEmpty) { throw RemotePeerError( type: RemotePeerErrorType.authFailed, @@ -135,41 +220,38 @@ class CompanionRemotePeerService with KeepaliveMixin { try { const int preferredPort = 48632; + late final HttpServer server; try { - _server = await HttpServer.bind(InternetAddress.anyIPv4, preferredPort); + server = await HttpServer.bind(InternetAddress.anyIPv4, preferredPort); appLogger.d('CompanionRemote: Server bound to port $preferredPort'); } catch (e) { appLogger.w('CompanionRemote: Port $preferredPort occupied, using random port'); - _server = await HttpServer.bind(InternetAddress.anyIPv4, 0); + server = await HttpServer.bind(InternetAddress.anyIPv4, 0); } + _server = server; + _acceptingHostConnections = true; + final localIps = await _getAllLocalIpAddresses(); - final port = _server!.port; + final port = server.port; final addresses = localIps.map((ip) => '$ip:$port').toList(); _hostAddress = addresses.first; appLogger.d('CompanionRemote: Host server started, addresses: $addresses'); - _server!.listen((HttpRequest request) async { - if (request.uri.path == '/ws') { - try { - final socket = await WebSocketTransformer.upgrade(request); - final sourceIp = request.connectionInfo?.remoteAddress.address ?? 'unknown'; - _handleNewWebSocketConnection(socket, deviceName, platform, authContexts, sourceIp); - } catch (e) { - appLogger.e('CompanionRemote: Failed to upgrade WebSocket', error: e); - } - } else { - request.response.statusCode = HttpStatus.notFound; - unawaited(request.response.close()); - } + server.listen((request) { + unawaited(_serveHostRequest(request, server, deviceName, platform, authContexts)); }); _connectionStateController.add(RemoteSessionStatus.connected); return (addresses: addresses, port: port); } catch (e) { + _acceptingHostConnections = false; + final failedServer = _server; + _server = null; + await _runDisconnectCleanup(failedServer?.close(force: true), 'failed server'); appLogger.e('CompanionRemote: Failed to create server', error: e); _errorController.add( RemotePeerError( @@ -182,208 +264,164 @@ class CompanionRemotePeerService with KeepaliveMixin { } } - void _handleNewWebSocketConnection( - WebSocket socket, + Future _serveHostRequest( + HttpRequest request, + HttpServer server, String hostDeviceName, String hostPlatform, List authContexts, - String sourceIp, - ) { - appLogger.d('CompanionRemote: New WebSocket connection from $sourceIp'); - - bool isAuthenticated = false; - Timer? authTimeout; - final auth = RemoteAuthService.instance; - final hostNonce = auth.generateNonce(); - - // Check rate limiting - final lockout = _authLockouts[sourceIp]; - if (lockout != null && DateTime.now().isBefore(lockout)) { - appLogger.w('CompanionRemote: Connection from $sourceIp rejected (rate limited)'); - socket.close(4005, 'Rate limited'); + ) async { + if (request.uri.path != '/ws') { + request.response.statusCode = HttpStatus.notFound; + await request.response.close(); return; } + final sourceIp = request.connectionInfo?.remoteAddress.address ?? 'unknown'; + final admission = _tryReserveHostAdmission(sourceIp, server); + if (admission == null) { + request.response.statusCode = HttpStatus.tooManyRequests; + request.response.headers.contentLength = 0; + await request.response.close(); + return; + } + + try { + final socket = await WebSocketTransformer.upgrade(request, compression: CompressionOptions.compressionOff); + admission.socket = socket; + admission.completeUpgrade(); + _afterHostUpgrade?.call(); + + if (!_isHostAdmissionLive(admission, phase: _HostAdmissionPhase.upgrading)) { + await _closeHostAdmissionSocket(admission); + return; + } + + _handleNewWebSocketConnection(admission, hostDeviceName, hostPlatform, authContexts); + } catch (e) { + admission.completeUpgrade(); + await _closeHostAdmissionSocket(admission); + try { + request.response.statusCode = HttpStatus.badRequest; + await request.response.close(); + } catch (_) { + // The upgrade path may already have committed and closed the response. + } + appLogger.d('CompanionRemote: WebSocket upgrade rejected', error: e); + } + } + + _HostAdmission? _tryReserveHostAdmission(String sourceIp, HttpServer server) { + if (!_acceptingHostConnections || !identical(_server, server) || _isSourceLockedOut(sourceIp)) { + return null; + } + + final sourceCount = _hostAdmissionsBySource[sourceIp] ?? 0; + if (_hostAdmissionCount >= _maxTotalHostConnections || sourceCount >= _maxHostConnectionsPerSource) { + return null; + } + + final admission = _HostAdmission(sourceIp: sourceIp, server: server); + _hostAdmissions.add(admission); + _hostAdmissionCount++; + _hostAdmissionsBySource[sourceIp] = sourceCount + 1; + return admission; + } + + bool _isSourceLockedOut(String sourceIp) { + final lockout = _authLockouts[sourceIp]; + if (lockout == null) return false; + if (DateTime.now().isBefore(lockout)) return true; + _authLockouts.remove(sourceIp); + _failedAuthAttempts.remove(sourceIp); + return false; + } + + bool _isHostAdmissionLive(_HostAdmission admission, {required _HostAdmissionPhase phase, int? commitGeneration}) { + final socket = admission.socket; + return !admission.released && + admission.phase == phase && + identical(_server, admission.server) && + _acceptingHostConnections && + socket != null && + socket.readyState == WebSocket.open && + (commitGeneration == null || commitGeneration == _authenticationCommitGeneration); + } + + void _handleNewWebSocketConnection( + _HostAdmission admission, + String hostDeviceName, + String hostPlatform, + List authContexts, + ) { + final socket = admission.socket!; + final auth = RemoteAuthService.instance; + final hostNonce = auth.generateNonce(); final primaryContext = authContexts.first; - // Send challenge: legacy hostClientId plus all selectable auth contexts. - socket.add( - jsonEncode({ - 'type': 'challenge', - 'nonce': base64Encode(hostNonce), - 'hostClientId': primaryContext.clientIdentifier, - 'authContexts': [ - for (final context in authContexts) {'id': context.id, 'hostClientId': context.clientIdentifier}, - ], - }), - ); + appLogger.d('CompanionRemote: New WebSocket connection from ${admission.sourceIp}'); - // Authentication timeout - authTimeout = Timer(const Duration(seconds: 10), () { - if (!isAuthenticated) { + admission.phase = _HostAdmissionPhase.awaitingAuth; + try { + socket.add( + jsonEncode({ + 'type': 'challenge', + 'nonce': base64Encode(hostNonce), + 'hostClientId': primaryContext.clientIdentifier, + 'authContexts': [ + for (final context in authContexts) {'id': context.id, 'hostClientId': context.clientIdentifier}, + ], + }), + ); + } catch (e) { + appLogger.d('CompanionRemote: Failed to send authentication challenge', error: e); + unawaited(_closeHostAdmissionSocket(admission)); + return; + } + + admission.authTimer = Timer(_authTimeout, () { + if (admission.phase == _HostAdmissionPhase.awaitingAuth || + admission.phase == _HostAdmissionPhase.authenticating) { appLogger.w('CompanionRemote: Authentication timeout'); - socket.close(4001, 'Authentication timeout'); + unawaited(_closeHostAdmissionSocket(admission, code: 4001, reason: 'Authentication timeout')); } }); late final StreamSubscription socketSubscription; socketSubscription = socket.listen( - (data) async { - try { - if (!isAuthenticated) { - final json = jsonDecode(data as String) as Map; - - if (json['type'] == 'auth') { - final authTag = json['authTag'] as String?; - final clientNonceB64 = json['clientNonce'] as String?; - final userUUID = json['userUUID'] as String?; - final clientIdentifier = json['clientIdentifier'] as String?; - final deviceName = json['deviceName'] as String?; - final platform = json['platform'] as String?; - final authContextId = json['authContextId'] as String?; - - if (authTag == null || - clientNonceB64 == null || - userUUID == null || - clientIdentifier == null || - deviceName == null || - platform == null) { - socket.add(jsonEncode({'type': 'authFailed'})); - unawaited(socket.close(4003, 'Authentication failed')); - return; - } - - final clientNonce = base64Decode(clientNonceB64); - RemoteAuthContext? selectedContext; - if (authContextId != null && authContextId.isNotEmpty) { - for (final context in authContexts) { - if (context.id == authContextId) { - selectedContext = context; - break; - } - } - } else if (authContexts.length == 1) { - selectedContext = primaryContext; - } - - if (selectedContext == null) { - _recordFailedAuth(sourceIp); - appLogger.w('CompanionRemote: Auth failed — unknown auth context'); - socket.add(jsonEncode({'type': 'authFailed'})); - unawaited(socket.close(4003, 'Authentication failed')); - return; - } - - // Verify userUUID is allowed for the selected profile connection. - if (selectedContext.allowedUserUuids.isNotEmpty && !selectedContext.allowedUserUuids.contains(userUUID)) { - _recordFailedAuth(sourceIp); - appLogger.w('CompanionRemote: Auth failed — unknown user'); - socket.add(jsonEncode({'type': 'authFailed'})); - unawaited(socket.close(4003, 'Authentication failed')); - return; - } - - // Verify auth tag - final valid = auth.verifyAuthTag( - authTag: authTag, - homeSecret: selectedContext.homeSecret, - hostNonce: hostNonce, - clientNonce: clientNonce, - hostClientId: selectedContext.clientIdentifier, - userUUID: userUUID, - clientIdentifier: clientIdentifier, - deviceName: deviceName, - platform: platform, - ); - - if (!valid) { - _recordFailedAuth(sourceIp); - appLogger.w('CompanionRemote: Auth failed — invalid auth tag'); - socket.add(jsonEncode({'type': 'authFailed'})); - unawaited(socket.close(4003, 'Authentication failed')); - return; - } - - // Auth success — derive per-session encryption key - _failedAuthAttempts.remove(sourceIp); - isAuthenticated = true; - authTimeout?.cancel(); - - final sessionEncKey = await auth.deriveSessionEncKey(selectedContext.homeSecret, hostNonce, clientNonce); - - // Close existing client if present - if (_clientSocket != null) { - appLogger.d('CompanionRemote: Replacing existing client connection'); - unawaited(_clientSocket!.close(4004, 'Replaced by new connection')); - } - - final previousSubscription = _clientSocketSubscription; - if (previousSubscription != null) { - try { - await previousSubscription.cancel(); - } catch (e) { - appLogger.d('CompanionRemote: previous client listener cancel ignored', error: e); - } - } - _clientSocketSubscription = socketSubscription; - _clientSocket = socket; - _sessionEncKey = sessionEncKey; - _sendCounter = 0; - _recvCounter = 0; - _isAuthenticated = true; - _selectedAuthContextId = selectedContext.id; - _selectedHostClientId = selectedContext.clientIdentifier; - - appLogger.d('CompanionRemote: Client authenticated: $deviceName ($platform)'); - - // Send encrypted authSuccess - await _sendEncryptedToSocket(socket, jsonEncode({'type': 'authSuccess'})); - - // Notify connection - final device = RemoteDevice( - id: 'remote-client', - name: deviceName, - platform: platform, - connectedAt: DateTime.now(), - ); - _deviceConnectedController.add(device); - _connectionStateController.add(RemoteSessionStatus.connected); - - // Send device info - sendDeviceInfo(hostDeviceName, hostPlatform); - } else { - appLogger.w('CompanionRemote: Expected auth, got ${json['type']}'); - unawaited(socket.close(4002, 'Authentication required')); - } - } else { - await _handleEncryptedCommand(data); - } - } catch (e) { - appLogger.e('CompanionRemote: Failed to process message', error: e); + (data) { + if (admission.phase == _HostAdmissionPhase.awaitingAuth) { + // Claim the only pre-auth message before parsing or awaiting. + admission.phase = _HostAdmissionPhase.authenticating; + unawaited( + _authenticateHostAdmission( + admission, + data, + hostNonce, + hostDeviceName, + hostPlatform, + authContexts, + primaryContext, + ), + ); + } else if (admission.phase == _HostAdmissionPhase.authenticated) { + unawaited( + _handleEncryptedCommand(data).catchError((Object error, StackTrace stackTrace) { + appLogger.e('CompanionRemote: Failed to process encrypted message', error: error, stackTrace: stackTrace); + }), + ); } }, onDone: () { - unawaited(socketSubscription.cancel()); - authTimeout?.cancel(); appLogger.d('CompanionRemote: WebSocket connection closed'); - // A replaced client's socket closes AFTER the new client already took - // over `_clientSocket`; only the socket that still owns the session may - // tear it down, or we'd clobber the live connection. - if (isAuthenticated && identical(_clientSocket, socket)) { - _clientSocket = null; - _clientSocketSubscription = null; - _sessionEncKey = null; - _isAuthenticated = false; - _selectedAuthContextId = null; - _selectedHostClientId = null; - _deviceDisconnectedController.add(null); - _connectionStateController.add(RemoteSessionStatus.disconnected); - stopKeepalive(); - } + admission.phase = _HostAdmissionPhase.terminal; + _releaseHostAdmission(admission); }, - onError: (error) { - authTimeout?.cancel(); - appLogger.e('CompanionRemote: WebSocket error', error: error); + onError: (Object error, StackTrace stackTrace) { + admission.phase = _HostAdmissionPhase.terminal; + admission.authTimer?.cancel(); + admission.authTimer = null; + appLogger.e('CompanionRemote: WebSocket error', error: error, stackTrace: stackTrace); _errorController.add( RemotePeerError( type: RemotePeerErrorType.dataChannelError, @@ -391,8 +429,300 @@ class CompanionRemotePeerService with KeepaliveMixin { originalError: error, ), ); + unawaited(_closeHostAdmissionSocket(admission)); }, + cancelOnError: true, ); + admission.subscription = socketSubscription; + } + + Future _authenticateHostAdmission( + _HostAdmission admission, + dynamic data, + List hostNonce, + String hostDeviceName, + String hostPlatform, + List authContexts, + RemoteAuthContext primaryContext, + ) async { + // dart:io delivers an assembled WebSocket message. This bounds application + // parsing and allocation after delivery, not the runtime's frame buffer. + if (data is! String || data.length > _maxPreAuthMessageBytes) { + _rejectHostAuthentication(admission); + return; + } + + late final Map message; + try { + if (utf8.encode(data).length > _maxPreAuthMessageBytes) { + _rejectHostAuthentication(admission); + return; + } + final decoded = jsonDecode(data); + if (decoded is! Map) { + _rejectHostAuthentication(admission); + return; + } + message = decoded; + } catch (_) { + _rejectHostAuthentication(admission); + return; + } + + if (message['type'] != 'auth') { + _rejectHostAuthentication(admission, closeCode: 4002, closeReason: 'Authentication required'); + return; + } + + final authTag = message['authTag']; + final clientNonceB64 = message['clientNonce']; + final userUuid = message['userUUID']; + final clientIdentifier = message['clientIdentifier']; + final deviceName = message['deviceName']; + final platform = message['platform']; + final authContextIdValue = message['authContextId']; + if (authTag is! String || + clientNonceB64 is! String || + userUuid is! String || + clientIdentifier is! String || + deviceName is! String || + platform is! String || + (authContextIdValue != null && authContextIdValue is! String)) { + _rejectHostAuthentication(admission); + return; + } + + late final List clientNonce; + try { + clientNonce = base64Decode(clientNonceB64); + } catch (_) { + _rejectHostAuthentication(admission); + return; + } + if (clientNonce.length != 32) { + _rejectHostAuthentication(admission); + return; + } + + final authContextId = authContextIdValue as String?; + RemoteAuthContext? selectedContext; + if (authContextId != null && authContextId.isNotEmpty) { + for (final context in authContexts) { + if (context.id == authContextId) { + selectedContext = context; + break; + } + } + } else if (authContexts.length == 1) { + selectedContext = primaryContext; + } + + if (selectedContext == null || + (selectedContext.allowedUserUuids.isNotEmpty && !selectedContext.allowedUserUuids.contains(userUuid))) { + _rejectHostAuthentication(admission); + return; + } + final authenticatedContext = selectedContext; + + final valid = RemoteAuthService.instance.verifyAuthTag( + authTag: authTag, + homeSecret: authenticatedContext.homeSecret, + hostNonce: hostNonce, + clientNonce: clientNonce, + hostClientId: authenticatedContext.clientIdentifier, + userUUID: userUuid, + clientIdentifier: clientIdentifier, + deviceName: deviceName, + platform: platform, + ); + if (!valid) { + _rejectHostAuthentication(admission); + return; + } + + admission.authTimer?.cancel(); + admission.authTimer = null; + _failedAuthAttempts.remove(admission.sourceIp); + + late final List sessionEncKey; + try { + sessionEncKey = await _deriveSessionEncKey(authenticatedContext.homeSecret, hostNonce, clientNonce); + } catch (e, stackTrace) { + appLogger.e('CompanionRemote: Failed to derive session key', error: e, stackTrace: stackTrace); + await _closeHostAdmissionSocket(admission, code: 4003, reason: 'Authentication failed'); + return; + } + + await _serializeHostAuthenticationCommit( + () => _commitAuthenticatedHostAdmission( + admission: admission, + sessionEncKey: sessionEncKey, + selectedContext: authenticatedContext, + deviceName: deviceName, + platform: platform, + hostDeviceName: hostDeviceName, + hostPlatform: hostPlatform, + ), + ); + } + + Future _serializeHostAuthenticationCommit(Future Function() commit) { + final operation = _hostAuthenticationCommitTail.then((_) => commit()); + _hostAuthenticationCommitTail = operation.catchError((Object _, StackTrace _) {}); + return operation; + } + + Future _commitAuthenticatedHostAdmission({ + required _HostAdmission admission, + required List sessionEncKey, + required RemoteAuthContext selectedContext, + required String deviceName, + required String platform, + required String hostDeviceName, + required String hostPlatform, + }) async { + if (!_isHostAdmissionLive(admission, phase: _HostAdmissionPhase.authenticating)) { + await _closeHostAdmissionSocket(admission); + return; + } + + final previousAdmission = _currentHostAdmission; + if (previousAdmission != null && !identical(previousAdmission, admission)) { + appLogger.d('CompanionRemote: Replacing existing client connection'); + await _retireReplacedHostAdmission(previousAdmission); + } + if (!_isHostAdmissionLive(admission, phase: _HostAdmissionPhase.authenticating)) { + await _closeHostAdmissionSocket(admission); + return; + } + + final commitGeneration = ++_authenticationCommitGeneration; + admission.commitGeneration = commitGeneration; + admission.phase = _HostAdmissionPhase.authenticated; + _currentHostAdmission = admission; + _clientSocket = admission.socket; + _sessionEncKey = sessionEncKey; + _sendCounter = 0; + _recvCounter = 0; + _isAuthenticated = true; + _selectedAuthContextId = selectedContext.id; + _selectedHostClientId = selectedContext.clientIdentifier; + + appLogger.d('CompanionRemote: Client authenticated: $deviceName ($platform)'); + try { + await _sendEncryptedToSocket(admission.socket!, jsonEncode({'type': 'authSuccess'})); + } catch (e, stackTrace) { + appLogger.e('CompanionRemote: Failed to send authentication result', error: e, stackTrace: stackTrace); + await _closeHostAdmissionSocket(admission); + return; + } + if (!_isHostAdmissionLive( + admission, + phase: _HostAdmissionPhase.authenticated, + commitGeneration: commitGeneration, + )) { + await _closeHostAdmissionSocket(admission); + return; + } + + final device = RemoteDevice(id: 'remote-client', name: deviceName, platform: platform, connectedAt: DateTime.now()); + _deviceConnectedController.add(device); + _connectionStateController.add(RemoteSessionStatus.connected); + sendDeviceInfo(hostDeviceName, hostPlatform); + } + + void _rejectHostAuthentication( + _HostAdmission admission, { + int closeCode = 4003, + String closeReason = 'Authentication failed', + }) { + if (admission.phase != _HostAdmissionPhase.authenticating) return; + admission.phase = _HostAdmissionPhase.terminal; + admission.authTimer?.cancel(); + admission.authTimer = null; + _recordFailedAuth(admission.sourceIp); + + final socket = admission.socket; + if (socket != null && socket.readyState == WebSocket.open) { + try { + socket.add(jsonEncode({'type': 'authFailed'})); + } catch (_) { + // The generic close below remains the terminal result. + } + } + unawaited(_closeHostAdmissionSocket(admission, code: closeCode, reason: closeReason)); + } + + Future _retireReplacedHostAdmission(_HostAdmission admission) async { + admission.phase = _HostAdmissionPhase.terminal; + admission.authTimer?.cancel(); + admission.authTimer = null; + _clearHostSessionIfOwned(admission, notify: false); + final subscription = admission.subscription; + admission.subscription = null; + await _runDisconnectCleanup(subscription?.cancel(), 'replaced client listener'); + unawaited(_closeHostAdmissionSocket(admission, code: 4004, reason: 'Replaced by new connection')); + } + + Future _closeHostAdmissionSocket(_HostAdmission admission, {int? code, String? reason}) { + final existing = admission.closeFuture; + if (existing != null) return existing; + final closeFuture = _closeHostAdmissionSocketOnce(admission, code: code, reason: reason); + admission.closeFuture = closeFuture; + return closeFuture; + } + + Future _closeHostAdmissionSocketOnce(_HostAdmission admission, {int? code, String? reason}) async { + admission.phase = _HostAdmissionPhase.terminal; + admission.authTimer?.cancel(); + admission.authTimer = null; + final socket = admission.socket; + if (socket != null) { + await _runDisconnectCleanup(socket.close(code, reason), 'host socket'); + } + _releaseHostAdmission(admission); + } + + void _releaseHostAdmission(_HostAdmission admission) { + if (admission.released) return; + admission.released = true; + admission.phase = _HostAdmissionPhase.terminal; + admission.authTimer?.cancel(); + admission.authTimer = null; + + final subscription = admission.subscription; + admission.subscription = null; + if (subscription != null) { + unawaited(_runDisconnectCleanup(subscription.cancel(), 'host listener')); + } + + if (_hostAdmissions.remove(admission)) { + _hostAdmissionCount--; + final sourceCount = (_hostAdmissionsBySource[admission.sourceIp] ?? 1) - 1; + if (sourceCount <= 0) { + _hostAdmissionsBySource.remove(admission.sourceIp); + } else { + _hostAdmissionsBySource[admission.sourceIp] = sourceCount; + } + } + + _clearHostSessionIfOwned(admission, notify: !_isDisconnecting); + admission.completeTerminal(); + } + + void _clearHostSessionIfOwned(_HostAdmission admission, {required bool notify}) { + if (!identical(_currentHostAdmission, admission)) return; + _currentHostAdmission = null; + _clientSocket = null; + _sessionEncKey = null; + _isAuthenticated = false; + _selectedAuthContextId = null; + _selectedHostClientId = null; + stopKeepalive(); + if (notify) { + _deviceDisconnectedController.add(null); + _connectionStateController.add(RemoteSessionStatus.disconnected); + } } void _recordFailedAuth(String sourceIp) { @@ -404,6 +734,10 @@ class CompanionRemotePeerService with KeepaliveMixin { } } + bool _ownsRemoteChannel(IOWebSocketChannel channel, int generation) { + return !_disposed && generation == _remoteConnectionGeneration && identical(_channel, channel); + } + /// Join a host session with any local auth context that the host also supports. Future joinSessionWithContexts( String deviceName, @@ -423,6 +757,7 @@ class CompanionRemotePeerService with KeepaliveMixin { if (_channel != null) { await disconnect(); } + final connectionGeneration = ++_remoteConnectionGeneration; _role = RemoteSessionRole.remote; _hostAddress = hostAddress; @@ -430,6 +765,7 @@ class CompanionRemotePeerService with KeepaliveMixin { final completer = Completer(); final auth = RemoteAuthService.instance; + IOWebSocketChannel? attemptedChannel; try { final url = 'ws://$hostAddress/ws'; @@ -437,22 +773,29 @@ class CompanionRemotePeerService with KeepaliveMixin { _connectionStateController.add(RemoteSessionStatus.connecting); - _channel = IOWebSocketChannel.connect(Uri.parse(url)); - await _channel!.ready; + final channel = IOWebSocketChannel.connect(Uri.parse(url)); + attemptedChannel = channel; + _channel = channel; + await channel.ready; + if (!_ownsRemoteChannel(channel, connectionGeneration)) { + unawaited(channel.sink.close()); + throw StateError('Companion Remote connection attempt became stale'); + } List? hostNonce; List? clientNonce; String? receivedHostClientId; - _channelSubscription = _channel!.stream.listen( + _channelSubscription = channel.stream.listen( (data) async { + if (!_ownsRemoteChannel(channel, connectionGeneration)) return; try { if (_isAuthenticated) { await _handleEncryptedCommand(data); } else if (_sessionEncKey != null) { // Keys derived, waiting for encrypted authSuccess final decrypted = await _decryptIncoming(data); - if (decrypted == null) return; + if (decrypted == null || !_ownsRemoteChannel(channel, connectionGeneration)) return; final json = jsonDecode(decrypted) as Map; if (json['type'] == 'authSuccess') { @@ -533,7 +876,7 @@ class CompanionRemotePeerService with KeepaliveMixin { ), ); } - unawaited(_channel?.sink.close(4003, 'Authentication failed')); + unawaited(channel.sink.close(4003, 'Authentication failed')); return; } @@ -550,7 +893,7 @@ class CompanionRemotePeerService with KeepaliveMixin { ), ); } - unawaited(_channel?.sink.close(4003, 'Authentication failed')); + unawaited(channel.sink.close(4003, 'Authentication failed')); return; } @@ -565,7 +908,7 @@ class CompanionRemotePeerService with KeepaliveMixin { platform: platform, ); - _channel!.sink.add( + channel.sink.add( jsonEncode({ 'type': 'auth', 'authContextId': selectedContext.id, @@ -578,7 +921,13 @@ class CompanionRemotePeerService with KeepaliveMixin { }), ); - _sessionEncKey = await auth.deriveSessionEncKey(selectedContext.homeSecret, hostNonce!, clientNonce!); + final sessionEncKey = await auth.deriveSessionEncKey( + selectedContext.homeSecret, + hostNonce!, + clientNonce!, + ); + if (!_ownsRemoteChannel(channel, connectionGeneration)) return; + _sessionEncKey = sessionEncKey; _sendCounter = 0; _recvCounter = 0; _selectedAuthContextId = selectedContext.id; @@ -607,10 +956,20 @@ class CompanionRemotePeerService with KeepaliveMixin { } }, onDone: () { + if (!_ownsRemoteChannel(channel, connectionGeneration)) return; appLogger.d('CompanionRemote: Connection closed'); + if (!completer.isCompleted) { + completer.completeError( + RemotePeerError( + type: RemotePeerErrorType.connectionFailed, + message: t.companionRemote.pairing.failedToConnect(error: 'Connection closed before authentication'), + ), + ); + } _deviceDisconnectedController.add(null); _connectionStateController.add(RemoteSessionStatus.disconnected); _isAuthenticated = false; + _channel = null; _channelSubscription = null; _sessionEncKey = null; _selectedAuthContextId = null; @@ -618,6 +977,7 @@ class CompanionRemotePeerService with KeepaliveMixin { stopKeepalive(); }, onError: (error) { + if (!_ownsRemoteChannel(channel, connectionGeneration)) return; appLogger.e('CompanionRemote: Connection error', error: error); if (!completer.isCompleted) { @@ -631,35 +991,48 @@ class CompanionRemotePeerService with KeepaliveMixin { originalError: error, ), ); + _isAuthenticated = false; + _channel = null; + _channelSubscription = null; + _sessionEncKey = null; + _selectedAuthContextId = null; + _selectedHostClientId = null; + stopKeepalive(); _connectionStateController.add(RemoteSessionStatus.error); }, ); } catch (e) { - appLogger.e('CompanionRemote: Failed to connect', error: e); - if (!completer.isCompleted) { completer.completeError(e); } - - _errorController.add( - RemotePeerError( - type: RemotePeerErrorType.connectionFailed, - message: t.companionRemote.pairing.failedToConnect(error: e.toString()), - originalError: e, - ), - ); + final channel = attemptedChannel; + if (channel != null && _ownsRemoteChannel(channel, connectionGeneration)) { + appLogger.e('CompanionRemote: Failed to connect', error: e); + _errorController.add( + RemotePeerError( + type: RemotePeerErrorType.connectionFailed, + message: t.companionRemote.pairing.failedToConnect(error: e.toString()), + originalError: e, + ), + ); + } } + final channel = attemptedChannel; + if (channel == null) return completer.future; + return completer.future.timeout( const Duration(seconds: 15), onTimeout: () async { - if (_channel != null) { + if (_ownsRemoteChannel(channel, connectionGeneration)) { try { - await _channel!.sink.close(); + await channel.sink.close(); } catch (e) { appLogger.d('CompanionRemote: channel close on timeout failed', error: e); } - _channel = null; + if (_ownsRemoteChannel(channel, connectionGeneration)) { + _channel = null; + } } throw RemotePeerError(type: RemotePeerErrorType.timeout, message: t.companionRemote.errors.joinTimedOut); }, @@ -698,61 +1071,66 @@ class CompanionRemotePeerService with KeepaliveMixin { // Race: try to connect to all addresses, first one to get a challenge wins final completer = Completer(); - final channels = []; - final subs = []; + final probes = <_RemoteAddressProbe>[]; - void cleanup() { - for (final sub in subs) { - sub.cancel(); - } - for (final ch in channels) { - try { - ch.sink.close(); - } catch (e) { - appLogger.d('CompanionRemote: race-loser close ignored', error: e); - } + Future cleanup() async { + // A probe consumes a host admission slot until its WebSocket has fully + // closed. Finish every cancellation and close before opening the managed + // connection so probes cannot reject that connection at the per-source + // limit. Cleanup errors are intentionally logged in address order and + // never replace the race/authentication result. + for (final probe in probes) { + await probe.close(); } } for (final address in hostAddresses) { + _RaceProbeConnection? connection; try { final url = 'ws://$address/ws'; - final channel = IOWebSocketChannel.connect(Uri.parse(url), connectTimeout: const Duration(seconds: 5)); - channels.add(channel); + connection = _raceProbeFactory(Uri.parse(url)); // Losing candidates fail their `ready` future (connect timeout, // no route to host, …); nothing awaits it here — the stream's // onError below is the visible signal — so swallow it or every // unreachable address becomes an unhandled async error. unawaited( - channel.ready.catchError((Object e) { + connection.ready.catchError((Object e) { appLogger.d('CompanionRemote: race candidate $address failed to connect', error: e); }), ); - final sub = channel.stream.listen( - (data) { - try { - final json = jsonDecode(data as String) as Map; - // First address to send us a challenge wins the race - if (json['type'] == 'challenge' && !completer.isCompleted) { - appLogger.d('CompanionRemote: Race winner: $address'); - completer.complete(address); + probes.add( + _RemoteAddressProbe( + requestClose: connection.close, + stream: connection.stream, + onData: (data) { + try { + final json = jsonDecode(data as String) as Map; + // First address to send us a challenge wins the race + if (json['type'] == 'challenge' && !completer.isCompleted) { + appLogger.d('CompanionRemote: Race winner: $address'); + completer.complete(address); + } + } catch (e) { + appLogger.d('CompanionRemote: race message parse skipped', error: e); } - } catch (e) { - appLogger.d('CompanionRemote: race message parse skipped', error: e); - } - }, - onError: (_) {}, - onDone: () {}, + }, + ), ); - subs.add(sub); } catch (e) { appLogger.d('CompanionRemote: Race candidate $address failed to start: $e'); + if (connection != null) { + try { + await connection.close(); + } catch (closeError) { + appLogger.d('CompanionRemote: failed race candidate close ignored', error: closeError); + } + } } } - if (channels.isEmpty) { + if (probes.isEmpty) { throw RemotePeerError( type: RemotePeerErrorType.connectionFailed, message: t.companionRemote.errors.failedToConnectAnyAddress, @@ -764,7 +1142,7 @@ class CompanionRemotePeerService with KeepaliveMixin { const Duration(seconds: 10), operation: 'CompanionRemote race connect', ); - cleanup(); + await cleanup(); // Set up the proper managed connection on the winning address await joinSessionWithContexts( @@ -777,7 +1155,7 @@ class CompanionRemotePeerService with KeepaliveMixin { ); return winner; } on TimeoutException { - cleanup(); + await cleanup(); throw RemotePeerError(type: RemotePeerErrorType.timeout, message: t.companionRemote.errors.joinTimedOut); } } @@ -926,26 +1304,42 @@ class CompanionRemotePeerService with KeepaliveMixin { } } - Future disconnect() async { - appLogger.d('CompanionRemote: Disconnecting'); + Future disconnect() { + if (_disposed) return Future.value(); + final existing = _disconnectInProgress; + if (existing != null) return existing; + late final Future tracked; + tracked = _disconnect().whenComplete(() { + if (identical(_disconnectInProgress, tracked)) { + _disconnectInProgress = null; + } + }); + _disconnectInProgress = tracked; + return tracked; + } + + Future _disconnect() async { + appLogger.d('CompanionRemote: Disconnecting'); + _remoteConnectionGeneration++; + + _acceptingHostConnections = false; + _isDisconnecting = true; _isAuthenticated = false; + _authenticationCommitGeneration++; stopKeepalive(); - final clientSocket = _clientSocket; final channel = _channel; final server = _server; - _server = null; + final serverClose = server?.close(force: true); + final admissions = List<_HostAdmission>.of(_hostAdmissions); try { - // Stop inbound callbacks before taking queue snapshots. New sends are - // already rejected by `_isAuthenticated = false`. - await _runDisconnectCleanup(_clientSocketSubscription?.cancel(), 'client listener'); - _clientSocketSubscription = null; + await Future.wait(admissions.map(_shutdownHostAdmission)); + await _runDisconnectCleanup(_channelSubscription?.cancel(), 'channel listener'); _channelSubscription = null; - final serverClose = server?.close(force: true); // Decrypted commands may enqueue acknowledgements, and sends enqueue // encryption, so drain in dependency order. @@ -953,10 +1347,16 @@ class CompanionRemotePeerService with KeepaliveMixin { await _sendQueue.settled; await _encryptQueue.settled; - await _runDisconnectCleanup(clientSocket?.close(), 'client socket'); await _runDisconnectCleanup(channel?.sink.close(), 'channel'); await _runDisconnectCleanup(serverClose, 'server'); } finally { + for (final admission in List<_HostAdmission>.of(_hostAdmissions)) { + _releaseHostAdmission(admission); + } + _hostAdmissions.clear(); + _hostAdmissionsBySource.clear(); + _hostAdmissionCount = 0; + _currentHostAdmission = null; _clientSocket = null; _channel = null; _myPeerId = null; @@ -972,20 +1372,128 @@ class CompanionRemotePeerService with KeepaliveMixin { _decryptQueue.reset(); _failedAuthAttempts.clear(); _authLockouts.clear(); + _isDisconnecting = false; _connectionStateController.add(RemoteSessionStatus.disconnected); } } + Future _shutdownHostAdmission(_HostAdmission admission) async { + admission.phase = _HostAdmissionPhase.terminal; + admission.authTimer?.cancel(); + admission.authTimer = null; + + await admission.upgradeFinished.future; + if (admission.released) return; + + final subscription = admission.subscription; + final socketClose = _closeHostAdmissionSocket(admission); + admission.subscription = null; + await _runDisconnectCleanup(subscription?.cancel(), 'host listener'); + await socketClose; + await admission.terminal.future; + } + /// Whether the HTTP server is currently running. bool get isServerRunning => _server != null; - Future dispose() async { + Future dispose() { + final existing = _disposeInProgress; + if (existing != null) return existing; + if (_disposed) return Future.value(); + + late final Future tracked; + tracked = _dispose().whenComplete(() { + if (identical(_disposeInProgress, tracked)) { + _disposeInProgress = null; + } + }); + _disposeInProgress = tracked; + return tracked; + } + + Future _dispose() async { await disconnect(); await _commandReceivedController.close(); await _deviceConnectedController.close(); await _deviceDisconnectedController.close(); await _errorController.close(); await _connectionStateController.close(); + _disposed = true; + } +} + +class _RemoteAddressProbe { + _RemoteAddressProbe({ + required this._requestClose, + required Stream stream, + required void Function(dynamic data) onData, + }) { + _subscription = stream.listen( + onData, + onError: (Object _, StackTrace _) => _completeTerminal(), + onDone: _completeTerminal, + ); + } + + static const _terminalTimeout = Duration(seconds: 5); + + final Future Function() _requestClose; + final Completer _terminal = Completer(); + late final StreamSubscription _subscription; + + void _completeTerminal() { + if (!_terminal.isCompleted) { + _terminal.complete(); + } + } + + Future close() async { + try { + await _requestClose(); + } catch (error) { + appLogger.d('CompanionRemote: race candidate close ignored', error: error); + } + try { + await _terminal.future.timeout(_terminalTimeout); + } on TimeoutException catch (error) { + appLogger.d('CompanionRemote: race candidate terminal close timed out', error: error); + } + try { + await _subscription.cancel(); + } catch (error) { + appLogger.d('CompanionRemote: race candidate listener cleanup ignored', error: error); + } + } +} + +enum _HostAdmissionPhase { upgrading, awaitingAuth, authenticating, authenticated, terminal } + +class _HostAdmission { + _HostAdmission({required this.sourceIp, required this.server}); + + final String sourceIp; + final HttpServer server; + final Completer upgradeFinished = Completer(); + final Completer terminal = Completer(); + + WebSocket? socket; + StreamSubscription? subscription; + Timer? authTimer; + _HostAdmissionPhase phase = _HostAdmissionPhase.upgrading; + int? commitGeneration; + bool released = false; + Future? closeFuture; + + void completeUpgrade() { + if (!upgradeFinished.isCompleted) { + upgradeFinished.complete(); + } + } + + void completeTerminal() { + if (!terminal.isCompleted) { + terminal.complete(); + } } } diff --git a/lib/services/settings_service.dart b/lib/services/settings_service.dart index 3c9dfefb..24d40a56 100644 --- a/lib/services/settings_service.dart +++ b/lib/services/settings_service.dart @@ -20,6 +20,8 @@ import '../models/transcode_quality_preset.dart'; import '../navigation/navigation_tabs.dart'; import '../utils/platform_detector.dart'; import 'trackers/tracker_constants.dart'; +import '../profiles/profile.dart'; +import '../watch_together/services/watch_together_relay_endpoint.dart'; enum ThemeMode { system, light, dark, oled } @@ -216,6 +218,15 @@ String? _trimEmptyAsNull(String? v) { return (t == null || t.isEmpty) ? null : t; } +String? _normalizeRelayBaseUrl(String? value) { + if (value == null || value.trim().isEmpty) return null; + final endpoint = WatchTogetherRelayEndpoint.tryParseCustom(value); + if (endpoint == null) { + throw FormatException('Invalid Watch Together relay base URL'); + } + return endpoint.canonicalBaseUrl; +} + String _legacyMpvEntriesToText(List entries) { final lines = []; for (final item in entries) { @@ -431,8 +442,15 @@ class SettingsService extends BaseSharedPreferencesService { static const appLocale = _AppLocalePref(); static const autoPip = _AutoPipPref(); static const customDownloadPath = NullableStringPref('custom_download_path'); - static final customRelayUrl = NullableStringPref('custom_relay_url', transform: _trimEmptyAsNull); - static const recentRooms = NullableStringPref('watch_together_recent_rooms'); + static final customRelayUrl = NullableStringPref('custom_relay_url', transform: _normalizeRelayBaseUrl); + + static NullableStringPref recentRoomsForProfile(String profileId) { + if (profileId.trim().isEmpty) { + throw ArgumentError.value(profileId, 'profileId', 'Must not be empty'); + } + return NullableStringPref(profileScopedPrefsKey(profileId, 'watch_together_recent_rooms')); + } + static final companionRemoteLastHostAddress = NullableStringPref( 'companion_remote_last_host_address', transform: _trimEmptyAsNull, @@ -576,6 +594,21 @@ class SettingsService extends BaseSharedPreferencesService { _cachedInstance = null; } + @override + Future onInit() async { + const legacyRecentRoomsKey = 'watch_together_recent_rooms'; + await prefs.remove(legacyRecentRoomsKey); + + final storedRelay = prefs.getString(customRelayUrl.key); + if (storedRelay == null) return; + final endpoint = WatchTogetherRelayEndpoint.tryParseCustom(storedRelay); + if (endpoint == null) { + await prefs.remove(customRelayUrl.key); + } else if (endpoint.canonicalBaseUrl != storedRelay) { + await prefs.setString(customRelayUrl.key, endpoint.canonicalBaseUrl); + } + } + /// Resolves a video mute toggle without replacing the saved volume with 0. /// /// `persistedVolume` is the non-zero value callers should keep in [volume], diff --git a/lib/services/trackers/oauth_proxy_client.dart b/lib/services/trackers/oauth_proxy_client.dart index dd57a832..eae29d6f 100644 --- a/lib/services/trackers/oauth_proxy_client.dart +++ b/lib/services/trackers/oauth_proxy_client.dart @@ -8,7 +8,7 @@ import '../../utils/app_logger.dart'; import '../../utils/platform_http_client_stub.dart' if (dart.library.io) '../../utils/platform_http_client_io.dart' as platform; -import '../../watch_together/services/watch_together_peer_service.dart'; +import '../../watch_together/services/watch_together_relay_endpoint.dart'; import 'tracker_constants.dart'; /// Client for the Plezy relay's `/auth/*` OAuth proxy. @@ -19,7 +19,7 @@ import 'tracker_constants.dart'; /// scheme is required — works identically on TVs without a browser. class OAuthProxyClient { /// Public base URL of the Plezy relay; colocated with Watch Together. - static String get baseUrl => WatchTogetherPeerService.defaultBaseUrl; + static String get baseUrl => WatchTogetherRelayEndpoint.defaultEndpoint.canonicalBaseUrl; final http.Client _http; diff --git a/lib/watch_together/models/sync_message.dart b/lib/watch_together/models/sync_message.dart index 7e26a9b6..f3b9bcb6 100644 --- a/lib/watch_together/models/sync_message.dart +++ b/lib/watch_together/models/sync_message.dart @@ -2,7 +2,7 @@ import 'dart:convert'; import 'playback_state.dart'; -/// Types of sync messages sent over the relay data channel (protocol v2). +/// Types of sync messages sent over the relay data channel (protocol v3). enum SyncMessageType { /// Authoritative playback state broadcast by the host state, @@ -37,7 +37,7 @@ class SyncMessage { /// Current sync protocol version, carried on join messages. Peers with a /// different version are excluded from readiness gating and surfaced as /// needing an update. - static const int protocolVersion = 2; + static const int protocolVersion = 3; /// Type of this message final SyncMessageType type; diff --git a/lib/watch_together/primitives.dart b/lib/watch_together/primitives.dart index cdc53430..299adbb8 100644 --- a/lib/watch_together/primitives.dart +++ b/lib/watch_together/primitives.dart @@ -1,7 +1,5 @@ int watchTogetherSystemNowMs() => DateTime.now().millisecondsSinceEpoch; -String watchTogetherHostPeerId(String sessionId) => 'wt-${sessionId.toUpperCase()}'; - bool orderedStringListsEqual(List first, List second) { if (first.length != second.length) return false; for (var i = 0; i < first.length; i++) { diff --git a/lib/watch_together/providers/watch_together_provider.dart b/lib/watch_together/providers/watch_together_provider.dart index ab335dfc..1154e1c8 100644 --- a/lib/watch_together/providers/watch_together_provider.dart +++ b/lib/watch_together/providers/watch_together_provider.dart @@ -5,21 +5,26 @@ import 'dart:math'; import 'package:flutter/foundation.dart'; import '../../mpv/mpv.dart'; -import '../../services/settings_service.dart'; import '../../utils/app_logger.dart'; import '../models/playback_state.dart'; import '../models/sync_message.dart'; import '../models/watch_session.dart'; -import '../primitives.dart'; import '../services/current_playback_dispatcher.dart'; +import '../services/relay_protocol.g.dart'; import '../services/watch_together_controller.dart'; import '../services/watch_together_peer_service.dart'; +import '../services/watch_together_relay_endpoint.dart'; /// Callback type for when media switches (for guest navigation). Returns /// whether the switch was handled; unhandled keys are re-dispatched on the /// host's next state heartbeat. typedef MediaSwitchCallback = Future Function(String ratingKey, ServerId serverId, String mediaTitle); +typedef WatchTogetherPeerServiceFactory = WatchTogetherPeerService Function({WatchTogetherRelayEndpoint? endpoint}); + +WatchTogetherPeerService _createWatchTogetherPeerService({WatchTogetherRelayEndpoint? endpoint}) => + WatchTogetherPeerService(endpoint: endpoint); + /// Provider for Watch Together functionality /// /// This provider manages: @@ -29,9 +34,15 @@ typedef MediaSwitchCallback = Future Function(String ratingKey, ServerId s /// - Participant list /// - Media switching across the session class WatchTogetherProvider with ChangeNotifier { + WatchTogetherProvider({WatchTogetherPeerServiceFactory peerServiceFactory = _createWatchTogetherPeerService}) + : _peerServiceFactory = peerServiceFactory; + + final WatchTogetherPeerServiceFactory _peerServiceFactory; + WatchSession? _session; WatchTogetherPeerService? _peerService; WatchTogetherController? _controller; + PeerError? _recoverableTransportError; final List _participants = []; bool _isSyncing = false; bool _isWaitingForPeers = false; @@ -88,6 +99,7 @@ class WatchTogetherProvider with ChangeNotifier { StreamSubscription? _peerDisconnectedSubscription; StreamSubscription? _messageSubscription; StreamSubscription? _errorSubscription; + StreamSubscription? _sessionEndedSubscription; // Getters bool get isInSession => _session != null && _session!.state != SessionState.disconnected; @@ -210,9 +222,36 @@ class WatchTogetherProvider with ChangeNotifier { _controller?.requestState(); } + void _requestCurrentPlaybackSnapshotIfHostConnected(WatchTogetherPeerService peerService) { + if (!identical(_peerService, peerService)) return; + final session = _session; + final hostPeerId = session?.hostPeerId; + if (session == null || session.isHost || hostPeerId == null || !peerService.connectedPeers.contains(hostPeerId)) { + return; + } + _controller?.updateSession(session); + requestCurrentPlaybackSnapshot(); + } + /// Wire up reconnection handler to re-announce join and re-sync state void _wireReconnectHandler() { - _peerService!.onReconnected = () { + final peerService = _peerService!; + peerService.onReconnected = () { + if (_disposed || !identical(_peerService, peerService)) return; + + final recoverableError = _recoverableTransportError; + final session = _session; + if (recoverableError != null && + session != null && + session.state == SessionState.error && + session.errorMessage == recoverableError.message) { + _session = session.copyWith(state: SessionState.connected, errorMessage: null); + _recoverableTransportError = null; + notifyListeners(); + } else { + _recoverableTransportError = null; + } + _controller?.announceJoin(_displayName); _controller?.onReconnected(); }; @@ -302,6 +341,7 @@ class WatchTogetherProvider with ChangeNotifier { /// Create a new watch together session as host Future createSession({ required ControlMode controlMode, + required WatchTogetherRelayEndpoint relayEndpoint, String? displayName, String? sessionId, String? mediaRatingKey, @@ -314,8 +354,7 @@ class WatchTogetherProvider with ChangeNotifier { appLogger.d('WatchTogether: Creating session with control mode: $controlMode'); - final customRelayUrl = SettingsService.instanceOrNull?.read(SettingsService.customRelayUrl); - _peerService = WatchTogetherPeerService(customBaseUrl: customRelayUrl); + _peerService = _peerServiceFactory(endpoint: relayEndpoint); _setupPeerServiceListeners(); try { @@ -323,7 +362,7 @@ class WatchTogetherProvider with ChangeNotifier { _session = WatchSession.createAsHost( sessionId: createdSessionId, - hostPeerId: _peerService!.myPeerId!, + hostPeerId: _peerService!.hostPeerId!, controlMode: controlMode, mediaRatingKey: mediaRatingKey, mediaServerId: mediaServerId, @@ -351,113 +390,139 @@ class WatchTogetherProvider with ChangeNotifier { } /// Join an existing session as guest - Future joinSession(String sessionId, {String? displayName}) async { + Future joinSession( + String sessionId, { + required WatchTogetherRelayEndpoint relayEndpoint, + String? displayName, + }) async { // Clean up any existing session await leaveSession(); _playbackDispatcher.reset(); appLogger.d('WatchTogether: Joining session: $sessionId'); - final customRelayUrl = SettingsService.instanceOrNull?.read(SettingsService.customRelayUrl); - _peerService = WatchTogetherPeerService(customBaseUrl: customRelayUrl); + final peerService = _peerServiceFactory(endpoint: relayEndpoint); + _peerService = peerService; _setupPeerServiceListeners(); - _session = WatchSession.joinAsGuest(sessionId: sessionId); + final joiningSession = WatchSession.joinAsGuest(sessionId: sessionId); + _session = joiningSession; notifyListeners(); try { - await _peerService!.joinSession(sessionId); + await peerService.joinSession(sessionId); + if (!identical(_peerService, peerService)) { + throw StateError('Watch Together join became stale'); + } - // Session will be fully configured when we receive sessionConfig from host - _session = _session!.copyWith(state: SessionState.connected, hostPeerId: watchTogetherHostPeerId(sessionId)); + // Host authority comes from the relay setup response rather than the + // public room code or a client-derived routing label. + _session = joiningSession.copyWith(state: SessionState.connected, hostPeerId: peerService.hostPeerId); _displayName = displayName ?? _generateDisplayName(); - _controller = WatchTogetherController(peerService: _peerService!, session: _session!); + _controller = WatchTogetherController(peerService: peerService, session: _session!); _wireController(); _wireReconnectHandler(); // Add self to participants - _participants.add(Participant(peerId: _peerService!.myPeerId!, displayName: _displayName, isHost: false)); + _participants.add(Participant(peerId: peerService.myPeerId!, displayName: _displayName, isHost: false)); // Announce join to other participants _controller!.announceJoin(_displayName); - requestCurrentPlaybackSnapshot(); + _requestCurrentPlaybackSnapshotIfHostConnected(peerService); notifyListeners(); appLogger.d('WatchTogether: Joined session successfully'); } catch (e) { appLogger.e('WatchTogether: Failed to join session', error: e); - await leaveSession(); + if (identical(_peerService, peerService)) { + await leaveSession(); + } rethrow; } } - /// Enter a room by code — joins if it exists, creates if empty. + /// Enter a room by code — joins any reserved room and creates only when the + /// relay reports that no room reservation exists. /// /// Returns `true` if the user became the host. - Future enterRoom(String sessionId, {ControlMode controlMode = ControlMode.anyone, String? displayName}) async { - // Probe the relay with a lightweight peer service to check room occupancy, - // then do a single createSession or joinSession. This avoids the crash-prone - // join→teardown→create cycle on the provider. - final customRelayUrl = SettingsService.instanceOrNull?.read(SettingsService.customRelayUrl); - final probe = WatchTogetherPeerService(customBaseUrl: customRelayUrl); - bool shouldBeHost; + Future enterRoom( + String sessionId, { + required WatchTogetherRelayEndpoint relayEndpoint, + ControlMode controlMode = ControlMode.anyone, + String? displayName, + }) async { + // A successful join reserves a durable guest identity. Release that probe + // identity before opening the provider's real connection. + final probe = _peerServiceFactory(endpoint: relayEndpoint); + var shouldBeHost = false; try { - await probe.joinSession(sessionId); - shouldBeHost = probe.connectedPeers.isEmpty; - } on PeerError catch (e) { - if (e.serverCode == 'room_not_found') { + try { + await probe.joinSession(sessionId); + } on PeerError catch (error) { + if (error.serverCode != RelayProtocol.roomNotFoundCode) rethrow; shouldBeHost = true; - } else { - await probe.disconnect(); - probe.dispose(); - rethrow; } + if (!shouldBeHost) { + await probe.releaseSession(); + } + } finally { + await probe.disconnect(); + probe.dispose(); } - await probe.disconnect(); - probe.dispose(); if (shouldBeHost) { - await createSession(controlMode: controlMode, displayName: displayName, sessionId: sessionId); + await createSession( + controlMode: controlMode, + relayEndpoint: relayEndpoint, + displayName: displayName, + sessionId: sessionId, + ); } else { - await joinSession(sessionId, displayName: displayName); + await joinSession(sessionId, relayEndpoint: relayEndpoint, displayName: displayName); } return shouldBeHost; } - /// Leave the current session + /// Leave the current session. Local callbacks and player bindings are + /// detached synchronously; relay release remains awaitable and observable. Future leaveSession() async { - if (_session == null) return; - + if (_session == null && _peerService == null) return; appLogger.d('WatchTogether: Leaving session'); + final peerService = _detachLocalSession(announceLeave: true); + if (peerService != null) { + await _finishPeerTeardown(peerService, release: true); + } + appLogger.d('WatchTogether: Session left'); + } - // Announce leave if connected - _controller?.announceLeave(); + WatchTogetherPeerService? _detachLocalSession({required bool announceLeave}) { + _recoverableTransportError = null; + if (announceLeave) _controller?.announceLeave(); - // Clean up subscriptions - unawaited(_peerConnectedSubscription?.cancel()); - unawaited(_peerDisconnectedSubscription?.cancel()); - unawaited(_messageSubscription?.cancel()); - unawaited(_errorSubscription?.cancel()); + final peerService = _peerService; + if (peerService != null) peerService.onReconnected = null; + _observeSubscriptionCancellation(_peerConnectedSubscription?.cancel()); + _observeSubscriptionCancellation(_peerDisconnectedSubscription?.cancel()); + _observeSubscriptionCancellation(_messageSubscription?.cancel()); + _observeSubscriptionCancellation(_errorSubscription?.cancel()); + _observeSubscriptionCancellation(_sessionEndedSubscription?.cancel()); _peerConnectedSubscription = null; _peerDisconnectedSubscription = null; _messageSubscription = null; _errorSubscription = null; + _sessionEndedSubscription = null; - // Cancel host reconnect grace period - _cancelHostReconnectGracePeriod(); + _hostReconnectTimer?.cancel(); + _hostReconnectTimer = null; + _isWaitingForHostReconnect = false; - // Clean up services _controller?.dispose(); _controller = null; - - await _peerService?.disconnect(); - _peerService?.dispose(); _peerService = null; - _session = null; _participants.clear(); _isSyncing = false; @@ -468,8 +533,51 @@ class WatchTogetherProvider with ChangeNotifier { _lastActionEventMs.clear(); _hostIntentionallyLeft = false; - notifyListeners(); - appLogger.d('WatchTogether: Session left'); + if (!_disposed) notifyListeners(); + return peerService; + } + + void _observeSubscriptionCancellation(Future? cancellation) { + if (cancellation == null) return; + unawaited( + cancellation.catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Failed to cancel local subscription', error: error, stackTrace: stackTrace); + }), + ); + } + + Future _finishPeerTeardown(WatchTogetherPeerService peerService, {required bool release}) async { + Object? failure; + StackTrace? failureStackTrace; + if (release) { + try { + await peerService.releaseSession(); + } catch (error, stackTrace) { + failure = error; + failureStackTrace = stackTrace; + appLogger.e('WatchTogether: Failed to release relay ownership', error: error, stackTrace: stackTrace); + } + } + try { + await peerService.disconnect(); + } catch (error, stackTrace) { + failure ??= error; + failureStackTrace ??= stackTrace; + appLogger.e('WatchTogether: Failed to disconnect relay transport', error: error, stackTrace: stackTrace); + } finally { + peerService.dispose(); + } + if (failure != null) { + Error.throwWithStackTrace(failure, failureStackTrace!); + } + } + + void _leaveSessionBestEffort(String reason) { + unawaited( + leaveSession().catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Best-effort $reason failed', error: error, stackTrace: stackTrace); + }), + ); } /// Attach a player to the sync controller for the given media. @@ -517,7 +625,9 @@ class WatchTogetherProvider with ChangeNotifier { /// Set up listeners for peer service events void _setupPeerServiceListeners() { - _peerConnectedSubscription = _peerService!.onPeerConnected.listen((peerId) { + final peerService = _peerService!; + _peerConnectedSubscription = peerService.onPeerConnected.listen((peerId) { + if (_disposed || !identical(_peerService, peerService)) return; appLogger.d('WatchTogether: Peer connected: $peerId'); // If host reconnected during grace period, cancel the timer @@ -526,14 +636,15 @@ class WatchTogetherProvider with ChangeNotifier { } if (!isHost && peerId == _session?.hostPeerId) { - requestCurrentPlaybackSnapshot(); + _requestCurrentPlaybackSnapshotIfHostConnected(peerService); } // Peer will announce themselves with a join message notifyListeners(); }); - _peerDisconnectedSubscription = _peerService!.onPeerDisconnected.listen((peerId) { + _peerDisconnectedSubscription = peerService.onPeerDisconnected.listen((peerId) { + if (_disposed || !identical(_peerService, peerService)) return; appLogger.d('WatchTogether: Peer disconnected: $peerId'); // Capture display name before removal for notification @@ -555,19 +666,45 @@ class WatchTogetherProvider with ChangeNotifier { notifyListeners(); }); - _messageSubscription = _peerService!.onMessageReceived.listen((message) { + _messageSubscription = peerService.onMessageReceived.listen((message) { + if (_disposed || !identical(_peerService, peerService)) return; _handleSyncMessage(message); }); - _errorSubscription = _peerService!.onError.listen((error) { + _errorSubscription = peerService.onError.listen((error) { + if (_disposed || !identical(_peerService, peerService)) return; + final hostPeerId = _session?.hostPeerId; + if (error.serverCode == 'not_in_room' && + !isHost && + hostPeerId != null && + !peerService.connectedPeers.contains(hostPeerId)) { + appLogger.d('WatchTogether: Declared host is not connected yet; keeping the retained-room join pending'); + return; + } appLogger.e('WatchTogether: Peer error: ${error.message}'); - // Update session state on error + // Only established transport stream errors are recoverable. Relay and + // terminal errors must survive any later reconnect callback. if (_session != null && _session!.state == SessionState.connected) { + _recoverableTransportError = error.originalError != null && error.serverCode == null ? error : null; _session = _session!.copyWith(state: SessionState.error, errorMessage: error.message); notifyListeners(); } }); + _sessionEndedSubscription = peerService.onSessionEnded.listen((_) { + if (_disposed || !identical(_peerService, peerService) || isHost) return; + appLogger.d('WatchTogether: Relay confirmed that the host ended the room'); + _hostIntentionallyLeft = true; + _handleHostExitedPlayer(); + final detachedPeerService = _detachLocalSession(announceLeave: false); + if (detachedPeerService != null) { + unawaited( + _finishPeerTeardown(detachedPeerService, release: false).catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Failed to close an ended relay session', error: error, stackTrace: stackTrace); + }), + ); + } + }); } /// Handle incoming sync messages for participant management @@ -622,7 +759,7 @@ class WatchTogetherProvider with ChangeNotifier { if (!isHost && message.peerId == _session?.hostPeerId) { _hostIntentionallyLeft = true; _handleHostExitedPlayer(); - leaveSession(); + _leaveSessionBestEffort('host-leave cleanup'); } notifyListeners(); @@ -772,6 +909,7 @@ class WatchTogetherProvider with ChangeNotifier { if (_isWaitingForHostReconnect) { appLogger.d('WatchTogether: Host reconnect grace period expired'); _isWaitingForHostReconnect = false; + _recoverableTransportError = null; _session = _session?.copyWith(state: SessionState.error, errorMessage: 'Host left the session'); onHostExitedPlayer?.call(); notifyListeners(); @@ -792,11 +930,18 @@ class WatchTogetherProvider with ChangeNotifier { @override void dispose() { + if (_disposed) return; _disposed = true; - _cancelHostReconnectGracePeriod(); - _participantEventController.close(); - leaveSession(); + final peerService = _detachLocalSession(announceLeave: true); + unawaited(_participantEventController.close()); super.dispose(); + if (peerService != null) { + unawaited( + _finishPeerTeardown(peerService, release: true).catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Failed to release session during dispose', error: error, stackTrace: stackTrace); + }), + ); + } } } diff --git a/lib/watch_together/screens/watch_together_screen.dart b/lib/watch_together/screens/watch_together_screen.dart index 50d440ee..cdc05bdb 100644 --- a/lib/watch_together/screens/watch_together_screen.dart +++ b/lib/watch_together/screens/watch_together_screen.dart @@ -26,7 +26,7 @@ import '../../widgets/overlay_sheet.dart'; import '../models/watch_session.dart'; import '../providers/watch_together_provider.dart'; import '../services/recent_rooms_service.dart'; -import '../services/watch_together_peer_service.dart'; +import '../services/watch_together_relay_endpoint.dart'; import '../widgets/join_session_dialog.dart'; import '../../widgets/loading_indicator_box.dart'; @@ -88,14 +88,18 @@ class _NotInSessionViewState extends State<_NotInSessionView> with MountedSetSta bool _isJoining = false; String? _enteringRoomCode; bool? _healthOk; - String? _customRelayUrl; + String? _profileId; + late WatchTogetherRelayEndpoint _relayEndpoint; List _recentRooms = []; @override void initState() { super.initState(); - _customRelayUrl = SettingsService.instanceOrNull?.read(SettingsService.customRelayUrl); - _recentRooms = RecentRoomsService.getRecentRooms(); + _profileId = context.read().activeId; + _relayEndpoint = WatchTogetherRelayEndpoint.resolve( + SettingsService.instanceOrNull?.read(SettingsService.customRelayUrl), + ); + _recentRooms = _loadRecentRooms(); _checkHealth(); } @@ -103,11 +107,17 @@ class _NotInSessionViewState extends State<_NotInSessionView> with MountedSetSta String? get _plexDisplayName => context.read().active?.displayName; + List _loadRecentRooms() { + final profileId = _profileId; + if (profileId == null || profileId.isEmpty) return const []; + return RecentRoomsService.getRecentRooms(profileId: profileId, endpoint: _relayEndpoint); + } + Future _checkHealth() async { final client = HttpClient(); try { client.connectionTimeout = const Duration(seconds: 5); - final request = await client.getUrl(Uri.parse(WatchTogetherPeerService.healthUrlFor(_customRelayUrl))); + final request = await client.getUrl(_relayEndpoint.healthUri); final response = await request.close().namedTimeout( const Duration(seconds: 5), operation: 'WatchTogether health check', @@ -226,10 +236,19 @@ class _NotInSessionViewState extends State<_NotInSessionView> with MountedSetSta try { final sessionId = await widget.watchTogether.createSession( controlMode: controlMode, + relayEndpoint: _relayEndpoint, displayName: _plexDisplayName, ); - await RecentRoomsService.addOrUpdateRoom(sessionId, controlMode: controlMode); - setStateIfMounted(() => _recentRooms = RecentRoomsService.getRecentRooms()); + final profileId = _profileId; + if (profileId != null && profileId.isNotEmpty) { + await RecentRoomsService.addOrUpdateRoom( + sessionId, + profileId: profileId, + endpoint: _relayEndpoint, + controlMode: controlMode, + ); + setStateIfMounted(() => _recentRooms = _loadRecentRooms()); + } } catch (e) { appLogger.e('Failed to create session', error: e); if (mounted) { @@ -271,9 +290,12 @@ class _NotInSessionViewState extends State<_NotInSessionView> with MountedSetSta setState(() => _isJoining = true); try { - await widget.watchTogether.joinSession(sessionId, displayName: _plexDisplayName); - await RecentRoomsService.addOrUpdateRoom(sessionId); - setStateIfMounted(() => _recentRooms = RecentRoomsService.getRecentRooms()); + await widget.watchTogether.joinSession(sessionId, relayEndpoint: _relayEndpoint, displayName: _plexDisplayName); + final profileId = _profileId; + if (profileId != null && profileId.isNotEmpty) { + await RecentRoomsService.addOrUpdateRoom(sessionId, profileId: profileId, endpoint: _relayEndpoint); + setStateIfMounted(() => _recentRooms = _loadRecentRooms()); + } } catch (e) { appLogger.e('Failed to join session', error: e); if (mounted) { @@ -292,11 +314,15 @@ class _NotInSessionViewState extends State<_NotInSessionView> with MountedSetSta try { await widget.watchTogether.enterRoom( room.code, + relayEndpoint: _relayEndpoint, controlMode: room.controlMode ?? ControlMode.anyone, displayName: _plexDisplayName, ); - await RecentRoomsService.addOrUpdateRoom(room.code); - setStateIfMounted(() => _recentRooms = RecentRoomsService.getRecentRooms()); + final profileId = _profileId; + if (profileId != null && profileId.isNotEmpty) { + await RecentRoomsService.addOrUpdateRoom(room.code, profileId: profileId, endpoint: _relayEndpoint); + setStateIfMounted(() => _recentRooms = _loadRecentRooms()); + } } catch (e) { appLogger.e('Failed to enter room', error: e); if (mounted) { @@ -316,13 +342,22 @@ class _NotInSessionViewState extends State<_NotInSessionView> with MountedSetSta ); if (name == null || !mounted) return; - await RecentRoomsService.renameRoom(room.code, name.isEmpty ? null : name); - setStateIfMounted(() => _recentRooms = RecentRoomsService.getRecentRooms()); + final profileId = _profileId; + if (profileId == null || profileId.isEmpty) return; + await RecentRoomsService.renameRoom( + room.code, + name.isEmpty ? null : name, + profileId: profileId, + endpoint: _relayEndpoint, + ); + setStateIfMounted(() => _recentRooms = _loadRecentRooms()); } Future _removeRoom(RecentRoom room) async { - await RecentRoomsService.removeRoom(room.code); - setStateIfMounted(() => _recentRooms = RecentRoomsService.getRecentRooms()); + final profileId = _profileId; + if (profileId == null || profileId.isEmpty) return; + await RecentRoomsService.removeRoom(room.code, profileId: profileId, endpoint: _relayEndpoint); + setStateIfMounted(() => _recentRooms = _loadRecentRooms()); } } @@ -633,8 +668,11 @@ class _ActiveSessionContent extends StatelessWidget { isDestructive: true, ); - if (confirmed) { + if (!confirmed) return; + try { await watchTogether.leaveSession(); + } catch (error, stackTrace) { + appLogger.e('WatchTogether: Session leave failed', error: error, stackTrace: stackTrace); } } } diff --git a/lib/watch_together/services/host_playback_coordinator.dart b/lib/watch_together/services/host_playback_coordinator.dart index 92f45690..e9c905d2 100644 --- a/lib/watch_together/services/host_playback_coordinator.dart +++ b/lib/watch_together/services/host_playback_coordinator.dart @@ -60,6 +60,8 @@ class HostPlaybackCoordinator { static const int seekDebounceMs = 200; static const int implicitJumpThresholdMs = 1500; static const int selfRecoveryMinBufferAheadMs = 2000; + static const double _minimumRemoteRate = 0.25; + static const double _maximumRemoteRate = 4.0; final String myPeerId; final void Function(PlaybackState state, {String? toPeerId}) _sendState; @@ -496,7 +498,9 @@ class HostPlaybackCoordinator { void _applyRemoteSeek(int targetMs, {required String actor}) { final player = _player; - if (player == null) return; + if (player == null || !player.seekable) return; + final durationMs = player.duration.inMilliseconds; + if (durationMs <= 0 || targetMs < 0 || targetMs > durationMs) return; _callbacks.onRemoteAction?.call(actor, PlaybackActionHint.seek); unawaited( player.seek(Duration(milliseconds: targetMs)).then((didSeek) { @@ -516,7 +520,7 @@ class HostPlaybackCoordinator { void _applyRemoteRate(double rate, {required String actor}) { final player = _player; - if (player == null) return; + if (player == null || !rate.isFinite || rate < _minimumRemoteRate || rate > _maximumRemoteRate) return; _callbacks.onRemoteAction?.call(actor, PlaybackActionHint.rate); unawaited( player.setRate(rate).then((didSet) { diff --git a/lib/watch_together/services/recent_rooms_service.dart b/lib/watch_together/services/recent_rooms_service.dart index 0e5291de..fd79e349 100644 --- a/lib/watch_together/services/recent_rooms_service.dart +++ b/lib/watch_together/services/recent_rooms_service.dart @@ -1,22 +1,26 @@ import 'dart:convert'; import 'package:json_annotation/json_annotation.dart'; +import 'package:crypto/crypto.dart'; +import 'package:flutter/foundation.dart'; import '../../services/settings_service.dart'; import '../models/watch_session.dart'; +import 'watch_together_relay_endpoint.dart'; part 'recent_rooms_service.g.dart'; @JsonSerializable(includeIfNull: false) class RecentRoom { final String code; + final String relayScope; final String? name; @JsonKey(fromJson: _dateTimeFromMillis, toJson: _dateTimeToMillis) final DateTime lastUsed; @JsonKey(fromJson: _controlModeFromIndex, toJson: _controlModeToIndex) final ControlMode? controlMode; - const RecentRoom({required this.code, this.name, required this.lastUsed, this.controlMode}); + const RecentRoom({required this.code, required this.relayScope, this.name, required this.lastUsed, this.controlMode}); Map toJson() => _$RecentRoomToJson(this); @@ -24,12 +28,14 @@ class RecentRoom { RecentRoom copyWith({ String? code, + String? relayScope, String? name, DateTime? lastUsed, ControlMode? controlMode, bool clearName = false, }) => RecentRoom( code: code ?? this.code, + relayScope: relayScope ?? this.relayScope, name: clearName ? null : (name ?? this.name), lastUsed: lastUsed ?? this.lastUsed, controlMode: controlMode ?? this.controlMode, @@ -50,12 +56,19 @@ int? _controlModeToIndex(ControlMode? value) => value?.index; class RecentRoomsService { static const int _maxRooms = 20; - static List getRecentRooms() { - final json = SettingsService.instanceOrNull?.read(SettingsService.recentRooms); + static List getRecentRooms({required String profileId, required WatchTogetherRelayEndpoint endpoint}) { + final scope = _scopeFor(endpoint); + return _load(profileId).where((room) => room.relayScope == scope).toList(growable: false); + } + + static List _load(String profileId) { + final settings = SettingsService.instanceOrNull; + if (settings == null) return []; + final json = settings.read(SettingsService.recentRoomsForProfile(profileId)); if (json == null) return []; try { final list = jsonDecode(json) as List; - final rooms = list.map((e) => RecentRoom.fromJson(e as Map)).toList(); + final rooms = list.map((entry) => RecentRoom.fromJson(entry as Map)).toList(); rooms.sort((a, b) => b.lastUsed.compareTo(a.lastUsed)); return rooms; } catch (_) { @@ -63,18 +76,27 @@ class RecentRoomsService { } } - static Future _save(List rooms) async { + static Future _save(String profileId, List rooms) async { rooms.sort((a, b) => b.lastUsed.compareTo(a.lastUsed)); - if (rooms.length > _maxRooms) rooms.removeRange(_maxRooms, rooms.length); + if (rooms.length > _maxRooms) { + rooms.removeRange(_maxRooms, rooms.length); + } await SettingsService.instanceOrNull?.write( - SettingsService.recentRooms, - jsonEncode(rooms.map((r) => r.toJson()).toList()), + SettingsService.recentRoomsForProfile(profileId), + jsonEncode(rooms.map((room) => room.toJson()).toList()), ); } - static Future addOrUpdateRoom(String code, {String? name, ControlMode? controlMode}) async { - final rooms = getRecentRooms(); - final index = rooms.indexWhere((r) => r.code == code); + static Future addOrUpdateRoom( + String code, { + required String profileId, + required WatchTogetherRelayEndpoint endpoint, + String? name, + ControlMode? controlMode, + }) async { + final scope = _scopeFor(endpoint); + final rooms = _load(profileId); + final index = rooms.indexWhere((room) => room.relayScope == scope && room.code == code); if (index >= 0) { rooms[index] = rooms[index].copyWith( lastUsed: DateTime.now(), @@ -82,23 +104,41 @@ class RecentRoomsService { controlMode: controlMode, ); } else { - rooms.add(RecentRoom(code: code, name: name, lastUsed: DateTime.now(), controlMode: controlMode)); + rooms.add( + RecentRoom(code: code, relayScope: scope, name: name, lastUsed: DateTime.now(), controlMode: controlMode), + ); } - await _save(rooms); + await _save(profileId, rooms); } - static Future removeRoom(String code) async { - final rooms = getRecentRooms(); - rooms.removeWhere((r) => r.code == code); - await _save(rooms); + static Future removeRoom( + String code, { + required String profileId, + required WatchTogetherRelayEndpoint endpoint, + }) async { + final scope = _scopeFor(endpoint); + final rooms = _load(profileId)..removeWhere((room) => room.relayScope == scope && room.code == code); + await _save(profileId, rooms); } - static Future renameRoom(String code, String? name) async { - final rooms = getRecentRooms(); - final index = rooms.indexWhere((r) => r.code == code); + static Future renameRoom( + String code, + String? name, { + required String profileId, + required WatchTogetherRelayEndpoint endpoint, + }) async { + final scope = _scopeFor(endpoint); + final rooms = _load(profileId); + final index = rooms.indexWhere((room) => room.relayScope == scope && room.code == code); if (index >= 0) { rooms[index] = rooms[index].copyWith(name: name, clearName: name == null); - await _save(rooms); + await _save(profileId, rooms); } } + + @visibleForTesting + static String relayScopeFor(WatchTogetherRelayEndpoint endpoint) => _scopeFor(endpoint); + + static String _scopeFor(WatchTogetherRelayEndpoint endpoint) => + sha256.convert(utf8.encode(endpoint.canonicalBaseUrl)).toString(); } diff --git a/lib/watch_together/services/recent_rooms_service.g.dart b/lib/watch_together/services/recent_rooms_service.g.dart index 118a5f89..88123dae 100644 --- a/lib/watch_together/services/recent_rooms_service.g.dart +++ b/lib/watch_together/services/recent_rooms_service.g.dart @@ -8,6 +8,7 @@ part of 'recent_rooms_service.dart'; RecentRoom _$RecentRoomFromJson(Map json) => RecentRoom( code: json['code'] as String, + relayScope: json['relayScope'] as String, name: json['name'] as String?, lastUsed: _dateTimeFromMillis((json['lastUsed'] as num).toInt()), controlMode: _controlModeFromIndex((json['controlMode'] as num?)?.toInt()), @@ -16,6 +17,7 @@ RecentRoom _$RecentRoomFromJson(Map json) => RecentRoom( Map _$RecentRoomToJson(RecentRoom instance) => { 'code': instance.code, + 'relayScope': instance.relayScope, 'name': ?instance.name, 'lastUsed': _dateTimeToMillis(instance.lastUsed), 'controlMode': ?_controlModeToIndex(instance.controlMode), diff --git a/lib/watch_together/services/relay_protocol.g.dart b/lib/watch_together/services/relay_protocol.g.dart index c815dc52..c733ec7a 100644 --- a/lib/watch_together/services/relay_protocol.g.dart +++ b/lib/watch_together/services/relay_protocol.g.dart @@ -1,11 +1,16 @@ // Generated by scripts/generate_relay_protocol.py. Do not edit. abstract final class RelayProtocol { + static const int protocolVersion = 2; + static const int legacyProtocolVersion = 0; + static const String create = 'create'; static const String join = 'join'; static const String broadcast = 'broadcast'; static const String sendTo = 'sendTo'; static const String ping = 'ping'; + static const String leave = 'leave'; + static const String endSession = 'endSession'; static const String created = 'created'; static const String joined = 'joined'; static const String peerJoined = 'peerJoined'; @@ -13,6 +18,8 @@ abstract final class RelayProtocol { static const String message = 'message'; static const String error = 'error'; static const String pong = 'pong'; + static const String left = 'left'; + static const String ended = 'ended'; static const String rateLimitedCode = 'rate_limited'; static const String invalidMessageCode = 'invalid_message'; static const String roomExistsCode = 'room_exists'; @@ -20,6 +27,8 @@ abstract final class RelayProtocol { static const String roomFullCode = 'room_full'; static const String notInRoomCode = 'not_in_room'; static const String alreadyInRoomCode = 'already_in_room'; + static const String peerIdUnavailableCode = 'peer_id_unavailable'; + static const String protocolMismatchCode = 'protocol_mismatch'; static const int maxRoomSize = 8; static const int maxMessageSize = 65536; diff --git a/lib/watch_together/services/watch_together_controller.dart b/lib/watch_together/services/watch_together_controller.dart index 0ab963d1..66f538fc 100644 --- a/lib/watch_together/services/watch_together_controller.dart +++ b/lib/watch_together/services/watch_together_controller.dart @@ -19,7 +19,7 @@ import 'watch_together_peer_service.dart'; /// switches or other attach gaps — the player attachment is just an output /// binding the role engine reconciles against. /// -/// Routes the v2 protocol between the relay and the role engine: +/// Routes the v3 protocol between the relay and the role engine: /// host → [HostPlaybackCoordinator] (single writer of [PlaybackState]), /// guest → [GuestPlaybackReconciler] (+ [ClockSync] against the host). class WatchTogetherController { @@ -165,7 +165,11 @@ class WatchTogetherController { _attachedPlayer = null; _coordinator?.detachPlayer(exiting: exiting); _reconciler?.detachPlayer(); - unawaited(attached.dispose()); + unawaited( + attached.dispose().catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Failed to detach player subscriptions', error: error, stackTrace: stackTrace); + }), + ); appLogger.d('WatchTogether: Player detached (exiting: $exiting)'); } @@ -231,7 +235,11 @@ class WatchTogetherController { _disposed = true; detachPlayer(exiting: true); for (final subscription in _subscriptions) { - unawaited(subscription.cancel()); + unawaited( + subscription.cancel().catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Failed to cancel controller subscription', error: error, stackTrace: stackTrace); + }), + ); } _subscriptions.clear(); _clockSync?.stop(); diff --git a/lib/watch_together/services/watch_together_peer_service.dart b/lib/watch_together/services/watch_together_peer_service.dart index 0512da94..bc591b7c 100644 --- a/lib/watch_together/services/watch_together_peer_service.dart +++ b/lib/watch_together/services/watch_together_peer_service.dart @@ -10,12 +10,17 @@ import '../../i18n/strings.g.dart'; import '../../services/base_peer_service.dart'; import '../../utils/app_logger.dart'; import '../models/sync_message.dart'; -import '../primitives.dart'; import 'relay_protocol.g.dart'; +import 'watch_together_relay_endpoint.dart'; // Re-export so existing callers that import from here keep working. export '../../services/base_peer_service.dart' show PeerError, PeerErrorType; +class _PinnedHostChangedError extends PeerError { + const _PinnedHostChangedError() + : super(type: PeerErrorType.serverError, message: 'Relay returned an invalid joined response'); +} + /// Service for managing Watch Together connections via a WebSocket relay /// /// This service handles: @@ -24,32 +29,31 @@ export '../../services/base_peer_service.dart' show PeerError, PeerErrorType; /// - Sending/receiving sync messages through the relay server /// - Reconnection on WebSocket drops class WatchTogetherPeerService with KeepaliveMixin { - static const String defaultBaseUrl = 'https://ice.plezy.app'; + final WatchTogetherRelayEndpoint endpoint; + final WebSocketChannel Function(Uri uri) _channelFactory; - final String _baseUrl; + /// Test synchronization point after relay setup succeeds but before the + /// reconnect is published to consumers. + final Future Function()? debugReconnectSetupSucceededBarrier; - static String get healthUrl => '$defaultBaseUrl/health'; - - static String healthUrlFor(String? customBaseUrl) { - final base = (customBaseUrl != null && customBaseUrl.trim().isNotEmpty) ? customBaseUrl.trim() : defaultBaseUrl; - return '$base/health'; - } - - String get _relayUrl { - final wsBase = _baseUrl.replaceFirst(RegExp(r'^https://'), 'wss://').replaceFirst(RegExp(r'^http://'), 'ws://'); - return '$wsBase/relay'; - } - - WatchTogetherPeerService({String? customBaseUrl}) - : _baseUrl = (customBaseUrl != null && customBaseUrl.trim().isNotEmpty) ? customBaseUrl.trim() : defaultBaseUrl; + WatchTogetherPeerService({ + WatchTogetherRelayEndpoint? endpoint, + this.debugReconnectSetupSucceededBarrier, + WebSocketChannel Function(Uri uri)? debugChannelFactory, + }) : endpoint = endpoint ?? WatchTogetherRelayEndpoint.defaultEndpoint, + _channelFactory = debugChannelFactory ?? ((uri) => WebSocketChannel.connect(uri)); + static const int _relayProtocolVersion = RelayProtocol.protocolVersion; WebSocketChannel? _channel; StreamSubscription? _channelSubscription; Completer? _setupCompleter; + String? _setupRequestType; final Set _connectedPeers = {}; String? _sessionId; String? _myPeerId; bool _isHost = false; + String? _reconnectToken; + String? _hostPeerId; // Stream controllers for events final _peerConnectedController = StreamController.broadcast(); @@ -57,6 +61,7 @@ class WatchTogetherPeerService with KeepaliveMixin { final _messageReceivedController = StreamController.broadcast(); final _errorController = StreamController.broadcast(); final _connectionStateController = StreamController.broadcast(); + final _sessionEndedController = StreamController.broadcast(); // Reconnection state int _reconnectAttempts = 0; @@ -64,6 +69,9 @@ class WatchTogetherPeerService with KeepaliveMixin { Timer? _reconnectTimer; int _connectionEpoch = 0; bool _disposed = false; + bool _initialSetupInProgress = false; + bool _teardownInProgress = false; + Future? _releaseFuture; /// Called after a successful reconnection so the provider can re-announce join. void Function()? onReconnected; @@ -93,12 +101,18 @@ class WatchTogetherPeerService with KeepaliveMixin { /// Stream of connection state changes (true = connected, false = disconnected) Stream get onConnectionStateChanged => _connectionStateController.stream; + /// Emitted when the host has durably ended the relay room. + Stream get onSessionEnded => _sessionEndedController.stream; + /// Current session ID (null if not in a session) String? get sessionId => _sessionId; /// This peer's ID String? get myPeerId => _myPeerId; + /// Relay-declared peer ID whose messages carry host authority. + String? get hostPeerId => _hostPeerId; + /// Whether this peer is the host bool get isHost => _isHost; @@ -115,21 +129,45 @@ class WatchTogetherPeerService with KeepaliveMixin { return String.fromCharCodes(List.generate(5, (_) => chars.codeUnitAt(random.nextInt(chars.length)))); } + /// Mint the reconnect capability before setup so a lost setup ACK can be + /// retried without relying on server-returned state. + static String _mintReconnectToken() { + final random = Random.secure(); + final bytes = List.generate(32, (_) => random.nextInt(256), growable: false); + return base64Url.encode(bytes).replaceAll('=', ''); + } + /// Connect to the relay WebSocket and set up the message listener. /// Returns a completer that completes when the expected response arrives. - Future _connectToRelay() async { - final uri = Uri.parse(_relayUrl); - final channel = WebSocketChannel.connect(uri); + Future _connectToRelay({Duration? timeout, String operation = 'WatchTogether connect'}) async { + final channel = _channelFactory(endpoint.webSocketUri); - // Wait for the connection to be established - await channel.ready; - - return channel; + try { + final ready = channel.ready; + if (timeout == null) { + await ready; + } else { + await ready.namedTimeout(timeout, operation: operation); + } + return channel; + } catch (_) { + try { + unawaited(channel.sink.close()); + } catch (error) { + appLogger.d('WatchTogether: pending channel close ignored', error: error); + } + rethrow; + } } /// Connect, listen, and send a room setup announcement. - Future> _connectAndAnnounce(String type, int epoch) async { - final channel = await _connectToRelay(); + Future> _connectAndAnnounce( + String type, + int epoch, { + Duration? connectTimeout, + String connectOperation = 'WatchTogether connect', + }) async { + final channel = await _connectToRelay(timeout: connectTimeout, operation: connectOperation); if (_disposed || epoch != _connectionEpoch || _sessionId == null) { unawaited(channel.sink.close()); throw StateError('Watch Together connection attempt became stale'); @@ -145,7 +183,15 @@ class WatchTogetherPeerService with KeepaliveMixin { Completer _announce(String type) { final completer = Completer(); _setupCompleter = completer; - _sendRaw({'type': type, 'sessionId': _sessionId, 'peerId': _myPeerId}); + _setupRequestType = type; + final reconnectToken = _reconnectToken; + _sendRaw({ + 'type': type, + 'sessionId': _sessionId, + 'peerId': _myPeerId, + 'reconnectToken': ?reconnectToken, + 'protocolVersion': _relayProtocolVersion, + }); return completer; } @@ -168,6 +214,7 @@ class WatchTogetherPeerService with KeepaliveMixin { if (_setupCompleter case final completer? when !completer.isCompleted) { completer.completeError(error); _setupCompleter = null; + _setupRequestType = null; } _handleWebSocketClosed(); }, @@ -179,12 +226,135 @@ class WatchTogetherPeerService with KeepaliveMixin { const PeerError(type: PeerErrorType.connectionFailed, message: 'WebSocket closed before setup completed'), ); _setupCompleter = null; + _setupRequestType = null; } _handleWebSocketClosed(); }, ); } + static final RegExp _reconnectTokenPattern = RegExp(r'^[A-Za-z0-9_-]{43}$'); + + PeerError _invalidSetupResponse(String type) { + return PeerError(type: PeerErrorType.serverError, message: 'Relay returned an invalid $type response'); + } + + List _acceptSetupResponse(Map msg, String type) { + final responseSessionId = msg['sessionId']; + final hostPeerId = msg['hostPeerId']; + final reconnectToken = msg['reconnectToken']; + final protocolVersion = msg['protocolVersion']; + if (responseSessionId is! String || + responseSessionId != _sessionId || + hostPeerId is! String || + !RelayProtocol.isValidPeerId(hostPeerId) || + (_isHost && hostPeerId != _myPeerId) || + reconnectToken is! String || + reconnectToken != _reconnectToken || + !_reconnectTokenPattern.hasMatch(reconnectToken) || + protocolVersion != _relayProtocolVersion) { + throw _invalidSetupResponse(type); + } + + final establishedHostPeerId = _hostPeerId; + if (!_isHost && _reconnectToken != null && establishedHostPeerId != null && hostPeerId != establishedHostPeerId) { + throw const _PinnedHostChangedError(); + } + + final rawPeers = msg['peers']; + final peers = []; + if (rawPeers != null) { + if (rawPeers is! List) throw _invalidSetupResponse(type); + for (final peerId in rawPeers) { + if (peerId is! String || !RelayProtocol.isValidPeerId(peerId)) { + throw _invalidSetupResponse(type); + } + peers.add(peerId); + } + } + + _hostPeerId = hostPeerId; + _reconnectToken = reconnectToken; + return peers; + } + + void _acceptTeardownResponse(Map msg, String type) { + if (msg['sessionId'] != _sessionId || msg['protocolVersion'] != _relayProtocolVersion) { + throw _invalidSetupResponse(type); + } + if (type == RelayProtocol.left && msg['peerId'] != _myPeerId) { + throw _invalidSetupResponse(type); + } + } + + void _failSetup(PeerError error) { + _safeAdd(_errorController, error); + if (_setupCompleter case final completer? when !completer.isCompleted) { + _setupCompleter = null; + _setupRequestType = null; + completer.completeError(error); + } + } + + void _rejectAdmittedGuestSetup(_PinnedHostChangedError error) { + _safeAdd(_errorController, error); + final rejectedSetup = _setupCompleter; + if (rejectedSetup == null || rejectedSetup.isCompleted) return; + + final leaveCompleter = Completer(); + _setupCompleter = leaveCompleter; + _setupRequestType = RelayProtocol.leave; + _sendRaw({ + 'type': RelayProtocol.leave, + 'sessionId': _sessionId, + 'peerId': _myPeerId, + 'reconnectToken': _reconnectToken, + 'protocolVersion': _relayProtocolVersion, + }); + unawaited(() async { + try { + await leaveCompleter.future.namedTimeout( + const Duration(seconds: 10), + operation: 'WatchTogether rejected reconnect leave', + ); + } catch (releaseError) { + appLogger.d('WatchTogether: rejected reconnect leave ignored', error: releaseError); + } finally { + if (identical(_setupCompleter, leaveCompleter)) { + _setupCompleter = null; + _setupRequestType = null; + } + if (!rejectedSetup.isCompleted) rejectedSetup.completeError(error); + } + }()); + } + + bool _isExhaustedGuestReconnectRoomNotFound(String code) => + code == RelayProtocol.roomNotFoundCode && + !_isHost && + !_initialSetupInProgress && + !_teardownInProgress && + _hostPeerId != null && + _setupRequestType == RelayProtocol.join && + _reconnectAttempts >= _maxReconnectAttempts; + + void _handleGuestSessionEnded() { + _teardownInProgress = true; + ++_connectionEpoch; + _reconnectTimer?.cancel(); + _reconnectTimer = null; + stopKeepalive(); + final error = const PeerError(type: PeerErrorType.invalidSession, message: 'Watch Together session ended'); + if (_setupCompleter case final completer? when !completer.isCompleted) { + _setupCompleter = null; + _setupRequestType = null; + _safeAdd(_errorController, error); + completer.completeError(error); + } + _safeAdd(_sessionEndedController, null); + _handleWebSocketClosed(); + } + /// Handle an incoming server message (JSON string). void _handleServerMessage(String raw) { try { @@ -193,15 +363,40 @@ class WatchTogetherPeerService with KeepaliveMixin { switch (type) { case RelayProtocol.created: - appLogger.d('WatchTogether: Room created: ${msg['sessionId']}'); + late final List peers; + try { + peers = _acceptSetupResponse(msg, RelayProtocol.created); + } on _PinnedHostChangedError catch (error) { + _rejectAdmittedGuestSetup(error); + break; + } on PeerError catch (error) { + _failSetup(error); + break; + } + appLogger.d('WatchTogether: Room created: ${msg['sessionId']} with peers: $peers'); + for (final peerId in peers) { + if (_connectedPeers.add(peerId)) { + _safeAdd(_peerConnectedController, peerId); + } + } _safeAdd(_connectionStateController, true); if (_setupCompleter case final completer? when !completer.isCompleted) { - completer.complete(); _setupCompleter = null; + _setupRequestType = null; + completer.complete(); } case RelayProtocol.joined: - final peers = (msg['peers'] as List?)?.cast() ?? []; + late final List peers; + try { + peers = _acceptSetupResponse(msg, RelayProtocol.joined); + } on _PinnedHostChangedError catch (error) { + _rejectAdmittedGuestSetup(error); + break; + } on PeerError catch (error) { + _failSetup(error); + break; + } appLogger.d('WatchTogether: Joined room ${msg['sessionId']} with peers: $peers'); for (final peerId in peers) { _connectedPeers.add(peerId); @@ -209,8 +404,9 @@ class WatchTogetherPeerService with KeepaliveMixin { } _safeAdd(_connectionStateController, true); if (_setupCompleter case final completer? when !completer.isCompleted) { - completer.complete(); _setupCompleter = null; + _setupRequestType = null; + completer.complete(); } case RelayProtocol.peerJoined: @@ -247,15 +443,55 @@ class WatchTogetherPeerService with KeepaliveMixin { } } + case RelayProtocol.left: + try { + _acceptTeardownResponse(msg, RelayProtocol.left); + } on PeerError catch (error) { + _failSetup(error); + break; + } + if (_setupCompleter case final completer? when !completer.isCompleted) { + _setupCompleter = null; + _setupRequestType = null; + completer.complete(); + } + + case RelayProtocol.ended: + try { + _acceptTeardownResponse(msg, RelayProtocol.ended); + } on PeerError catch (error) { + _failSetup(error); + break; + } + final expectedTeardown = + _setupRequestType == RelayProtocol.endSession || _setupRequestType == RelayProtocol.leave; + if (expectedTeardown) { + if (_setupCompleter case final completer? when !completer.isCompleted) { + _setupCompleter = null; + _setupRequestType = null; + completer.complete(); + } + } else if (!_isHost) { + _handleGuestSessionEnded(); + } else { + _failSetup(_invalidSetupResponse(RelayProtocol.ended)); + } + case RelayProtocol.error: final code = msg['code'] as String? ?? 'unknown'; final message = msg['message'] as String? ?? t.common.unknown; + if (_isExhaustedGuestReconnectRoomNotFound(code)) { + appLogger.d('WatchTogether: Room gone after reconnect retries; guest session ended'); + _handleGuestSessionEnded(); + break; + } appLogger.e('WatchTogether: Server error: $code - $message'); final error = PeerError(type: PeerErrorType.serverError, message: '$code: $message', serverCode: code); _safeAdd(_errorController, error); if (_setupCompleter case final completer? when !completer.isCompleted) { - completer.completeError(error); _setupCompleter = null; + _setupRequestType = null; + completer.completeError(error); } case RelayProtocol.pong: @@ -265,8 +501,13 @@ class WatchTogetherPeerService with KeepaliveMixin { default: appLogger.w('WatchTogether: Unknown server message type: $type'); } - } catch (e) { - appLogger.e('WatchTogether: Failed to parse server message', error: e); + } catch (_) { + appLogger.e('WatchTogether: Failed to parse server message'); + if (_setupCompleter case final completer? when !completer.isCompleted) { + _failSetup( + const PeerError(type: PeerErrorType.serverError, message: 'Relay returned an invalid setup response'), + ); + } } } @@ -295,7 +536,8 @@ class WatchTogetherPeerService with KeepaliveMixin { /// Handle the WebSocket being closed unexpectedly — attempt reconnection. void _handleWebSocketClosed() { final channel = _channel; - ++_connectionEpoch; + final shouldReconnect = !_initialSetupInProgress && !_teardownInProgress; + if (shouldReconnect) ++_connectionEpoch; stopKeepalive(); unawaited(_channelSubscription?.cancel()); _channelSubscription = null; @@ -308,7 +550,7 @@ class WatchTogetherPeerService with KeepaliveMixin { _connectedPeers.clear(); _safeAdd(_connectionStateController, false); - if (!_disposed && _sessionId != null) { + if (shouldReconnect && !_disposed && _sessionId != null) { _attemptReconnect(_connectionEpoch); } } @@ -355,6 +597,8 @@ class WatchTogetherPeerService with KeepaliveMixin { } } + await debugReconnectSetupSucceededBarrier?.call(); + if (_disposed || epoch != _connectionEpoch) return; _reconnectAttempts = 0; appLogger.d('WatchTogether: Reconnected successfully'); @@ -371,12 +615,67 @@ class WatchTogetherPeerService with KeepaliveMixin { }); } + bool _isRetryableInitialSetupError(Object error) => + error is TimeoutException || + (error is PeerError && + (error.type == PeerErrorType.connectionFailed || + error.type == PeerErrorType.networkError || + error.type == PeerErrorType.timeout)) || + error is! PeerError; + + Future _resetTransportForInitialRetry() async { + stopKeepalive(); + final subscription = _channelSubscription; + final channel = _channel; + _channelSubscription = null; + _channel = null; + _setupCompleter = null; + _setupRequestType = null; + await subscription?.cancel(); + try { + await channel?.sink.close(); + } catch (error) { + appLogger.d('WatchTogether: setup retry close ignored', error: error); + } + } + + Future _performInitialSetup(String type, int epoch, PeerError timeoutError) async { + _initialSetupInProgress = true; + try { + for (var attempt = 0; attempt < _maxReconnectAttempts; attempt++) { + try { + final completer = await _connectAndAnnounce(type, epoch); + await completer.future.timeout(const Duration(seconds: 10), onTimeout: () => throw timeoutError); + return; + } catch (error) { + if (_disposed || epoch != _connectionEpoch) rethrow; + await _resetTransportForInitialRetry(); + if (!_isRetryableInitialSetupError(error) || attempt + 1 >= _maxReconnectAttempts) { + rethrow; + } + await Future.delayed(Duration(milliseconds: 250 * (attempt + 1))); + } + } + } finally { + _initialSetupInProgress = false; + } + } + + Future _bestEffortReleaseFailedSetup(Object setupError) async { + if (!_isRetryableInitialSetupError(setupError)) return; + try { + await releaseSession(); + } catch (releaseError) { + appLogger.d('WatchTogether: failed setup reservation release ignored', error: releaseError); + } + } + /// Create a new session as host /// /// Returns the session ID that others can use to join. /// If [sessionId] is provided, uses that instead of generating a new one. Future createSession({String? sessionId}) async { - if (_channel != null) { + if (_sessionId != null || _channel != null) { await disconnect(); } @@ -390,24 +689,23 @@ class WatchTogetherPeerService with KeepaliveMixin { } _isHost = true; _sessionId = resolvedSessionId; - _myPeerId = watchTogetherHostPeerId(resolvedSessionId); + _myPeerId = const Uuid().v4(); + _reconnectToken = _mintReconnectToken(); _reconnectAttempts = 0; final epoch = ++_connectionEpoch; try { - final completer = await _connectAndAnnounce(RelayProtocol.create, epoch); - - await completer.future.timeout( - const Duration(seconds: 10), - onTimeout: () { - throw const PeerError(type: PeerErrorType.timeout, message: 'Timed out creating session'); - }, + await _performInitialSetup( + RelayProtocol.create, + epoch, + const PeerError(type: PeerErrorType.timeout, message: 'Timed out creating session'), ); appLogger.d('WatchTogether: Session created: $_sessionId'); return _sessionId!; } catch (e) { appLogger.e('WatchTogether: Failed to create session', error: e); + await _bestEffortReleaseFailedSetup(e); await disconnect(); rethrow; } @@ -415,7 +713,7 @@ class WatchTogetherPeerService with KeepaliveMixin { /// Join an existing session as guest. Future joinSession(String sessionId) async { - if (_channel != null) { + if (_sessionId != null || _channel != null) { await disconnect(); } @@ -430,22 +728,21 @@ class WatchTogetherPeerService with KeepaliveMixin { _isHost = false; _sessionId = resolvedSessionId; _myPeerId = const Uuid().v4(); + _reconnectToken = _mintReconnectToken(); _reconnectAttempts = 0; final epoch = ++_connectionEpoch; try { - final completer = await _connectAndAnnounce(RelayProtocol.join, epoch); - - await completer.future.timeout( - const Duration(seconds: 10), - onTimeout: () { - throw PeerError(type: PeerErrorType.timeout, message: t.watchTogether.failedToJoin); - }, + await _performInitialSetup( + RelayProtocol.join, + epoch, + PeerError(type: PeerErrorType.timeout, message: t.watchTogether.failedToJoin), ); appLogger.d('WatchTogether: Joined session: $_sessionId'); } catch (e) { appLogger.e('WatchTogether: Failed to join session', error: e); + await _bestEffortReleaseFailedSetup(e); await disconnect(); rethrow; } @@ -466,7 +763,82 @@ class WatchTogetherPeerService with KeepaliveMixin { _sendRaw({'type': RelayProtocol.sendTo, 'to': peerId, 'payload': payload}); } - /// Disconnect from all peers and close the session + /// Explicitly release this peer's relay ownership. Hosts destroy the room; + /// 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() async { + if (_sessionId == null || _myPeerId == null || _reconnectToken == null) return; + + _teardownInProgress = true; + _reconnectTimer?.cancel(); + _reconnectTimer = null; + final epoch = ++_connectionEpoch; + try { + if (_setupCompleter case final completer? when !completer.isCompleted) { + await _resetTransportForInitialRetry(); + } + for (var attempt = 0; attempt < _maxReconnectAttempts; attempt++) { + try { + if (_channel == null) { + final reconnectCompleter = await _connectAndAnnounce( + RelayProtocol.join, + epoch, + connectTimeout: const Duration(seconds: 10), + connectOperation: 'WatchTogether release reconnect', + ); + await reconnectCompleter.future.namedTimeout( + const Duration(seconds: 10), + operation: 'WatchTogether release reconnect', + ); + } + + final releaseCompleter = Completer(); + _setupCompleter = releaseCompleter; + _setupRequestType = _isHost ? RelayProtocol.endSession : RelayProtocol.leave; + _sendRaw({ + 'type': _isHost ? RelayProtocol.endSession : RelayProtocol.leave, + 'sessionId': _sessionId, + 'peerId': _myPeerId, + 'reconnectToken': _reconnectToken, + 'protocolVersion': _relayProtocolVersion, + }); + await releaseCompleter.future.namedTimeout( + const Duration(seconds: 10), + operation: _isHost ? 'WatchTogether end session' : 'WatchTogether leave session', + ); + return; + } catch (error) { + if (error is PeerError && + (error.serverCode == RelayProtocol.roomNotFoundCode || + error.serverCode == RelayProtocol.notInRoomCode || + (!_isHost && error.serverCode == RelayProtocol.peerIdUnavailableCode))) { + return; + } + await _resetTransportForInitialRetry(); + if (!_isRetryableInitialSetupError(error) || attempt + 1 >= _maxReconnectAttempts) { + rethrow; + } + await Future.delayed(Duration(milliseconds: 250 * (attempt + 1))); + } + } + } finally { + _teardownInProgress = false; + } + } + + /// Close and forget local relay state without sending a release. Established + /// intentional exits call [releaseSession] before this cleanup step. Future disconnect() async { appLogger.d('WatchTogether: Disconnecting...'); ++_connectionEpoch; @@ -480,14 +852,18 @@ class WatchTogetherPeerService with KeepaliveMixin { _channel = null; final setupCompleter = _setupCompleter; _setupCompleter = null; + _setupRequestType = null; if (setupCompleter != null && !setupCompleter.isCompleted) { setupCompleter.completeError(StateError('Watch Together connection cancelled')); } _connectedPeers.clear(); _sessionId = null; _myPeerId = null; + _reconnectToken = null; + _hostPeerId = null; _isHost = false; _reconnectAttempts = 0; + _teardownInProgress = false; unawaited(subscription?.cancel()); try { @@ -509,5 +885,6 @@ class WatchTogetherPeerService with KeepaliveMixin { _messageReceivedController.close(); _errorController.close(); _connectionStateController.close(); + _sessionEndedController.close(); } } diff --git a/lib/watch_together/services/watch_together_relay_endpoint.dart b/lib/watch_together/services/watch_together_relay_endpoint.dart new file mode 100644 index 00000000..b49b2c78 --- /dev/null +++ b/lib/watch_together/services/watch_together_relay_endpoint.dart @@ -0,0 +1,85 @@ +/// Canonical HTTP(S) base endpoint for a Watch Together relay. +/// +/// A base may include a reverse-proxy path prefix. The concrete health and +/// WebSocket routes are always appended as path segments. +final class WatchTogetherRelayEndpoint { + WatchTogetherRelayEndpoint._(this._baseUri); + + static const String defaultBaseUrl = 'https://ice.plezy.app'; + + static final WatchTogetherRelayEndpoint defaultEndpoint = WatchTogetherRelayEndpoint._(Uri.parse(defaultBaseUrl)); + + final Uri _baseUri; + + String get canonicalBaseUrl => _baseUri.toString(); + + Uri get healthUri => _appendPathSegment('health'); + + Uri get webSocketUri => _appendPathSegment('relay').replace(scheme: _baseUri.scheme == 'https' ? 'wss' : 'ws'); + + static WatchTogetherRelayEndpoint resolve(String? value) { + if (value == null || value.trim().isEmpty) return defaultEndpoint; + final endpoint = tryParseCustom(value); + if (endpoint == null) { + throw FormatException('Invalid Watch Together relay base URL'); + } + return endpoint; + } + + static WatchTogetherRelayEndpoint? tryParseCustom(String value) { + final trimmed = value.trim(); + if (trimmed.isEmpty) return null; + + final Uri uri; + try { + uri = Uri.parse(trimmed); + if (uri.hasPort && (uri.port < 1 || uri.port > 65535)) { + return null; + } + } on FormatException { + return null; + } + + if ((uri.scheme != 'http' && uri.scheme != 'https') || + !uri.isAbsolute || + !uri.hasAuthority || + uri.host.isEmpty || + uri.userInfo.isNotEmpty || + uri.hasQuery || + uri.hasFragment) { + return null; + } + + final normalized = uri.normalizePath(); + var path = normalized.path; + while (path.endsWith('/')) { + path = path.substring(0, path.length - 1); + } + final hasDefaultPort = + normalized.hasPort && + ((normalized.scheme == 'http' && normalized.port == 80) || + (normalized.scheme == 'https' && normalized.port == 443)); + final canonical = Uri( + scheme: normalized.scheme, + host: normalized.host, + port: normalized.hasPort && !hasDefaultPort ? normalized.port : null, + path: path, + ); + return WatchTogetherRelayEndpoint._(canonical); + } + + Uri _appendPathSegment(String segment) { + final prefix = _baseUri.path; + return _baseUri.replace(path: '$prefix/$segment'); + } + + @override + bool operator ==(Object other) => + identical(this, other) || other is WatchTogetherRelayEndpoint && other.canonicalBaseUrl == canonicalBaseUrl; + + @override + int get hashCode => canonicalBaseUrl.hashCode; + + @override + String toString() => canonicalBaseUrl; +} diff --git a/lib/watch_together/widgets/watch_together_overlay.dart b/lib/watch_together/widgets/watch_together_overlay.dart index b8968f7a..4356786c 100644 --- a/lib/watch_together/widgets/watch_together_overlay.dart +++ b/lib/watch_together/widgets/watch_together_overlay.dart @@ -273,7 +273,11 @@ class _SessionMenuSheet extends StatelessWidget { if (context.mounted) { OverlaySheetController.closeAdaptive(context); } - unawaited(provider.leaveSession()); + unawaited( + provider.leaveSession().catchError((Object error, StackTrace stackTrace) { + appLogger.e('WatchTogether: Overlay leave failed', error: error, stackTrace: stackTrace); + }), + ); onLeaveSession?.call(); } } diff --git a/relay_protocol.json b/relay_protocol.json index 8a9a8b8c..47a79c2e 100644 --- a/relay_protocol.json +++ b/relay_protocol.json @@ -1,10 +1,14 @@ { + "protocolVersion": 2, + "legacyProtocolVersion": 0, "clientMessageTypes": { "create": "create", "join": "join", "broadcast": "broadcast", "sendTo": "sendTo", - "ping": "ping" + "ping": "ping", + "leave": "leave", + "endSession": "endSession" }, "serverMessageTypes": { "created": "created", @@ -13,7 +17,9 @@ "peerLeft": "peerLeft", "message": "message", "error": "error", - "pong": "pong" + "pong": "pong", + "left": "left", + "ended": "ended" }, "errorCodes": { "rateLimited": "rate_limited", @@ -22,7 +28,9 @@ "roomNotFound": "room_not_found", "roomFull": "room_full", "notInRoom": "not_in_room", - "alreadyInRoom": "already_in_room" + "alreadyInRoom": "already_in_room", + "peerIdUnavailable": "peer_id_unavailable", + "protocolMismatch": "protocol_mismatch" }, "limits": { "maxRoomSize": 8, diff --git a/scripts/ci_checks.sh b/scripts/ci_checks.sh index 2057f949..d02a8596 100755 --- a/scripts/ci_checks.sh +++ b/scripts/ci_checks.sh @@ -90,6 +90,7 @@ if python3 scripts/check_build_workflow.py && python3 scripts/check_workflow_action_pins.py && python3 scripts/test_check_workflow_action_pins.py && python3 scripts/test_check_codegen.py && + python3 scripts/test_generate_relay_protocol.py && python3 scripts/test_format_native.py && python3 scripts/check_update_packages_workflow.py && python3 scripts/test_pubspec_version.py && diff --git a/scripts/generate_relay_protocol.py b/scripts/generate_relay_protocol.py index f08845c5..0876c243 100755 --- a/scripts/generate_relay_protocol.py +++ b/scripts/generate_relay_protocol.py @@ -11,17 +11,41 @@ SPEC_PATH = ROOT / "relay_protocol.json" DART_PATH = ROOT / "lib/watch_together/services/relay_protocol.g.dart" GO_PATH = ROOT / "server/relay_protocol_gen.go" +SUPPORTED_ID_PATTERN = r"^[A-Za-z0-9_-]+$" + def camel_to_pascal(value: str) -> str: return value[:1].upper() + value[1:] +def validated_id_pattern(spec: dict) -> str: + try: + pattern = spec["idPattern"] + except KeyError: + raise ValueError("idPattern is required") from None + if not isinstance(pattern, str): + raise ValueError("idPattern must be a string") + if pattern != SUPPORTED_ID_PATTERN: + raise ValueError( + f"unsupported idPattern {pattern!r}; expected {SUPPORTED_ID_PATTERN!r}" + ) + return pattern + + def dart_source(spec: dict) -> str: + id_pattern = validated_id_pattern(spec) lines = [ "// Generated by scripts/generate_relay_protocol.py. Do not edit.", "", "abstract final class RelayProtocol {", ] + lines.extend( + [ + f" static const int protocolVersion = {spec['protocolVersion']};", + f" static const int legacyProtocolVersion = {spec['legacyProtocolVersion']};", + "", + ] + ) for group in ("clientMessageTypes", "serverMessageTypes"): for name, value in spec[group].items(): lines.append(f" static const String {name} = {value!r};") @@ -33,7 +57,7 @@ def dart_source(spec: dict) -> str: lines.extend( [ "", - " static final RegExp _idPattern = RegExp(r'^[A-Za-z0-9_-]+$');", + f" static final RegExp _idPattern = RegExp(r{id_pattern!r});", "", " static bool isValidSessionId(String value) =>", " value.isNotEmpty && value.length <= maxSessionIdLength && _idPattern.hasMatch(value);", @@ -48,6 +72,7 @@ def dart_source(spec: dict) -> str: def go_source(spec: dict) -> str: + validated_id_pattern(spec) lines = [ "// Code generated by scripts/generate_relay_protocol.py. DO NOT EDIT.", "", @@ -55,6 +80,13 @@ def go_source(spec: dict) -> str: "", "const (", ] + lines.extend( + [ + f"\trelayProtocolVersion = {spec['protocolVersion']}", + f"\tlegacyRelayProtocolVersion = {spec['legacyProtocolVersion']}", + "", + ] + ) protocol_constants = [] for group in ("clientMessageTypes", "serverMessageTypes"): protocol_constants.extend( @@ -109,8 +141,10 @@ def go_source(spec: dict) -> str: def main() -> None: spec = json.loads(SPEC_PATH.read_text(encoding="utf-8")) - DART_PATH.write_text(dart_source(spec), encoding="utf-8") - GO_PATH.write_text(go_source(spec), encoding="utf-8") + dart_output = dart_source(spec) + go_output = go_source(spec) + DART_PATH.write_text(dart_output, encoding="utf-8") + GO_PATH.write_text(go_output, encoding="utf-8") if __name__ == "__main__": diff --git a/scripts/test_generate_relay_protocol.py b/scripts/test_generate_relay_protocol.py new file mode 100644 index 00000000..92f2db8c --- /dev/null +++ b/scripts/test_generate_relay_protocol.py @@ -0,0 +1,72 @@ +import copy +import json +import tempfile +import unittest +from pathlib import Path +from unittest import mock + +import generate_relay_protocol as generator + + +class RelayProtocolGeneratorTest(unittest.TestCase): + def setUp(self) -> None: + self.spec = json.loads(generator.SPEC_PATH.read_text(encoding="utf-8")) + + def test_supported_pattern_renders_both_targets(self) -> None: + dart_output = generator.dart_source(copy.deepcopy(self.spec)) + go_output = generator.go_source(copy.deepcopy(self.spec)) + + self.assertIn( + f"RegExp(r{generator.SUPPORTED_ID_PATTERN!r})", + dart_output, + ) + self.assertIn("func validRelayID(value string, maxLength int) bool", go_output) + + def test_changed_pattern_fails_before_writing_either_target(self) -> None: + changed_spec = copy.deepcopy(self.spec) + changed_spec["idPattern"] = r"^[A-Za-z0-9_.-]+$" + + for renderer in (generator.dart_source, generator.go_source): + with self.subTest(renderer=renderer.__name__): + with self.assertRaisesRegex(ValueError, "idPattern"): + renderer(changed_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(changed_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, "idPattern"): + 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_missing_pattern_is_rejected(self) -> None: + spec = copy.deepcopy(self.spec) + del spec["idPattern"] + + with self.assertRaisesRegex(ValueError, "idPattern is required"): + generator.validated_id_pattern(spec) + + def test_non_string_pattern_is_rejected(self) -> None: + for value in (None, 42, [generator.SUPPORTED_ID_PATTERN]): + with self.subTest(value=value): + spec = copy.deepcopy(self.spec) + spec["idPattern"] = value + + with self.assertRaisesRegex(ValueError, "idPattern must be a string"): + generator.validated_id_pattern(spec) + + +if __name__ == "__main__": + unittest.main() diff --git a/server/client_ip.go b/server/client_ip.go new file mode 100644 index 00000000..cb6e732a --- /dev/null +++ b/server/client_ip.go @@ -0,0 +1,149 @@ +package main + +import ( + "errors" + "net" + "net/http" + "net/netip" + "strings" +) + +var errInvalidClientAddress = errors.New("invalid client address") + +const ( + maxForwardedForBytes = 4 * 1024 + maxForwardedForHops = 32 +) + +type clientIPResolver struct { + trustedProxies []netip.Prefix +} + +func newClientIPResolver(trustedProxies []netip.Prefix) clientIPResolver { + return clientIPResolver{trustedProxies: append([]netip.Prefix(nil), trustedProxies...)} +} + +func parseTrustedProxyCIDRs(value string) ([]netip.Prefix, error) { + if strings.TrimSpace(value) == "" { + return nil, nil + } + + parts := strings.Split(value, ",") + prefixes := make([]netip.Prefix, 0, len(parts)) + for _, part := range parts { + part = strings.TrimSpace(part) + if part == "" { + return nil, errInvalidClientAddress + } + prefix, err := netip.ParsePrefix(part) + if err != nil { + return nil, errInvalidClientAddress + } + if prefix.Addr().Is4In6() { + if prefix.Bits() < 96 { + return nil, errInvalidClientAddress + } + prefix = netip.PrefixFrom(prefix.Addr().Unmap(), prefix.Bits()-96) + } + prefixes = append(prefixes, prefix.Masked()) + } + return prefixes, nil +} + +func (r clientIPResolver) resolve(req *http.Request) (string, error) { + peer, err := parseRemoteAddress(req.RemoteAddr) + if err != nil { + return "", errInvalidClientAddress + } + peer = peer.Unmap() + + if !r.trusted(peer) { + return normalizeClientAddress(peer), nil + } + + values := req.Header.Values("X-Forwarded-For") + totalBytes := 0 + hopCount := 0 + for _, value := range values { + totalBytes += len(value) + if totalBytes > maxForwardedForBytes { + return "", errInvalidClientAddress + } + hopCount += strings.Count(value, ",") + 1 + if hopCount > maxForwardedForHops { + return "", errInvalidClientAddress + } + } + if len(values) == 0 { + return normalizeClientAddress(peer), nil + } + + selected := peer + useForwardedHop := true + parsedHops := 0 + for valueIndex := len(values) - 1; valueIndex >= 0; valueIndex-- { + value := values[valueIndex] + end := len(value) + for { + separator := strings.LastIndexByte(value[:end], ',') + element := strings.TrimSpace(value[separator+1 : end]) + if element == "" { + return "", errInvalidClientAddress + } + addr, parseErr := netip.ParseAddr(element) + if parseErr != nil || addr.Zone() != "" { + return "", errInvalidClientAddress + } + parsedHops++ + if useForwardedHop { + if r.trusted(selected) { + selected = addr.Unmap() + } else { + useForwardedHop = false + } + } + if separator < 0 { + break + } + end = separator + } + } + if parsedHops == 0 { + return normalizeClientAddress(peer), nil + } + return normalizeClientAddress(selected), nil +} + +func (r clientIPResolver) trusted(addr netip.Addr) bool { + addr = addr.Unmap() + for _, prefix := range r.trustedProxies { + if prefix.Contains(addr) { + return true + } + } + return false +} + +func parseRemoteAddress(remote string) (netip.Addr, error) { + host, _, err := net.SplitHostPort(remote) + if err == nil { + addr, parseErr := netip.ParseAddr(host) + if parseErr != nil || addr.Zone() != "" { + return netip.Addr{}, errInvalidClientAddress + } + return addr, nil + } + addr, parseErr := netip.ParseAddr(remote) + if parseErr != nil || addr.Zone() != "" { + return netip.Addr{}, errInvalidClientAddress + } + return addr, nil +} + +func normalizeClientAddress(addr netip.Addr) string { + addr = addr.Unmap() + if addr.Is6() { + return netip.PrefixFrom(addr, 64).Masked().Addr().String() + } + return addr.String() +} diff --git a/server/docker-compose.yml b/server/docker-compose.yml index 5babe160..6b53cb17 100644 --- a/server/docker-compose.yml +++ b/server/docker-compose.yml @@ -10,6 +10,7 @@ services: - "127.0.0.1:8080:8080" environment: OAUTH_BASE_URL: https://ice.plezy.app + TRUSTED_PROXY_CIDRS: ${TRUSTED_PROXY_CIDRS:-} MAL_CLIENT_ID: ${MAL_CLIENT_ID:-} ANILIST_CLIENT_ID: ${ANILIST_CLIENT_ID:-} ANILIST_CLIENT_SECRET: ${ANILIST_CLIENT_SECRET:-} diff --git a/server/main.go b/server/main.go index a0c78d96..8ad3c545 100644 --- a/server/main.go +++ b/server/main.go @@ -3,9 +3,13 @@ package main import ( "context" "crypto/rand" + "crypto/sha256" + "crypto/subtle" + "encoding/base64" "encoding/json" "errors" "flag" + "fmt" "io" "io/fs" "log" @@ -25,33 +29,47 @@ import ( ) const ( - rateBurst = 30 - rateSustained = 10 - cleanupInterval = 5 * time.Minute - emptyRoomMaxAge = 5 * time.Minute - roomMaxAge = 24 * time.Hour - writeWait = 10 * time.Second - pongWait = 60 * time.Second - pingInterval = 30 * time.Second - maxLogSize = 1 * 1024 * 1024 // 1MB - logMaxAge = 3 * 24 * time.Hour - logIDLength = 5 - logRateInterval = 1 * time.Minute - maxLogEntries = 500 - maxPosterSize = 5 * 1024 * 1024 // 5MB - maxPosterStoreSize = int64(1 * 1024 * 1024 * 1024) - posterMaxAge = 3 * time.Hour - posterIDLength = 16 - maxConnsPerIP = 5 - maxGlobalConns = 100 - maxRoomsPerIP = 3 - connRateBurst = 5 - connRateSustained = 1 - - snapshotFormatVersion = 1 - snapshotDebounce = 100 * time.Millisecond - snapshotFlushTimeout = 5 * time.Second - snapshotMaxFileSize = 1 * 1024 * 1024 + rateBurst = 30 + rateSustained = 10 + cleanupInterval = 5 * time.Minute + emptyRoomMaxAge = 5 * time.Minute + roomMaxAge = 24 * time.Hour + writeWait = 10 * time.Second + httpResponseWriteMargin = 10 * time.Second + httpResponseWriteTimeout = oauthResultWait + httpResponseWriteMargin + pongWait = 60 * time.Second + pingInterval = 30 * time.Second + maxLogSize = 1 * 1024 * 1024 // 1MB + logMaxAge = 3 * 24 * time.Hour + logIDLength = 25 + logRateInterval = 1 * time.Minute + logLookupRateBurst = 10 + logLookupRateSustained = 1 + maxLogEntries = 500 + maxFailedLogLookupSources = 4096 + maxConcurrentLogLookups = 32 + maxHTTPHeaderBytes = 64 * 1024 + maxPosterSize = 5 * 1024 * 1024 // 5MB + maxPosterStoreSize = int64(1 * 1024 * 1024 * 1024) + posterMaxAge = 3 * time.Hour + posterIDLength = 16 + posterPerIPRateBurst = 3 + posterPerIPRateSustained = 1 + posterGlobalRateBurst = 8 + posterGlobalRateSustained = 2 + maxConcurrentPosterUploads = 4 + posterUploadReadTimeout = 30 * time.Second + maxConnsPerIP = 5 + maxGlobalConns = 100 + maxRoomsPerIP = 3 + maxRetainedRooms = 2000 + connRateBurst = 5 + connRateSustained = 1 + reconnectTokenSize = 32 + snapshotFormatVersion = 3 + snapshotDebounce = 100 * time.Millisecond + snapshotFlushTimeout = 5 * time.Second + snapshotMaxFileSize = 4 * 1024 * 1024 ) var upgrader = websocket.Upgrader{ @@ -63,52 +81,80 @@ var upgrader = websocket.Upgrader{ // --- Messages --- type clientMsg struct { - Type string `json:"type"` - SessionID string `json:"sessionId,omitempty"` - PeerID string `json:"peerId,omitempty"` - To string `json:"to,omitempty"` - Payload json.RawMessage `json:"payload,omitempty"` + Type string `json:"type"` + SessionID string `json:"sessionId,omitempty"` + PeerID string `json:"peerId,omitempty"` + ReconnectToken string `json:"reconnectToken,omitempty"` + ProtocolVersion int `json:"protocolVersion,omitempty"` + To string `json:"to,omitempty"` + Payload json.RawMessage `json:"payload,omitempty"` } type serverMsg struct { - Type string `json:"type"` - SessionID string `json:"sessionId,omitempty"` - PeerID string `json:"peerId,omitempty"` - From string `json:"from,omitempty"` - Peers []string `json:"peers,omitempty"` - Code string `json:"code,omitempty"` - Message string `json:"message,omitempty"` - Payload json.RawMessage `json:"payload,omitempty"` + Type string `json:"type"` + SessionID string `json:"sessionId,omitempty"` + PeerID string `json:"peerId,omitempty"` + HostPeerID string `json:"hostPeerId,omitempty"` + ReconnectToken string `json:"reconnectToken,omitempty"` + ProtocolVersion int `json:"protocolVersion,omitempty"` + From string `json:"from,omitempty"` + Peers []string `json:"peers,omitempty"` + Code string `json:"code,omitempty"` + Message string `json:"message,omitempty"` + Payload json.RawMessage `json:"payload,omitempty"` } // --- Client (serializes writes to a single goroutine) --- +type outboundFrame struct { + data []byte + written chan bool +} + type Client struct { conn *websocket.Conn - send chan []byte + send chan outboundFrame done chan struct{} closeOnce sync.Once } func newClient(conn *websocket.Conn) *Client { - c := &Client{conn: conn, send: make(chan []byte, 64), done: make(chan struct{})} + c := &Client{conn: conn, send: make(chan outboundFrame, 64), done: make(chan struct{})} go c.writePump() return c } func (c *Client) writePump() { ticker := time.NewTicker(pingInterval) - defer ticker.Stop() + defer func() { + ticker.Stop() + c.close() + }() for { select { - case data := <-c.send: - c.conn.SetWriteDeadline(time.Now().Add(writeWait)) - c.conn.WriteMessage(websocket.TextMessage, data) + case frame := <-c.send: + if err := c.conn.SetWriteDeadline(time.Now().Add(writeWait)); err != nil { + if frame.written != nil { + frame.written <- false + } + return + } + if err := c.conn.WriteMessage(websocket.TextMessage, frame.data); err != nil { + if frame.written != nil { + frame.written <- false + } + return + } + if frame.written != nil { + frame.written <- true + } case <-c.done: - c.conn.WriteMessage(websocket.CloseMessage, nil) + _ = c.conn.WriteMessage(websocket.CloseMessage, nil) return case <-ticker.C: - c.conn.SetWriteDeadline(time.Now().Add(writeWait)) + if err := c.conn.SetWriteDeadline(time.Now().Add(writeWait)); err != nil { + return + } if err := c.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } @@ -116,20 +162,51 @@ func (c *Client) writePump() { } } -func (c *Client) trySend(data []byte) { +func (c *Client) enqueueFrame(frame outboundFrame) bool { select { - case c.send <- data: case <-c.done: + return false default: } + + select { + case <-c.done: + return false + case c.send <- frame: + return true + default: + c.close() + return false + } } -func (c *Client) sendJSON(msg serverMsg) { +func (c *Client) enqueue(data []byte) bool { + return c.enqueueFrame(outboundFrame{data: data}) +} + +func (c *Client) sendJSON(msg serverMsg) bool { data, err := json.Marshal(msg) if err != nil { - return + return false + } + return c.enqueue(data) +} + +func (c *Client) sendJSONAndWait(msg serverMsg) bool { + data, err := json.Marshal(msg) + if err != nil { + return false + } + written := make(chan bool, 1) + if !c.enqueueFrame(outboundFrame{data: data, written: written}) { + return false + } + select { + case ok := <-written: + return ok + case <-time.After(writeWait): + return false } - c.trySend(data) } func (c *Client) close() { @@ -141,22 +218,32 @@ func (c *Client) close() { // --- Room --- +type reconnectVerifier [sha256.Size]byte + type Room struct { - SessionID string - HostPeerID string - Peers map[string]*Client `json:"-"` - mu sync.RWMutex `json:"-"` - CreatedAt time.Time - LastActivityAt time.Time + SessionID string + HostPeerID string + ProtocolVersion int + hostVerifier reconnectVerifier + peerVerifiers map[string]reconnectVerifier + Peers map[string]*Client `json:"-"` + quotaOwnerKey string `json:"-"` + mu sync.RWMutex `json:"-"` + closing bool `json:"-"` + CreatedAt time.Time + LastActivityAt time.Time } // --- Snapshot types (on-disk JSON format) --- type roomSnapshot struct { - SessionID string `json:"sessionId"` - HostPeerID string `json:"hostPeerId"` - CreatedAt time.Time `json:"createdAt"` - LastActivityAt time.Time `json:"lastActivityAt"` + SessionID string `json:"sessionId"` + HostPeerID string `json:"hostPeerId"` + ProtocolVersion int `json:"protocolVersion,omitempty"` + HostReconnectVerifier string `json:"hostReconnectVerifier"` + PeerReconnectVerifiers map[string]string `json:"peerReconnectVerifiers,omitempty"` + CreatedAt time.Time `json:"createdAt"` + LastActivityAt time.Time `json:"lastActivityAt"` } type stateSnapshot struct { @@ -165,6 +252,43 @@ type stateSnapshot struct { Rooms []roomSnapshot `json:"rooms"` } +func mintReconnectToken() (string, reconnectVerifier, error) { + raw := make([]byte, reconnectTokenSize) + if _, err := rand.Read(raw); err != nil { + return "", reconnectVerifier{}, err + } + return base64.RawURLEncoding.EncodeToString(raw), sha256.Sum256(raw), nil +} + +func reconnectVerifierFromToken(token string) (reconnectVerifier, bool) { + if len(token) != base64.RawURLEncoding.EncodedLen(reconnectTokenSize) { + return reconnectVerifier{}, false + } + raw, err := base64.RawURLEncoding.DecodeString(token) + if err != nil || len(raw) != reconnectTokenSize { + return reconnectVerifier{}, false + } + return sha256.Sum256(raw), true +} + +func reconnectVerifierFromSnapshot(encoded string) (reconnectVerifier, bool) { + raw, err := base64.RawURLEncoding.DecodeString(encoded) + if err != nil || len(raw) != sha256.Size { + return reconnectVerifier{}, false + } + var verifier reconnectVerifier + copy(verifier[:], raw) + return verifier, true +} + +func encodeReconnectVerifier(verifier reconnectVerifier) string { + return base64.RawURLEncoding.EncodeToString(verifier[:]) +} + +func reconnectVerifierMatches(expected, presented reconnectVerifier) bool { + return subtle.ConstantTimeCompare(expected[:], presented[:]) == 1 +} + func (r *Room) peerIDs() []string { ids := make([]string, 0, len(r.Peers)) for id := range r.Peers { @@ -190,29 +314,150 @@ func (r *Room) broadcastExcept(senderID string, msg serverMsg) { r.mu.Unlock() for _, client := range targets { - client.trySend(data) + client.enqueue(data) } } -func (r *Room) sendTo(targetID string, msg serverMsg) bool { +func (r *Room) broadcastFrom(senderID string, sender *Client, msg serverMsg) bool { data, err := json.Marshal(msg) if err != nil { return false } r.mu.Lock() - client, ok := r.Peers[targetID] + if r.Peers[senderID] != sender { + r.mu.Unlock() + return false + } + if r.closing { + r.mu.Unlock() + return true + } + targets := make([]*Client, 0, len(r.Peers)-1) + r.LastActivityAt = time.Now() + for id, client := range r.Peers { + if id != senderID { + targets = append(targets, client) + } + } + r.mu.Unlock() + + for _, target := range targets { + target.enqueue(data) + } + return true +} + +type directedSendResult uint8 + +const ( + directedSenderUnavailable directedSendResult = iota + directedTargetMissing + directedTargetFound + directedSendSuppressed +) + +func (r *Room) sendFrom(senderID string, sender *Client, targetID string, msg serverMsg) directedSendResult { + data, err := json.Marshal(msg) + if err != nil { + return directedSenderUnavailable + } + r.mu.Lock() + if r.Peers[senderID] != sender { + r.mu.Unlock() + return directedSenderUnavailable + } + if r.closing { + r.mu.Unlock() + return directedSendSuppressed + } + target, ok := r.Peers[targetID] if ok { r.LastActivityAt = time.Now() } r.mu.Unlock() if !ok { - return false + return directedTargetMissing } - client.trySend(data) - return true + target.enqueue(data) + return directedTargetFound } // --- Log store --- +type artifactRemovalError struct { + err error +} + +func (e *artifactRemovalError) Error() string { + return "artifact removal failed" +} + +func (e *artifactRemovalError) Unwrap() error { + return e.err +} + +var errArtifactOutsideStore = errors.New("artifact path outside store") + +func classifyRemovalError(err error) error { + if err == nil || errors.Is(err, fs.ErrNotExist) { + return nil + } + return err +} + +func removeArtifact(removeFile func(string) error, root, path string) error { + err := classifyRemovalError(removeFile(path)) + if err == nil { + return nil + } + if !errors.Is(err, syscall.ENOTEMPTY) && !errors.Is(err, syscall.EEXIST) { + return &artifactRemovalError{err: err} + } + if err := removeConfinedDirectory(root, path); err != nil { + return &artifactRemovalError{err: err} + } + return nil +} + +func removeConfinedDirectory(root, path string) error { + info, err := os.Lstat(path) + if err != nil { + return classifyRemovalError(err) + } + if !info.IsDir() { + return syscall.ENOTDIR + } + + rootPath, err := filepath.Abs(root) + if err != nil { + return errArtifactOutsideStore + } + rootPath, err = filepath.EvalSymlinks(rootPath) + if err != nil { + return errArtifactOutsideStore + } + artifactPath, err := filepath.Abs(path) + if err != nil { + return errArtifactOutsideStore + } + artifactPath, err = filepath.EvalSymlinks(artifactPath) + if err != nil { + return errArtifactOutsideStore + } + relative, err := filepath.Rel(rootPath, artifactPath) + if err != nil || + relative == "." || + relative == ".." || + filepath.IsAbs(relative) || + strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + return errArtifactOutsideStore + } + return os.RemoveAll(artifactPath) +} + +type pendingRemoval struct { + size int64 + sizeKnown bool +} type logEntry struct { Size int @@ -223,24 +468,35 @@ type logEntry struct { var errLogStoreFull = errors.New("log store full") type logStore struct { - entries map[string]logEntry - rateLimit map[string]time.Time // IP -> last upload time - dir string - generateID func() string - mu sync.RWMutex + entries map[string]logEntry + pendingRemovals map[string]pendingRemoval + rateLimit map[string]time.Time // IP -> last upload time + failedLookupRate map[string]*rateLimiter + dir string + generateID func() string + removeFile func(string) error + startupErr error + mu sync.RWMutex } func newLogStore(dir string) *logStore { + return newLogStoreWithRemover(dir, os.Remove) +} + +func newLogStoreWithRemover(dir string, removeFile func(string) error) *logStore { if err := os.MkdirAll(dir, 0755); err != nil { log.Fatalf("failed to create log dir %s: %v", dir, err) } ls := &logStore{ - entries: make(map[string]logEntry), - rateLimit: make(map[string]time.Time), - dir: dir, - generateID: generateLogID, + entries: make(map[string]logEntry), + pendingRemovals: make(map[string]pendingRemoval), + rateLimit: make(map[string]time.Time), + failedLookupRate: make(map[string]*rateLimiter), + dir: dir, + generateID: generateLogID, + removeFile: removeFile, } - ls.loadExisting(time.Now()) + ls.startupErr = ls.loadExisting(time.Now()) return ls } @@ -271,45 +527,73 @@ func logIDFromFilename(filename string) (string, bool) { return id, validID(id, logIDLength) } -func (ls *logStore) loadExisting(now time.Time) { +func (ls *logStore) loadExisting(now time.Time) error { ls.mu.Lock() defer ls.mu.Unlock() files, err := os.ReadDir(ls.dir) if err != nil { log.Printf("logs: failed to read dir %s: %v", ls.dir, err) - return + return nil } + var removalErr error for _, file := range files { filename := file.Name() - path := filepath.Join(ls.dir, filename) if file.IsDir() || strings.HasSuffix(filename, ".tmp") { - os.RemoveAll(path) + removalErr = errors.Join(removalErr, ls.removeUntrackedLocked(filename)) continue } id, ok := logIDFromFilename(filename) if !ok { - os.Remove(path) + removalErr = errors.Join(removalErr, ls.removeUntrackedLocked(filename)) continue } - info, err := file.Info() - if err != nil || info.Size() <= 0 || info.Size() > maxLogSize { - os.Remove(path) + info, infoErr := file.Info() + if infoErr != nil || !info.Mode().IsRegular() || info.Size() <= 0 || info.Size() > maxLogSize { + removalErr = errors.Join(removalErr, ls.removeUntrackedLocked(filename)) continue } createdAt := info.ModTime() - expiresAt := createdAt.Add(logMaxAge) - if !now.Before(expiresAt) { - os.Remove(path) - continue - } ls.entries[id] = logEntry{ Size: int(info.Size()), CreatedAt: createdAt, - ExpiresAt: expiresAt, + ExpiresAt: createdAt.Add(logMaxAge), } } - ls.evictOldestLocked(maxLogEntries) + removalErr = errors.Join(removalErr, ls.cleanupExpiredLocked(now)) + removalErr = errors.Join(removalErr, ls.evictOldestLocked(maxLogEntries)) + return removalErr +} + +func (ls *logStore) removeUntrackedLocked(filename string) error { + if err := removeArtifact(ls.removeFile, ls.dir, filepath.Join(ls.dir, filename)); err != nil { + if _, exists := ls.pendingRemovals[filename]; !exists { + ls.pendingRemovals[filename] = pendingRemoval{} + } + return err + } + delete(ls.pendingRemovals, filename) + return nil +} + +func (ls *logStore) retryPendingLocked() error { + var removalErr error + for filename := range ls.pendingRemovals { + if err := removeArtifact(ls.removeFile, ls.dir, filepath.Join(ls.dir, filename)); err != nil { + removalErr = errors.Join(removalErr, err) + continue + } + delete(ls.pendingRemovals, filename) + } + return removalErr +} + +func (ls *logStore) cleanupFailedTempLocked(tmpPath string) { + _ = ls.removeUntrackedLocked(filepath.Base(tmpPath)) +} + +func (ls *logStore) artifactCountLocked() int { + return len(ls.entries) + len(ls.pendingRemovals) } func (ls *logStore) store(data []byte, now time.Time) (string, logEntry, error) { @@ -322,8 +606,12 @@ func (ls *logStore) store(data []byte, now time.Time) (string, logEntry, error) ls.mu.Lock() defer ls.mu.Unlock() - ls.cleanupExpiredLocked(now) - if len(ls.entries) >= maxLogEntries { + _ = ls.retryPendingLocked() + _ = ls.cleanupExpiredLocked(now) + if err := ls.evictOldestLocked(maxLogEntries); err != nil { + return "", logEntry{}, err + } + if ls.artifactCountLocked() >= maxLogEntries { return "", logEntry{}, errLogStoreFull } @@ -340,11 +628,11 @@ func (ls *logStore) store(data []byte, now time.Time) (string, logEntry, error) path := ls.filePath(id) tmpPath := path + ".tmp" if err := os.WriteFile(tmpPath, data, 0644); err != nil { - os.Remove(tmpPath) + ls.cleanupFailedTempLocked(tmpPath) return "", logEntry{}, err } if err := os.Rename(tmpPath, path); err != nil { - os.Remove(tmpPath) + ls.cleanupFailedTempLocked(tmpPath) return "", logEntry{}, err } _ = os.Chtimes(path, now, now) @@ -358,33 +646,52 @@ func (ls *logStore) store(data []byte, now time.Time) (string, logEntry, error) return id, entry, nil } -func (ls *logStore) lookup(id string, now time.Time) (logEntry, bool) { +func (ls *logStore) lookup(id string, now time.Time) (logEntry, bool, error) { if !validID(id, logIDLength) { - return logEntry{}, false + return logEntry{}, false, nil } ls.mu.Lock() defer ls.mu.Unlock() entry, ok := ls.entries[id] if !ok { - return logEntry{}, false + return logEntry{}, false, nil } if !now.Before(entry.ExpiresAt) { - ls.deleteEntryLocked(id) - return logEntry{}, false + if err := ls.deleteEntryLocked(id); err != nil { + return logEntry{}, false, err + } + return logEntry{}, false, nil } - return entry, true + return entry, true, nil } -func (ls *logStore) cleanupExpiredLocked(now time.Time) { +func (ls *logStore) allowFailedLookup(source string, now time.Time) bool { + ls.mu.Lock() + defer ls.mu.Unlock() + limiter := ls.failedLookupRate[source] + if limiter == nil { + cleanupRateLimiters(ls.failedLookupRate, now, nil) + if len(ls.failedLookupRate) >= maxFailedLogLookupSources { + return false + } + limiter = newRateLimiterAt(logLookupRateBurst, logLookupRateSustained, now) + ls.failedLookupRate[source] = limiter + } + return limiter.allowAt(now) +} + +func (ls *logStore) cleanupExpiredLocked(now time.Time) error { + var removalErr error for id, entry := range ls.entries { if !now.Before(entry.ExpiresAt) { - ls.deleteEntryLocked(id) + removalErr = errors.Join(removalErr, ls.deleteEntryLocked(id)) } } + return removalErr } -func (ls *logStore) evictOldestLocked(limit int) { - for len(ls.entries) > limit { +func (ls *logStore) evictOldestLocked(limit int) error { + for ls.artifactCountLocked() > limit { var oldestID string var oldest logEntry for id, entry := range ls.entries { @@ -394,23 +701,35 @@ func (ls *logStore) evictOldestLocked(limit int) { } } if oldestID == "" { - return + return nil + } + if err := ls.deleteEntryLocked(oldestID); err != nil { + return err } - ls.deleteEntryLocked(oldestID) } + return nil } -func (ls *logStore) deleteEntryLocked(id string) { - os.Remove(ls.filePath(id)) +func (ls *logStore) deleteEntryLocked(id string) error { + if _, ok := ls.entries[id]; !ok { + return nil + } + if err := removeArtifact(ls.removeFile, ls.dir, ls.filePath(id)); err != nil { + return err + } delete(ls.entries, id) + return nil } -func (ls *logStore) cleanup() { +func (ls *logStore) cleanup(now time.Time) error { ls.mu.Lock() defer ls.mu.Unlock() - now := time.Now() - ls.cleanupExpiredLocked(now) + removalErr := ls.retryPendingLocked() + removalErr = errors.Join(removalErr, ls.cleanupExpiredLocked(now)) + removalErr = errors.Join(removalErr, ls.evictOldestLocked(maxLogEntries)) cleanupRateWindows(ls.rateLimit, now, logRateInterval) + cleanupRateLimiters(ls.failedLookupRate, now, nil) + return removalErr } // --- Poster store --- @@ -424,25 +743,41 @@ type posterEntry struct { } type posterStore struct { - entries map[string]posterEntry - dir string - maxBytes int64 - maxAge time.Duration - totalBytes int64 - mu sync.RWMutex + entries map[string]posterEntry + pendingRemovals map[string]pendingRemoval + dir string + maxBytes int64 + maxAge time.Duration + totalBytes int64 + pendingBytes int64 + unknownPending int + removeFile func(string) error + startupErr error + mu sync.RWMutex } func newPosterStore(dir string, maxBytes int64, maxAge time.Duration) *posterStore { + return newPosterStoreWithRemover(dir, maxBytes, maxAge, os.Remove) +} + +func newPosterStoreWithRemover( + dir string, + maxBytes int64, + maxAge time.Duration, + removeFile func(string) error, +) *posterStore { if err := os.MkdirAll(dir, 0755); err != nil { log.Fatalf("failed to create poster dir %s: %v", dir, err) } ps := &posterStore{ - entries: make(map[string]posterEntry), - dir: dir, - maxBytes: maxBytes, - maxAge: maxAge, + entries: make(map[string]posterEntry), + pendingRemovals: make(map[string]pendingRemoval), + dir: dir, + maxBytes: maxBytes, + maxAge: maxAge, + removeFile: removeFile, } - ps.loadExisting(time.Now()) + ps.startupErr = ps.loadExisting(time.Now()) return ps } @@ -511,50 +846,130 @@ func posterIDFromFilename(filename string) (string, bool) { return id, true } -func (ps *posterStore) loadExisting(now time.Time) { +func (ps *posterStore) loadExisting(now time.Time) error { ps.mu.Lock() defer ps.mu.Unlock() files, err := os.ReadDir(ps.dir) if err != nil { log.Printf("posters: failed to read dir %s: %v", ps.dir, err) - return + return nil } - for _, f := range files { - filename := f.Name() - path := ps.filePath(filename) - if f.IsDir() || strings.HasSuffix(filename, ".tmp") { - os.RemoveAll(path) + var removalErr error + for _, file := range files { + filename := file.Name() + if file.IsDir() || strings.HasSuffix(filename, ".tmp") { + size, known := posterArtifactSize(file) + removalErr = errors.Join( + removalErr, + ps.removeUntrackedLocked(filename, size, known), + ) continue } id, ok := posterIDFromFilename(filename) if !ok { - os.Remove(path) + size, known := posterArtifactSize(file) + removalErr = errors.Join( + removalErr, + ps.removeUntrackedLocked(filename, size, known), + ) continue } - info, err := f.Info() - if err != nil { - os.Remove(path) + info, infoErr := file.Info() + if infoErr != nil || !info.Mode().IsRegular() { + removalErr = errors.Join( + removalErr, + ps.removeUntrackedLocked(filename, 0, false), + ) + continue + } + if _, duplicate := ps.entries[id]; duplicate { + removalErr = errors.Join( + removalErr, + ps.removeUntrackedLocked(filename, info.Size(), true), + ) continue } createdAt := info.ModTime() - expiresAt := createdAt.Add(ps.maxAge) - if !now.Before(expiresAt) { - os.Remove(path) - continue - } contentType, _ := posterContentTypeForExt(filepath.Ext(filename)) entry := posterEntry{ Filename: filename, Size: info.Size(), ContentType: contentType, CreatedAt: createdAt, - ExpiresAt: expiresAt, + ExpiresAt: createdAt.Add(ps.maxAge), } ps.entries[id] = entry ps.totalBytes += entry.Size } - ps.evictOldestLocked(0) + removalErr = errors.Join(removalErr, ps.cleanupExpiredLocked(now)) + removalErr = errors.Join(removalErr, ps.evictOldestLocked(0)) + return removalErr +} + +func posterArtifactSize(file fs.DirEntry) (int64, bool) { + info, err := file.Info() + if err != nil || !info.Mode().IsRegular() { + return 0, false + } + return info.Size(), true +} + +func (ps *posterStore) addPendingLocked(filename string, size int64, known bool) { + if _, exists := ps.pendingRemovals[filename]; exists { + return + } + ps.pendingRemovals[filename] = pendingRemoval{size: size, sizeKnown: known} + if known { + ps.pendingBytes += size + } else { + ps.unknownPending++ + } +} + +func (ps *posterStore) removeUntrackedLocked(filename string, size int64, known bool) error { + if err := removeArtifact(ps.removeFile, ps.dir, ps.filePath(filename)); err != nil { + ps.addPendingLocked(filename, size, known) + return err + } + return nil +} + +func (ps *posterStore) retryPendingLocked(knownOnly bool) error { + var removalErr error + for filename, pending := range ps.pendingRemovals { + if knownOnly && !pending.sizeKnown { + continue + } + if err := removeArtifact(ps.removeFile, ps.dir, ps.filePath(filename)); err != nil { + removalErr = errors.Join(removalErr, err) + continue + } + delete(ps.pendingRemovals, filename) + if pending.sizeKnown { + ps.pendingBytes -= pending.size + } else { + ps.unknownPending-- + } + } + return removalErr +} + +func (ps *posterStore) cleanupFailedTempLocked(tmpPath string) { + if err := removeArtifact(ps.removeFile, ps.dir, tmpPath); err == nil { + return + } + info, statErr := os.Stat(tmpPath) + known := statErr == nil && info.Mode().IsRegular() + var size int64 + if known { + size = info.Size() + } + ps.addPendingLocked(filepath.Base(tmpPath), size, known) +} + +func (ps *posterStore) accountedBytesLocked() int64 { + return ps.totalBytes + ps.pendingBytes } func (ps *posterStore) store(data []byte, contentType string, now time.Time) (string, posterEntry, error) { @@ -573,9 +988,16 @@ func (ps *posterStore) store(data []byte, contentType string, now time.Time) (st ps.mu.Lock() defer ps.mu.Unlock() - ps.cleanupExpiredLocked(now) - ps.evictOldestLocked(entrySize) - if ps.totalBytes+entrySize > ps.maxBytes { + // Known regular-file debt counts against quota and is retried on demand. + // Unknown artifacts are left to periodic cleanup: their size cannot be + // accounted safely, and a permanent directory or stat failure must not + // deny otherwise capacity-safe uploads. + _ = ps.retryPendingLocked(true) + _ = ps.cleanupExpiredLocked(now) + if err := ps.evictOldestLocked(entrySize); err != nil { + return "", posterEntry{}, err + } + if ps.accountedBytesLocked()+entrySize > ps.maxBytes { return "", posterEntry{}, errors.New("poster store full") } @@ -593,11 +1015,11 @@ func (ps *posterStore) store(data []byte, contentType string, now time.Time) (st path := ps.filePath(filename) tmpPath := path + ".tmp" if err := os.WriteFile(tmpPath, data, 0644); err != nil { - os.Remove(tmpPath) + ps.cleanupFailedTempLocked(tmpPath) return "", posterEntry{}, err } if err := os.Rename(tmpPath, path); err != nil { - os.Remove(tmpPath) + ps.cleanupFailedTempLocked(tmpPath) return "", posterEntry{}, err } _ = os.Chtimes(path, now, now) @@ -614,42 +1036,48 @@ func (ps *posterStore) store(data []byte, contentType string, now time.Time) (st return id, entry, nil } -func (ps *posterStore) lookup(filename string, now time.Time) (posterEntry, bool) { +func (ps *posterStore) lookup(filename string, now time.Time) (posterEntry, bool, error) { id, ok := posterIDFromFilename(filename) if !ok { - return posterEntry{}, false + return posterEntry{}, false, nil } ps.mu.Lock() defer ps.mu.Unlock() entry, ok := ps.entries[id] if !ok || entry.Filename != filename { - return posterEntry{}, false + return posterEntry{}, false, nil } if !now.Before(entry.ExpiresAt) { - ps.deleteEntryLocked(id, entry) - return posterEntry{}, false + if err := ps.deleteEntryLocked(id); err != nil { + return posterEntry{}, false, err + } + return posterEntry{}, false, nil } - return entry, true + return entry, true, nil } -func (ps *posterStore) cleanup(now time.Time) { +func (ps *posterStore) cleanup(now time.Time) error { ps.mu.Lock() defer ps.mu.Unlock() - ps.cleanupExpiredLocked(now) - ps.evictOldestLocked(0) + removalErr := ps.retryPendingLocked(false) + removalErr = errors.Join(removalErr, ps.cleanupExpiredLocked(now)) + removalErr = errors.Join(removalErr, ps.evictOldestLocked(0)) + return removalErr } -func (ps *posterStore) cleanupExpiredLocked(now time.Time) { +func (ps *posterStore) cleanupExpiredLocked(now time.Time) error { + var removalErr error for id, entry := range ps.entries { if !now.Before(entry.ExpiresAt) { - ps.deleteEntryLocked(id, entry) + removalErr = errors.Join(removalErr, ps.deleteEntryLocked(id)) } } + return removalErr } -func (ps *posterStore) evictOldestLocked(extraBytes int64) { - for ps.totalBytes+extraBytes > ps.maxBytes && len(ps.entries) > 0 { +func (ps *posterStore) evictOldestLocked(extraBytes int64) error { + for ps.accountedBytesLocked()+extraBytes > ps.maxBytes && len(ps.entries) > 0 { var oldestID string var oldest posterEntry first := true @@ -661,19 +1089,26 @@ func (ps *posterStore) evictOldestLocked(extraBytes int64) { } } if oldestID == "" { - return + return nil + } + if err := ps.deleteEntryLocked(oldestID); err != nil { + return err } - ps.deleteEntryLocked(oldestID, oldest) } + return nil } -func (ps *posterStore) deleteEntryLocked(id string, entry posterEntry) { - os.Remove(ps.filePath(entry.Filename)) +func (ps *posterStore) deleteEntryLocked(id string) error { + entry, ok := ps.entries[id] + if !ok { + return nil + } + if err := removeArtifact(ps.removeFile, ps.dir, ps.filePath(entry.Filename)); err != nil { + return err + } delete(ps.entries, id) ps.totalBytes -= entry.Size - if ps.totalBytes < 0 { - ps.totalBytes = 0 - } + return nil } // --- Snapshotter (single-writer, debounced, atomic disk persistence) --- @@ -686,6 +1121,7 @@ type snapshotter struct { done chan struct{} exited chan struct{} build func() stateSnapshot + persist func([]byte) error writeMu sync.Mutex stopOnce sync.Once @@ -698,7 +1134,7 @@ func newSnapshotter(path string, build func() stateSnapshot) *snapshotter { if err := os.MkdirAll(dir, 0755); err != nil { log.Printf("snapshot: mkdir %s: %v", dir, err) } - return &snapshotter{ + sn := &snapshotter{ path: path, dir: dir, trigger: make(chan struct{}, 1), @@ -707,6 +1143,8 @@ func newSnapshotter(path string, build func() stateSnapshot) *snapshotter { exited: make(chan struct{}), build: build, } + sn.persist = sn.persistAtomic + return sn } func (sn *snapshotter) schedule() { @@ -753,7 +1191,13 @@ func (sn *snapshotter) write() error { if err != nil { return err } + if len(data) > snapshotMaxFileSize { + return fmt.Errorf("snapshot exceeds maximum size: %d > %d bytes", len(data), snapshotMaxFileSize) + } + return sn.persist(data) +} +func (sn *snapshotter) persistAtomic(data []byte) error { tmpPath := sn.path + ".tmp" f, err := os.OpenFile(tmpPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0644) if err != nil { @@ -777,12 +1221,17 @@ func (sn *snapshotter) write() error { os.Remove(tmpPath) return err } - // Best-effort dir fsync so the rename is durable after host crash. - if d, err := os.Open(sn.dir); err == nil { - d.Sync() - d.Close() + // The directory entry must reach stable storage before a terminal + // acknowledgement can rely on the rename surviving a host crash. + d, err := os.Open(sn.dir) + if err != nil { + return err } - return nil + if err := d.Sync(); err != nil { + d.Close() + return err + } + return d.Close() } func (sn *snapshotter) flushAndStop(timeout time.Duration) error { @@ -824,25 +1273,80 @@ func (sn *snapshotter) logWriteErr(err error) { } // --- Server --- - -type Server struct { - rooms map[string]*Room - logs *logStore - posters *posterStore - conns *connTracker - snap *snapshotter - oauth *oauthProxy // nil when OAUTH_BASE_URL is unset - mu sync.RWMutex +type removalErrorThrottle struct { + mu sync.Mutex + lastLog map[string]time.Time } -func newServer(logDir, stateFile, posterDir string) *Server { - s := &Server{ - rooms: make(map[string]*Room), - logs: newLogStore(logDir), - posters: newPosterStore(posterDir, maxPosterStoreSize, posterMaxAge), - conns: newConnTracker(), +func (s *Server) logRemovalError(store, operation string, err error) { + key := store + ":" + operation + s.removalErrors.mu.Lock() + defer s.removalErrors.mu.Unlock() + if last := s.removalErrors.lastLog[key]; !last.IsZero() && time.Since(last) < time.Hour { + return } - if p, ok := oauthConfigFromEnv(); ok { + if s.removalErrors.lastLog == nil { + s.removalErrors.lastLog = make(map[string]time.Time) + } + s.removalErrors.lastLog[key] = time.Now() + category := "other" + switch { + case errors.Is(err, errArtifactOutsideStore): + category = "confinement" + case errors.Is(err, fs.ErrPermission): + category = "permission" + case errors.Is(err, syscall.ENOSPC): + category = "capacity" + case errors.Is(err, syscall.EROFS): + category = "read_only" + case errors.Is(err, syscall.EBUSY): + category = "busy" + case errors.Is(err, syscall.ENOTEMPTY), errors.Is(err, syscall.EEXIST): + category = "not_empty" + } + var errno syscall.Errno + if errors.As(err, &errno) { + log.Printf("%s: %s removal failed: category=%s errno=%d", store, operation, category, errno) + return + } + log.Printf("%s: %s removal failed: category=%s errno=unknown", store, operation, category) +} + +type Server struct { + rooms map[string]*Room + logs *logStore + posters *posterStore + posterUploads *posterUploadLimiter + logLookups chan struct{} + posterBodyReadTimeout time.Duration + conns *connTracker + clientIPs clientIPResolver + snap *snapshotter + oauth *oauthProxy // nil when OAUTH_BASE_URL is unset + removalErrors removalErrorThrottle + beforeJoinRoomLock func() // test-only deterministic admission barrier + beforeTerminalDelivery func() // test-only post-persistence, pre-delivery barrier + mu sync.RWMutex +} + +func newServer(logDir, stateFile, posterDir string, clientIPs clientIPResolver) *Server { + s := &Server{ + rooms: make(map[string]*Room), + logs: newLogStore(logDir), + posters: newPosterStore(posterDir, maxPosterStoreSize, posterMaxAge), + posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, time.Now()), + logLookups: make(chan struct{}, maxConcurrentLogLookups), + posterBodyReadTimeout: posterUploadReadTimeout, + conns: newConnTracker(), + clientIPs: clientIPs, + } + if s.logs.startupErr != nil { + s.logRemovalError("logs", "startup", s.logs.startupErr) + } + if s.posters.startupErr != nil { + s.logRemovalError("posters", "startup", s.posters.startupErr) + } + if p, ok := oauthConfigFromEnv(clientIPs); ok { s.oauth = p log.Printf("oauth: proxy enabled (base=%s, services=%d)", p.baseURL, len(p.services)) } @@ -855,6 +1359,20 @@ func newServer(logDir, stateFile, posterDir string) *Server { return s } +// removeRoomLocked removes room only while it is still the authoritative map +// entry. The caller must hold s.mu. A current-process quota reservation follows +// the retained room and is returned exactly once by successful removal. +func (s *Server) removeRoomLocked(sessionID string, room *Room) bool { + if s.rooms[sessionID] != room { + return false + } + delete(s.rooms, sessionID) + if room.quotaOwnerKey != "" { + s.conns.releaseRoom(room.quotaOwnerKey) + } + return true +} + // buildSnapshot copies room identity into a serializable value with no locks held during marshal. // Lock order: s.mu before room.mu, matching cleanupLoop. func (s *Server) buildSnapshot() stateSnapshot { @@ -867,11 +1385,21 @@ func (s *Server) buildSnapshot() stateSnapshot { } for _, room := range s.rooms { room.mu.RLock() + var peerVerifiers map[string]string + if len(room.peerVerifiers) != 0 { + peerVerifiers = make(map[string]string, len(room.peerVerifiers)) + for peerID, verifier := range room.peerVerifiers { + peerVerifiers[peerID] = encodeReconnectVerifier(verifier) + } + } snap.Rooms = append(snap.Rooms, roomSnapshot{ - SessionID: room.SessionID, - HostPeerID: room.HostPeerID, - CreatedAt: room.CreatedAt, - LastActivityAt: room.LastActivityAt, + SessionID: room.SessionID, + HostPeerID: room.HostPeerID, + ProtocolVersion: room.ProtocolVersion, + HostReconnectVerifier: encodeReconnectVerifier(room.hostVerifier), + PeerReconnectVerifiers: peerVerifiers, + CreatedAt: room.CreatedAt, + LastActivityAt: room.LastActivityAt, }) room.mu.RUnlock() } @@ -899,37 +1427,78 @@ func (s *Server) loadSnapshot(path string) error { log.Printf("snapshot: corrupt file at %s, starting fresh: %v", path, err) return nil } - if snap.Version != snapshotFormatVersion { + if snap.Version != 2 && snap.Version != snapshotFormatVersion { log.Printf("snapshot: unknown version %d, starting fresh", snap.Version) return nil } now := time.Now() - loaded, skipped := 0, 0 - s.mu.Lock() + restored := make(map[string]*Room, min(len(snap.Rooms), maxRetainedRooms)) + skipped := 0 for _, r := range snap.Rooms { if !validRelayID(r.SessionID, maxSessionIDLength) || !validRelayID(r.HostPeerID, maxPeerIDLength) { skipped++ continue } - if now.Sub(r.CreatedAt) > roomMaxAge { + hostVerifier, ok := reconnectVerifierFromSnapshot(r.HostReconnectVerifier) + if !ok { skipped++ continue } - if now.Sub(r.LastActivityAt) > emptyRoomMaxAge { + if r.ProtocolVersion != legacyRelayProtocolVersion && r.ProtocolVersion != relayProtocolVersion { skipped++ continue } - s.rooms[r.SessionID] = &Room{ - SessionID: r.SessionID, - HostPeerID: r.HostPeerID, - Peers: make(map[string]*Client), - CreatedAt: r.CreatedAt, - LastActivityAt: r.LastActivityAt, + var peerVerifiers map[string]reconnectVerifier + if len(r.PeerReconnectVerifiers) != 0 { + if r.ProtocolVersion == legacyRelayProtocolVersion { + skipped++ + continue + } + peerVerifiers = make(map[string]reconnectVerifier, len(r.PeerReconnectVerifiers)) + validPeerVerifiers := true + for peerID, encodedVerifier := range r.PeerReconnectVerifiers { + verifier, verifierOK := reconnectVerifierFromSnapshot(encodedVerifier) + if !validRelayID(peerID, maxPeerIDLength) || + peerID == r.HostPeerID || + !verifierOK || + len(peerVerifiers) >= maxRoomSize-1 { + validPeerVerifiers = false + break + } + peerVerifiers[peerID] = verifier + } + if !validPeerVerifiers { + skipped++ + continue + } + } + if now.Sub(r.CreatedAt) > roomMaxAge || now.Sub(r.LastActivityAt) > emptyRoomMaxAge { + skipped++ + continue + } + if _, duplicate := restored[r.SessionID]; duplicate { + skipped++ + continue + } + if len(restored) >= maxRetainedRooms { + log.Printf("snapshot: too many retained rooms, starting fresh") + return nil + } + restored[r.SessionID] = &Room{ + SessionID: r.SessionID, + HostPeerID: r.HostPeerID, + ProtocolVersion: r.ProtocolVersion, + hostVerifier: hostVerifier, + peerVerifiers: peerVerifiers, + Peers: make(map[string]*Client), + CreatedAt: r.CreatedAt, + LastActivityAt: r.LastActivityAt, } - loaded++ } + s.mu.Lock() + s.rooms = restored s.mu.Unlock() - log.Printf("snapshot: loaded %d rooms, skipped %d expired", loaded, skipped) + log.Printf("snapshot: loaded %d rooms, skipped %d invalid or expired rooms", len(restored), skipped) return nil } @@ -946,23 +1515,25 @@ func (s *Server) runCleanupStep(now time.Time) { changed := false var expiredClients []*Client for id, room := range s.rooms { - room.mu.RLock() + room.mu.Lock() empty := len(room.Peers) == 0 age := now.Sub(room.CreatedAt) idle := now.Sub(room.LastActivityAt) expired := age > roomMaxAge - if expired && !empty { - for _, client := range room.Peers { - expiredClients = append(expiredClients, client) + remove := (empty && idle > emptyRoomMaxAge) || expired + if remove { + room.closing = true + if expired && !empty { + for _, client := range room.Peers { + expiredClients = append(expiredClients, client) + } + clear(room.Peers) } - } - room.mu.RUnlock() - - if (empty && idle > emptyRoomMaxAge) || expired { log.Printf("cleanup: removing room %s (empty=%v, idle=%v, age=%v)", id, empty, idle, age) - delete(s.rooms, id) + s.removeRoomLocked(id, room) changed = true } + room.mu.Unlock() } roomCount := len(s.rooms) s.mu.Unlock() @@ -973,8 +1544,13 @@ func (s *Server) runCleanupStep(now time.Time) { if changed { s.snap.schedule() } - s.logs.cleanup() - s.posters.cleanup(now) + if err := s.logs.cleanup(now); err != nil { + s.logRemovalError("logs", "cleanup", err) + } + if err := s.posters.cleanup(now); err != nil { + s.logRemovalError("posters", "cleanup", err) + } + s.posterUploads.cleanup(now) s.conns.cleanup(now) if s.oauth != nil { s.oauth.cleanup() @@ -986,34 +1562,17 @@ func (s *Server) runCleanupStep(now time.Time) { s.conns.mu.Unlock() } -func clientIP(r *http.Request) string { - var raw string - if fwd := r.Header.Get("X-Forwarded-For"); fwd != "" { - raw = strings.TrimSpace(strings.SplitN(fwd, ",", 2)[0]) - } else { - host, _, err := net.SplitHostPort(r.RemoteAddr) - if err != nil { - raw = r.RemoteAddr - } else { - raw = host - } - } - // Normalize IPv6 to /64 prefix to prevent per-address bypass - ip := net.ParseIP(raw) - if ip != nil && ip.To4() == nil { - mask := net.CIDRMask(64, 128) - return ip.Mask(mask).String() - } - return raw -} - func (s *Server) handlePostLogs(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) return } - ip := clientIP(r) + ip, err := s.clientIPs.resolve(r) + if err != nil { + http.Error(w, "Invalid client address", http.StatusBadRequest) + return + } s.logs.mu.Lock() if last, ok := s.logs.rateLimit[ip]; ok && time.Since(last) < logRateInterval { s.logs.mu.Unlock() @@ -1043,44 +1602,112 @@ func (s *Server) handlePostLogs(w http.ResponseWriter, r *http.Request) { http.Error(w, "Log store full", http.StatusServiceUnavailable) return } - log.Printf("logs: failed to store from %s: %v", ip, err) + var removalErr *artifactRemovalError + if errors.As(err, &removalErr) { + s.logRemovalError("logs", "store", err) + } else { + log.Printf("logs: failed to store from %s: %v", ip, err) + } http.Error(w, "Failed to store log", http.StatusInternalServerError) return } - log.Printf("logs: stored %s (%d bytes) from %s", id, entry.Size, ip) + log.Printf("logs: stored %d bytes from %s", entry.Size, ip) w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(map[string]string{"id": id}) } func (s *Server) handleGetLogs(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "private, no-store") if r.Method != http.MethodGet { http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) return } - id := strings.TrimPrefix(r.URL.Path, "/logs/") - if id == "" || len(id) != logIDLength { - http.Error(w, "Not found", http.StatusNotFound) - return + type lookupResult struct { + entry logEntry + data []byte + status int + message string + } + lookup := func() lookupResult { + select { + case s.logLookups <- struct{}{}: + defer func() { <-s.logLookups }() + default: + return lookupResult{ + status: http.StatusTooManyRequests, + message: "Too many concurrent lookups", + } + } + + source, err := s.clientIPs.resolve(r) + if err != nil { + return lookupResult{ + status: http.StatusBadRequest, + message: "Invalid client address", + } + } + id := strings.TrimPrefix(r.URL.Path, "/logs/") + entry, ok, err := s.logs.lookup(id, time.Now()) + if err != nil { + s.logRemovalError("logs", "lookup", err) + return lookupResult{ + status: http.StatusInternalServerError, + message: "Failed to retrieve log", + } + } + if !ok { + if !s.logs.allowFailedLookup(source, time.Now()) { + return lookupResult{ + status: http.StatusTooManyRequests, + message: "Too many failed lookups", + } + } + return lookupResult{status: http.StatusNotFound, message: "Not found"} + } + + data, err := os.ReadFile(s.logs.filePath(id)) + if err != nil { + return lookupResult{status: http.StatusNotFound, message: "Not found"} + } + return lookupResult{entry: entry, data: data} + }() + + controller := http.NewResponseController(w) + if err := controller.SetWriteDeadline(time.Now().Add(httpResponseWriteTimeout)); err == nil { + defer controller.SetWriteDeadline(time.Time{}) + } else if !errors.Is(err, http.ErrNotSupported) { + log.Printf("logs: failed to set response write deadline") } - entry, ok := s.logs.lookup(id, time.Now()) - if !ok { - http.Error(w, "Not found", http.StatusNotFound) - return - } - - data, err := os.ReadFile(s.logs.filePath(id)) - if err != nil { - http.Error(w, "Not found", http.StatusNotFound) + if lookup.status != 0 { + http.Error(w, lookup.message, lookup.status) return } w.Header().Set("Content-Type", "text/plain; charset=utf-8") - w.Header().Set("Content-Length", strconv.Itoa(entry.Size)) - w.Write(data) + w.Header().Set("Content-Length", strconv.Itoa(lookup.entry.Size)) + if written, err := w.Write(lookup.data); err != nil || written != len(lookup.data) { + log.Printf("logs: response write failed") + } +} + +var errPosterBodyReadTimeout = errors.New("poster body read timeout") + +func readPosterBody(body io.ReadCloser, maxBytes int64, timeout time.Duration) ([]byte, error) { + timedOut := make(chan struct{}) + timer := time.AfterFunc(timeout, func() { + close(timedOut) + _ = body.Close() + }) + data, err := io.ReadAll(io.LimitReader(body, maxBytes+1)) + if timer.Stop() { + return data, err + } + <-timedOut + return nil, errPosterBodyReadTimeout } func (s *Server) handlePostPosters(w http.ResponseWriter, r *http.Request) { @@ -1089,7 +1716,31 @@ func (s *Server) handlePostPosters(w http.ResponseWriter, r *http.Request) { return } - body, err := io.ReadAll(io.LimitReader(r.Body, maxPosterSize+1)) + ip, err := s.clientIPs.resolve(r) + if err != nil { + http.Error(w, "Invalid client address", http.StatusBadRequest) + return + } + if !s.posterUploads.tryStart(ip, time.Now()) { + http.Error(w, "Too many poster uploads", http.StatusTooManyRequests) + return + } + defer s.posterUploads.finish() + + timeout := s.posterBodyReadTimeout + if timeout <= 0 { + timeout = posterUploadReadTimeout + } + readDeadline := http.NewResponseController(w) + if err := readDeadline.SetReadDeadline(time.Now().Add(timeout)); err == nil { + defer readDeadline.SetReadDeadline(time.Time{}) + } + body, err := readPosterBody(r.Body, maxPosterSize, timeout) + var timeoutErr net.Error + if errors.Is(err, errPosterBodyReadTimeout) || errors.As(err, &timeoutErr) && timeoutErr.Timeout() { + http.Error(w, "Request body timeout", http.StatusRequestTimeout) + return + } if err != nil { http.Error(w, "Failed to read body", http.StatusBadRequest) return @@ -1111,13 +1762,18 @@ func (s *Server) handlePostPosters(w http.ResponseWriter, r *http.Request) { id, entry, err := s.posters.store(body, contentType, time.Now()) if err != nil { - log.Printf("posters: failed to store from %s: %v", clientIP(r), err) + var removalErr *artifactRemovalError + if errors.As(err, &removalErr) { + s.logRemovalError("posters", "store", err) + } else { + log.Printf("posters: failed to store from %s: %v", ip, err) + } http.Error(w, "Failed to store poster", http.StatusInternalServerError) return } url := "/posters/" + entry.Filename - log.Printf("posters: stored %s (%d bytes) from %s", id, entry.Size, clientIP(r)) + log.Printf("posters: stored %s (%d bytes) from %s", id, entry.Size, ip) w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(map[string]any{ @@ -1134,7 +1790,12 @@ func (s *Server) handleGetPosters(w http.ResponseWriter, r *http.Request) { } filename := strings.TrimPrefix(r.URL.Path, "/posters/") - entry, ok := s.posters.lookup(filename, time.Now()) + entry, ok, err := s.posters.lookup(filename, time.Now()) + if err != nil { + s.logRemovalError("posters", "lookup", err) + http.Error(w, "Failed to retrieve poster", http.StatusInternalServerError) + return + } if !ok { http.Error(w, "Not found", http.StatusNotFound) return @@ -1157,7 +1818,13 @@ func (s *Server) handleGetPosters(w http.ResponseWriter, r *http.Request) { http.ServeContent(w, r, entry.Filename, entry.CreatedAt, f) } func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { - ip := clientIP(r) + ip, err := s.clientIPs.resolve(r) + if err != nil { + http.Error(w, "Invalid client address", http.StatusBadRequest) + return + } + // Retained-room ownership uses the same canonical source key as connection admission. + quotaOwnerKey := ip if !s.conns.tryConnect(ip) { http.Error(w, "Too many connections", http.StatusTooManyRequests) @@ -1179,14 +1846,12 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { return nil }) - // Client wraps the conn with a serialized write channel + ping ticker client := newClient(conn) defer client.close() rl := newRateLimiter(rateBurst, rateSustained) var currentRoom *Room var currentPeerID string - var isHost bool rejectRoomTransition := func() bool { if currentRoom == nil { return false @@ -1205,23 +1870,25 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { defer func() { if currentRoom != nil && currentPeerID != "" { currentRoom.mu.Lock() + closing := currentRoom.closing stale := currentRoom.Peers[currentPeerID] != client - if !stale { + if !closing && !stale { delete(currentRoom.Peers, currentPeerID) + if currentRoom.ProtocolVersion == legacyRelayProtocolVersion && + currentPeerID != currentRoom.HostPeerID { + delete(currentRoom.peerVerifiers, currentPeerID) + } currentRoom.LastActivityAt = time.Now() } currentRoom.mu.Unlock() - if !stale { + if !closing && !stale { currentRoom.broadcastExcept(currentPeerID, serverMsg{ Type: relayTypePeerLeft, PeerID: currentPeerID, }) s.snap.schedule() } - if isHost { - s.conns.releaseRoom(ip) - } - log.Printf("peer %s left room %s (stale=%v)", currentPeerID, currentRoom.SessionID, stale) + log.Printf("peer %s left room %s (closing=%v, stale=%v)", currentPeerID, currentRoom.SessionID, closing, stale) } }() @@ -1251,42 +1918,150 @@ 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 { + client.sendJSON(serverMsg{ + Type: relayTypeError, + Code: relayErrorProtocolMismatch, + Message: "Unsupported relay protocol version", + ProtocolVersion: relayProtocolVersion, + }) + continue + } if rejectRoomTransition() { continue } - if !s.conns.tryCreateRoom(ip) { - client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many rooms created"}) - continue - } - s.mu.Lock() - if existing, exists := s.rooms[msg.SessionID]; exists { - existing.mu.RLock() - empty := len(existing.Peers) == 0 - existing.mu.RUnlock() - if !empty { - s.mu.Unlock() - s.conns.releaseRoom(ip) - client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomExists, Message: "Room already exists"}) + + reconnectToken := msg.ReconnectToken + var hostVerifier reconnectVerifier + if reconnectToken == "" { + if msg.ProtocolVersion != legacyRelayProtocolVersion { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Modern room creation requires a reconnect token"}) continue } - // Empty stale room — reclaim the ID - delete(s.rooms, msg.SessionID) + var err error + reconnectToken, hostVerifier, err = mintReconnectToken() + if err != nil { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to create room"}) + continue + } + } else { + var ok bool + hostVerifier, ok = reconnectVerifierFromToken(reconnectToken) + if !ok { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid reconnect token"}) + continue + } + } + + var rejection *serverMsg + var oldHostClient *Client + var hostWasAbsent bool + s.mu.Lock() + existing := s.rooms[msg.SessionID] + if existing != nil { + existing.mu.Lock() + idempotentModernCreate := + !existing.closing && + msg.ProtocolVersion == relayProtocolVersion && + existing.ProtocolVersion == relayProtocolVersion && + msg.PeerID == existing.HostPeerID && + reconnectVerifierMatches(existing.hostVerifier, hostVerifier) + if idempotentModernCreate { + oldHostClient = existing.Peers[msg.PeerID] + hostWasAbsent = oldHostClient == nil + existing.Peers[msg.PeerID] = client + existing.LastActivityAt = time.Now() + peers := existing.peerIDs() + existing.mu.Unlock() + s.mu.Unlock() + + currentRoom = existing + currentPeerID = msg.PeerID + if oldHostClient != nil && oldHostClient != client { + oldHostClient.close() + } + existingPeers := make([]string, 0, len(peers)-1) + for _, peerID := range peers { + if peerID != msg.PeerID { + existingPeers = append(existingPeers, peerID) + } + } + client.sendJSON(serverMsg{ + Type: relayTypeCreated, + SessionID: msg.SessionID, + HostPeerID: msg.PeerID, + ReconnectToken: reconnectToken, + ProtocolVersion: relayProtocolVersion, + Peers: existingPeers, + }) + if hostWasAbsent { + existing.broadcastExcept(msg.PeerID, serverMsg{ + Type: relayTypePeerJoined, + PeerID: msg.PeerID, + }) + } + s.snap.schedule() + continue + } + authorizedLegacyReplacement := + len(existing.Peers) == 0 && + !existing.closing && + msg.ProtocolVersion == legacyRelayProtocolVersion && + existing.ProtocolVersion == legacyRelayProtocolVersion && + msg.PeerID == existing.HostPeerID && + msg.ReconnectToken != "" && + reconnectVerifierMatches(existing.hostVerifier, hostVerifier) + existing.mu.Unlock() + if !authorizedLegacyReplacement { + rejection = &serverMsg{Type: relayTypeError, Code: relayErrorRoomExists, Message: "Room already exists"} + } + } else if len(s.rooms) >= maxRetainedRooms { + rejection = &serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many retained rooms"} + } + if rejection == nil { + var reserved bool + if existing == nil { + reserved = s.conns.tryCreateRoom(quotaOwnerKey) + } else { + reserved = s.conns.tryCreateRoomReplacing(quotaOwnerKey, existing.quotaOwnerKey) + } + if !reserved { + rejection = &serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many rooms created"} + } + } + if rejection != nil { + s.mu.Unlock() + client.sendJSON(*rejection) + continue + } + if existing != nil { + s.removeRoomLocked(msg.SessionID, existing) } now := time.Now() room := &Room{ - SessionID: msg.SessionID, - HostPeerID: msg.PeerID, - Peers: map[string]*Client{msg.PeerID: client}, - CreatedAt: now, - LastActivityAt: now, + SessionID: msg.SessionID, + HostPeerID: msg.PeerID, + ProtocolVersion: msg.ProtocolVersion, + hostVerifier: hostVerifier, + Peers: map[string]*Client{msg.PeerID: client}, + quotaOwnerKey: quotaOwnerKey, + CreatedAt: now, + LastActivityAt: now, } + s.rooms[msg.SessionID] = room s.mu.Unlock() + currentRoom = room currentPeerID = msg.PeerID - isHost = true log.Printf("room %s created by %s", msg.SessionID, msg.PeerID) - client.sendJSON(serverMsg{Type: relayTypeCreated, SessionID: msg.SessionID}) + client.sendJSON(serverMsg{ + Type: relayTypeCreated, + SessionID: msg.SessionID, + HostPeerID: msg.PeerID, + ReconnectToken: reconnectToken, + ProtocolVersion: msg.ProtocolVersion, + }) s.snap.schedule() case relayTypeJoin: @@ -1294,51 +2069,302 @@ 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 { + client.sendJSON(serverMsg{ + Type: relayTypeError, + Code: relayErrorProtocolMismatch, + Message: "Unsupported relay protocol version", + ProtocolVersion: relayProtocolVersion, + }) + continue + } if rejectRoomTransition() { continue } + + newToken, newVerifier, err := mintReconnectToken() + if err != nil { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to join room"}) + continue + } + presentedVerifier, tokenValid := reconnectVerifierFromToken(msg.ReconnectToken) + s.mu.RLock() room, exists := s.rooms[msg.SessionID] - s.mu.RUnlock() if !exists { + s.mu.RUnlock() client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"}) continue } + if s.beforeJoinRoomLock != nil { + s.beforeJoinRoomLock() + } room.mu.Lock() - if len(room.Peers) >= maxRoomSize { + if s.rooms[msg.SessionID] != room { room.mu.Unlock() - client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomFull, Message: "Room is full"}) + s.mu.RUnlock() + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"}) + continue + } + s.mu.RUnlock() + if room.closing { + room.mu.Unlock() + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"}) + continue + } + if room.ProtocolVersion != msg.ProtocolVersion { + room.mu.Unlock() + client.sendJSON(serverMsg{ + Type: relayTypeError, + Code: relayErrorProtocolMismatch, + Message: "Client and room protocol versions are incompatible", + ProtocolVersion: room.ProtocolVersion, + }) continue } - room.Peers[msg.PeerID] = client - room.LastActivityAt = time.Now() - peers := room.peerIDs() - room.mu.Unlock() - currentRoom = room - currentPeerID = msg.PeerID - log.Printf("peer %s joined room %s", msg.PeerID, msg.SessionID) - // Tell the joiner who's already here (excluding themselves) - existingPeers := make([]string, 0, len(peers)-1) - for _, p := range peers { - if p != msg.PeerID { - existingPeers = append(existingPeers, p) + existingClient, occupied := room.Peers[msg.PeerID] + expectedVerifier, identityReserved := room.peerVerifiers[msg.PeerID] + responseToken := newToken + responseVerifier := newVerifier + authorized := false + if room.ProtocolVersion == relayProtocolVersion { + switch { + case msg.PeerID == room.HostPeerID: + authorized = tokenValid && reconnectVerifierMatches(room.hostVerifier, presentedVerifier) + responseToken = msg.ReconnectToken + responseVerifier = room.hostVerifier + case identityReserved: + authorized = tokenValid && reconnectVerifierMatches(expectedVerifier, presentedVerifier) + responseToken = msg.ReconnectToken + responseVerifier = expectedVerifier + case occupied: + authorized = false + default: + authorized = tokenValid + responseToken = msg.ReconnectToken + responseVerifier = presentedVerifier + } + } else if msg.PeerID == room.HostPeerID { + switch { + case tokenValid && reconnectVerifierMatches(room.hostVerifier, presentedVerifier): + authorized = true + responseToken = msg.ReconnectToken + responseVerifier = room.hostVerifier + case msg.ReconnectToken == "" && + !occupied && + room.quotaOwnerKey != "" && + room.quotaOwnerKey == quotaOwnerKey: + // Tokenless host reconnect is retained only for unversioned rooms, + // only within this process, and only from the creating source. + authorized = true + responseToken = "" + responseVerifier = room.hostVerifier + } + } else { + // Legacy guests have no durable proof. Never let one replace a live + // identity; disconnected identity reuse remains confined to legacy rooms. + authorized = !occupied + } + + if !authorized { + room.mu.Unlock() + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorPeerIdUnavailable, Message: "Peer ID is unavailable"}) + continue + } + if !occupied { + roomFull := false + if room.ProtocolVersion == relayProtocolVersion { + if msg.PeerID != room.HostPeerID && !identityReserved { + roomFull = len(room.peerVerifiers) >= maxRoomSize-1 + } + } else { + admissionLimit := maxRoomSize + _, hostConnected := room.Peers[room.HostPeerID] + if msg.PeerID != room.HostPeerID && !hostConnected { + admissionLimit-- + } + roomFull = len(room.Peers) >= admissionLimit + } + if roomFull { + room.mu.Unlock() + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomFull, Message: "Room is full"}) + continue } } - client.sendJSON(serverMsg{Type: relayTypeJoined, SessionID: msg.SessionID, Peers: existingPeers}) + + if room.peerVerifiers == nil { + room.peerVerifiers = make(map[string]reconnectVerifier) + } + room.Peers[msg.PeerID] = client + if room.ProtocolVersion == relayProtocolVersion && msg.PeerID != room.HostPeerID { + room.peerVerifiers[msg.PeerID] = responseVerifier + } + room.LastActivityAt = time.Now() + peers := room.peerIDs() + hostPeerID := room.HostPeerID + roomProtocolVersion := room.ProtocolVersion + room.mu.Unlock() + + currentRoom = room + currentPeerID = msg.PeerID + if occupied && existingClient != client { + existingClient.close() + } + log.Printf("peer %s joined room %s", msg.PeerID, msg.SessionID) + + existingPeers := make([]string, 0, len(peers)-1) + for _, peerID := range peers { + if peerID != msg.PeerID { + existingPeers = append(existingPeers, peerID) + } + } + client.sendJSON(serverMsg{ + Type: relayTypeJoined, + SessionID: msg.SessionID, + HostPeerID: hostPeerID, + ReconnectToken: responseToken, + ProtocolVersion: roomProtocolVersion, + Peers: existingPeers, + }) room.broadcastExcept(msg.PeerID, serverMsg{Type: relayTypePeerJoined, PeerID: msg.PeerID}) s.snap.schedule() + case relayTypeLeave: + if currentRoom == nil { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"}) + continue + } + room := currentRoom + presentedVerifier, tokenValid := reconnectVerifierFromToken(msg.ReconnectToken) + room.mu.Lock() + currentClient := room.Peers[currentPeerID] == client + isGuest := currentPeerID != room.HostPeerID + authorized := currentClient && isGuest && !room.closing && msg.ProtocolVersion == room.ProtocolVersion + if authorized && room.ProtocolVersion == relayProtocolVersion { + expectedVerifier, ok := room.peerVerifiers[currentPeerID] + authorized = ok && tokenValid && reconnectVerifierMatches(expectedVerifier, presentedVerifier) + } + if !authorized { + room.mu.Unlock() + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorPeerIdUnavailable, Message: "Unable to release peer identity"}) + continue + } + releasedPeerID := currentPeerID + persistedMembershipChanged := room.ProtocolVersion == relayProtocolVersion + delete(room.Peers, releasedPeerID) + delete(room.peerVerifiers, releasedPeerID) + room.LastActivityAt = time.Now() + room.mu.Unlock() + + currentRoom = nil + currentPeerID = "" + if persistedMembershipChanged { + if err := s.snap.writeAndLog(); err != nil { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to persist released peer identity"}) + continue + } + } else { + s.snap.schedule() + } + if s.beforeTerminalDelivery != nil { + s.beforeTerminalDelivery() + } + client.sendJSON(serverMsg{ + Type: relayTypeLeft, + SessionID: room.SessionID, + PeerID: releasedPeerID, + ProtocolVersion: room.ProtocolVersion, + }) + room.broadcastExcept(releasedPeerID, serverMsg{Type: relayTypePeerLeft, PeerID: releasedPeerID}) + + case relayTypeEndSession: + if currentRoom == nil { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"}) + continue + } + room := currentRoom + presentedVerifier, tokenValid := reconnectVerifierFromToken(msg.ReconnectToken) + s.mu.Lock() + room.mu.Lock() + authorized := + s.rooms[room.SessionID] == room && + !room.closing && + currentPeerID == room.HostPeerID && + room.Peers[currentPeerID] == client && + msg.ProtocolVersion == room.ProtocolVersion + if authorized && room.ProtocolVersion == relayProtocolVersion { + authorized = tokenValid && reconnectVerifierMatches(room.hostVerifier, presentedVerifier) + } + if !authorized { + room.mu.Unlock() + s.mu.Unlock() + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorPeerIdUnavailable, Message: "Unable to end room"}) + continue + } + room.closing = true + guests := make([]*Client, 0, len(room.Peers)-1) + for peerID, peerClient := range room.Peers { + if peerID != currentPeerID { + guests = append(guests, peerClient) + } + } + s.removeRoomLocked(room.SessionID, room) + room.mu.Unlock() + s.mu.Unlock() + + // Persist after releasing both locks and before any success frame. + // The room remains undiscoverable while existing membership stays + // authoritative for orderly terminal delivery. + if err := s.snap.writeAndLog(); err != nil { + client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to persist ended room"}) + room.mu.Lock() + clear(room.Peers) + clear(room.peerVerifiers) + room.mu.Unlock() + currentRoom = nil + currentPeerID = "" + for _, guest := range guests { + guest.close() + } + continue + } + if s.beforeTerminalDelivery != nil { + s.beforeTerminalDelivery() + } + endedMessage := serverMsg{ + Type: relayTypeEnded, + SessionID: room.SessionID, + ProtocolVersion: room.ProtocolVersion, + } + client.sendJSON(endedMessage) + for _, guest := range guests { + guest.sendJSONAndWait(endedMessage) + } + + room.mu.Lock() + clear(room.Peers) + clear(room.peerVerifiers) + room.mu.Unlock() + currentRoom = nil + currentPeerID = "" + for _, guest := range guests { + guest.close() + } + case relayTypeBroadcast: if currentRoom == nil { client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"}) continue } - currentRoom.broadcastExcept(currentPeerID, serverMsg{ + if !currentRoom.broadcastFrom(currentPeerID, client, serverMsg{ Type: relayTypeMessage, From: currentPeerID, Payload: msg.Payload, - }) + }) { + client.close() + return + } case relayTypeSendTo: if currentRoom == nil { @@ -1349,11 +2375,15 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid to field"}) continue } - if !currentRoom.sendTo(msg.To, serverMsg{ + switch currentRoom.sendFrom(currentPeerID, client, msg.To, serverMsg{ Type: relayTypeMessage, From: currentPeerID, Payload: msg.Payload, }) { + case directedSenderUnavailable: + client.close() + return + case directedTargetMissing: client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Target peer not found"}) } @@ -1366,6 +2396,16 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { } } +func newHTTPServer(addr string, handler http.Handler) *http.Server { + return &http.Server{ + Addr: addr, + Handler: handler, + ReadTimeout: posterUploadReadTimeout, + WriteTimeout: httpResponseWriteTimeout, + MaxHeaderBytes: maxHTTPHeaderBytes, + } +} + func main() { addr := flag.String("addr", ":8080", "Listen address") logDir := flag.String("log-dir", "/data/logs", "Directory for log file storage") @@ -1373,7 +2413,12 @@ func main() { stateFile := flag.String("state-file", "/data/rooms.json", "Path to room snapshot file") flag.Parse() - srv := newServer(*logDir, *stateFile, *posterDir) + trustedProxyCIDRs, err := parseTrustedProxyCIDRs(os.Getenv("TRUSTED_PROXY_CIDRS")) + if err != nil { + log.Fatalf("invalid TRUSTED_PROXY_CIDRS") + } + clientIPs := newClientIPResolver(trustedProxyCIDRs) + srv := newServer(*logDir, *stateFile, *posterDir, clientIPs) mux := http.NewServeMux() mux.HandleFunc("/relay", srv.handleWS) @@ -1387,7 +2432,7 @@ func main() { mux.HandleFunc("/posters/", srv.handleGetPosters) registerOAuthRoutes(mux, srv.oauth) - httpSrv := &http.Server{Addr: *addr, Handler: mux} + httpSrv := newHTTPServer(*addr, mux) serveErr := make(chan error, 1) go func() { diff --git a/server/main_test.go b/server/main_test.go index 0c5f75e0..af7d5e42 100644 --- a/server/main_test.go +++ b/server/main_test.go @@ -1,10 +1,15 @@ package main import ( + "bufio" "bytes" "encoding/json" + "errors" "fmt" "io" + "io/fs" + "log" + "math" "net" "net/http" "net/http/httptest" @@ -14,12 +19,37 @@ import ( "strings" "sync" "sync/atomic" + "syscall" "testing" "time" "github.com/gorilla/websocket" ) +func TestGeneratedRelayProtocolVersionsMatchSpec(t *testing.T) { + data, err := os.ReadFile(filepath.Join("..", "relay_protocol.json")) + if err != nil { + t.Fatalf("read relay protocol spec: %v", err) + } + var spec struct { + ProtocolVersion int `json:"protocolVersion"` + LegacyProtocolVersion int `json:"legacyProtocolVersion"` + } + if err := json.Unmarshal(data, &spec); err != nil { + t.Fatalf("decode relay protocol spec: %v", err) + } + if relayProtocolVersion != spec.ProtocolVersion || + legacyRelayProtocolVersion != spec.LegacyProtocolVersion { + t.Fatalf( + "generated versions=(%d,%d), spec=(%d,%d)", + relayProtocolVersion, + legacyRelayProtocolVersion, + spec.ProtocolVersion, + spec.LegacyProtocolVersion, + ) + } +} + // newTestServer builds a Server wired for tests: no goroutines, no network, // logs scratched to a per-test temp dir. The snapshotter is constructed but // its goroutine is NOT started — tests drive it synchronously via write() @@ -27,31 +57,70 @@ import ( func newTestServer(t *testing.T, stateFile string) *Server { t.Helper() s := &Server{ - rooms: make(map[string]*Room), - logs: newLogStore(t.TempDir()), - posters: newPosterStore(t.TempDir(), maxPosterStoreSize, posterMaxAge), - conns: newConnTracker(), + rooms: make(map[string]*Room), + logs: newLogStore(t.TempDir()), + posters: newPosterStore(t.TempDir(), maxPosterStoreSize, posterMaxAge), + posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, time.Now()), + conns: newConnTracker(), + clientIPs: newClientIPResolver(nil), } s.snap = newSnapshotter(stateFile, s.buildSnapshot) return s } +func mustReconnectToken(t *testing.T) (string, reconnectVerifier) { + t.Helper() + token, verifier, err := mintReconnectToken() + if err != nil { + t.Fatalf("mint reconnect token: %v", err) + } + return token, verifier +} + +func makeRoomSnapshots(count int, maximumLengthIDs bool, now time.Time) []roomSnapshot { + rooms := make([]roomSnapshot, 0, count) + for i := range count { + suffix := fmt.Sprintf("%04d", i) + sessionID := "S" + suffix + hostPeerID := "H" + if maximumLengthIDs { + sessionID = suffix + strings.Repeat("S", maxSessionIDLength-len(suffix)) + hostPeerID = strings.Repeat("H", maxPeerIDLength) + } + rooms = append(rooms, roomSnapshot{ + SessionID: sessionID, + HostPeerID: hostPeerID, + HostReconnectVerifier: encodeReconnectVerifier(reconnectVerifier{1}), + CreatedAt: now.Add(-time.Minute), + LastActivityAt: now, + }) + } + return rooms +} + func TestSnapshotRoundTrip(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "rooms.json") s := newTestServer(t, path) now := time.Now().UTC().Truncate(time.Second) + _, verifier1 := mustReconnectToken(t) + _, verifier2 := mustReconnectToken(t) s.rooms["ABC12"] = &Room{ SessionID: "ABC12", HostPeerID: "host-1", + hostVerifier: verifier1, + peerVerifiers: make(map[string]reconnectVerifier), Peers: map[string]*Client{}, CreatedAt: now.Add(-time.Minute), LastActivityAt: now, + quotaOwnerKey: "203.0.113.44", } s.rooms["XYZ99"] = &Room{ SessionID: "XYZ99", HostPeerID: "host-2", + hostVerifier: verifier2, + peerVerifiers: make(map[string]reconnectVerifier), Peers: map[string]*Client{}, CreatedAt: now.Add(-time.Hour), LastActivityAt: now.Add(-time.Second), @@ -60,6 +129,13 @@ func TestSnapshotRoundTrip(t *testing.T) { if err := s.snap.write(); err != nil { t.Fatalf("write: %v", err) } + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read snapshot: %v", err) + } + if bytes.Contains(data, []byte("203.0.113.44")) || bytes.Contains(data, []byte("quotaOwner")) { + t.Fatalf("snapshot persisted process-local client identity: %s", data) + } // Reconstruct into a fresh Server and verify identity. s2 := newTestServer(t, path) @@ -78,12 +154,18 @@ func TestSnapshotRoundTrip(t *testing.T) { if r.HostPeerID != orig.HostPeerID { t.Errorf("%s: HostPeerID=%q want %q", id, r.HostPeerID, orig.HostPeerID) } + if !reconnectVerifierMatches(r.hostVerifier, orig.hostVerifier) { + t.Errorf("%s: host reconnect verifier did not round-trip", id) + } if !r.CreatedAt.Equal(orig.CreatedAt) { t.Errorf("%s: CreatedAt=%v want %v", id, r.CreatedAt, orig.CreatedAt) } if !r.LastActivityAt.Equal(orig.LastActivityAt) { t.Errorf("%s: LastActivityAt=%v want %v", id, r.LastActivityAt, orig.LastActivityAt) } + if r.quotaOwnerKey != "" { + t.Errorf("%s: quotaOwnerKey=%q after reload, want empty", id, r.quotaOwnerKey) + } if r.Peers == nil { t.Errorf("%s: Peers map nil after reload", id) } @@ -91,22 +173,29 @@ func TestSnapshotRoundTrip(t *testing.T) { t.Errorf("%s: expected empty Peers, got %d", id, len(r.Peers)) } } + s2.conns.mu.Lock() + defer s2.conns.mu.Unlock() + if len(s2.conns.roomsPerIP) != 0 { + t.Fatalf("reload restored process-local room quota: %v", s2.conns.roomsPerIP) + } } func TestLoadSkipsExpired(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "rooms.json") now := time.Now() + _, verifier := mustReconnectToken(t) + encodedVerifier := encodeReconnectVerifier(verifier) snap := stateSnapshot{ Version: snapshotFormatVersion, SavedAt: now, Rooms: []roomSnapshot{ - {SessionID: "FRESH", HostPeerID: "h", CreatedAt: now.Add(-time.Minute), LastActivityAt: now.Add(-30 * time.Second)}, - {SessionID: "OLD24", HostPeerID: "h", CreatedAt: now.Add(-25 * time.Hour), LastActivityAt: now.Add(-time.Second)}, - {SessionID: "IDLE6", HostPeerID: "h", CreatedAt: now.Add(-2 * time.Hour), LastActivityAt: now.Add(-6 * time.Minute)}, - {SessionID: "", HostPeerID: "h", CreatedAt: now, LastActivityAt: now}, - {SessionID: "NOHOS", HostPeerID: "", CreatedAt: now, LastActivityAt: now}, + {SessionID: "FRESH", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now.Add(-time.Minute), LastActivityAt: now.Add(-30 * time.Second)}, + {SessionID: "OLD24", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now.Add(-25 * time.Hour), LastActivityAt: now.Add(-time.Second)}, + {SessionID: "IDLE6", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now.Add(-2 * time.Hour), LastActivityAt: now.Add(-6 * time.Minute)}, + {SessionID: "", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now, LastActivityAt: now}, + {SessionID: "NOHOS", HostPeerID: "", HostReconnectVerifier: encodedVerifier, CreatedAt: now, LastActivityAt: now}, }, } data, err := json.Marshal(snap) @@ -160,18 +249,22 @@ func TestLoadHandlesMissing(t *testing.T) { } } -func TestLoadHandlesUnknownVersion(t *testing.T) { - dir := t.TempDir() - path := filepath.Join(dir, "rooms.json") - if err := os.WriteFile(path, []byte(`{"version":99,"rooms":[{"sessionId":"X"}]}`), 0644); err != nil { - t.Fatalf("write: %v", err) - } - s := newTestServer(t, path) - if err := s.loadSnapshot(path); err != nil { - t.Fatalf("loadSnapshot: %v", err) - } - if len(s.rooms) != 0 { - t.Fatalf("expected empty rooms for unknown version, got %d", len(s.rooms)) +func TestLoadRejectsSnapshotsWithoutHostAuthority(t *testing.T) { + for _, version := range []int{1, 99} { + t.Run(fmt.Sprintf("version_%d", version), func(t *testing.T) { + path := filepath.Join(t.TempDir(), "rooms.json") + data := []byte(fmt.Sprintf(`{"version":%d,"rooms":[{"sessionId":"X","hostPeerId":"H"}]}`, version)) + if err := os.WriteFile(path, data, 0644); err != nil { + t.Fatalf("write: %v", err) + } + s := newTestServer(t, path) + if err := s.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot: %v", err) + } + if len(s.rooms) != 0 { + t.Fatalf("expected empty rooms for version %d, got %d", version, len(s.rooms)) + } + }) } } @@ -217,6 +310,45 @@ func TestCleanupUsesIdleNotAge(t *testing.T) { } } +func TestRemoveRoomLockedReleasesOwnedQuotaExactlyOnce(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + ownerKey := "203.0.113.8" + room := &Room{ + SessionID: "OWNED", + HostPeerID: "H", + Peers: map[string]*Client{}, + quotaOwnerKey: ownerKey, + } + s.rooms[room.SessionID] = room + if !s.conns.tryCreateRoom(ownerKey) { + t.Fatal("reserve room quota") + } + + s.mu.Lock() + firstRemoval := s.removeRoomLocked(room.SessionID, room) + secondRemoval := s.removeRoomLocked(room.SessionID, room) + s.mu.Unlock() + if !firstRemoval || secondRemoval { + t.Fatalf("first removal=%v second removal=%v, want true then false", firstRemoval, secondRemoval) + } + s.conns.mu.Lock() + remaining := s.conns.roomsPerIP[ownerKey] + s.conns.mu.Unlock() + if remaining != 0 { + t.Fatalf("quota after repeated removal=%d, want 0", remaining) + } + + replacement := &Room{SessionID: room.SessionID, Peers: map[string]*Client{}} + s.mu.Lock() + s.rooms[room.SessionID] = replacement + staleRemoval := s.removeRoomLocked(room.SessionID, room) + authoritative := s.rooms[room.SessionID] + s.mu.Unlock() + if staleRemoval || authoritative != replacement { + t.Fatal("stale pointer removed the authoritative replacement") + } +} + func TestSnapshotAtomicWriteSurvivesRenameFailure(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "rooms.json") @@ -269,6 +401,134 @@ func TestSnapshotAtomicWriteSurvivesRenameFailure(t *testing.T) { } } +func TestSnapshotWriteRejectsOversizeAndPreservesLastValidFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "rooms.json") + s := newTestServer(t, path) + now := time.Now() + s.rooms["ORIG"] = &Room{ + SessionID: "ORIG", + HostPeerID: "H", + Peers: map[string]*Client{}, + CreatedAt: now, + LastActivityAt: now, + } + if err := s.snap.write(); err != nil { + t.Fatalf("write valid snapshot: %v", err) + } + original, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read valid snapshot: %v", err) + } + + s.snap.build = func() stateSnapshot { + return stateSnapshot{ + Version: snapshotFormatVersion, + SavedAt: now, + Rooms: []roomSnapshot{{ + SessionID: strings.Repeat("S", snapshotMaxFileSize), + HostPeerID: "H", + CreatedAt: now, + LastActivityAt: now, + }}, + } + } + if err := s.snap.write(); err == nil { + t.Fatal("oversized snapshot write succeeded") + } + after, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read preserved snapshot: %v", err) + } + if !bytes.Equal(after, original) { + t.Fatal("oversized write replaced the last valid snapshot") + } + + reloaded := newTestServer(t, path) + if err := reloaded.loadSnapshot(path); err != nil { + t.Fatalf("load preserved snapshot: %v", err) + } + if _, ok := reloaded.rooms["ORIG"]; !ok { + t.Fatal("last valid snapshot did not survive oversized write") + } +} + +func TestSnapshotAtRetainedRoomCapFitsAndReloads(t *testing.T) { + path := filepath.Join(t.TempDir(), "rooms.json") + s := newTestServer(t, path) + now := time.Now().UTC() + for _, room := range makeRoomSnapshots(maxRetainedRooms, true, now) { + s.rooms[room.SessionID] = &Room{ + SessionID: room.SessionID, + HostPeerID: room.HostPeerID, + Peers: map[string]*Client{}, + CreatedAt: room.CreatedAt, + LastActivityAt: room.LastActivityAt, + } + } + + snapshot := s.buildSnapshot() + data, err := json.Marshal(snapshot) + if err != nil { + t.Fatalf("marshal maximum snapshot: %v", err) + } + if len(data) > snapshotMaxFileSize { + t.Fatalf("maximum admitted snapshot is %d bytes, exceeds %d", len(data), snapshotMaxFileSize) + } + if err := s.snap.write(); err != nil { + t.Fatalf("write maximum snapshot: %v", err) + } + + reloaded := newTestServer(t, path) + if err := reloaded.loadSnapshot(path); err != nil { + t.Fatalf("load maximum snapshot: %v", err) + } + if got := len(reloaded.rooms); got != maxRetainedRooms { + t.Fatalf("reloaded rooms=%d, want %d", got, maxRetainedRooms) + } + for _, expected := range snapshot.Rooms { + room := reloaded.rooms[expected.SessionID] + if room == nil || room.Peers == nil { + t.Fatalf("room %q is not available for joins after reload", expected.SessionID) + } + } +} + +func TestLoadRejectsSnapshotOverRetainedRoomCapWithoutPartialState(t *testing.T) { + path := filepath.Join(t.TempDir(), "rooms.json") + now := time.Now().UTC() + snapshot := stateSnapshot{ + Version: snapshotFormatVersion, + SavedAt: now, + Rooms: makeRoomSnapshots(maxRetainedRooms+1, false, now), + } + data, err := json.Marshal(snapshot) + if err != nil { + t.Fatalf("marshal over-count snapshot: %v", err) + } + if len(data) > snapshotMaxFileSize { + t.Fatalf("over-count fixture is %d bytes, must exercise count limit below %d", len(data), snapshotMaxFileSize) + } + if err := os.WriteFile(path, data, 0644); err != nil { + t.Fatalf("write over-count snapshot: %v", err) + } + + s := newTestServer(t, path) + if err := s.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot: %v", err) + } + if len(s.rooms) != 0 { + t.Fatalf("over-count snapshot partially loaded %d rooms", len(s.rooms)) + } + s.conns.mu.Lock() + defer s.conns.mu.Unlock() + if len(s.conns.roomsPerIP) != 0 { + t.Fatalf("over-count snapshot changed process quota state: %v", s.conns.roomsPerIP) + } + if _, err := os.Stat(path); err != nil { + t.Fatalf("rejected snapshot was not preserved: %v", err) + } +} + func TestSnapshotDebounceCoalesces(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "rooms.json") @@ -277,10 +537,15 @@ func TestSnapshotDebounceCoalesces(t *testing.T) { buildCount int countMu sync.Mutex ) + built := make(chan struct{}, 1) sn := newSnapshotter(path, func() stateSnapshot { countMu.Lock() buildCount++ countMu.Unlock() + select { + case built <- struct{}{}: + default: + } return stateSnapshot{Version: snapshotFormatVersion, SavedAt: time.Now(), Rooms: nil} }) go sn.run() @@ -290,8 +555,18 @@ func TestSnapshotDebounceCoalesces(t *testing.T) { for i := 0; i < 20; i++ { sn.schedule() } - // Give the debounce window + a small buffer to actually run. - time.Sleep(snapshotDebounce + 50*time.Millisecond) + select { + case <-built: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for debounced snapshot build") + } + quiet := time.NewTimer(5 * snapshotDebounce) + defer quiet.Stop() + select { + case <-built: + t.Fatal("debounced burst produced an unexpected trailing snapshot build") + case <-quiet.C: + } countMu.Lock() got := buildCount @@ -314,18 +589,65 @@ type relayHarness struct { baseURL string } +func mustClientIPResolver(t *testing.T, cidrs string) clientIPResolver { + t.Helper() + prefixes, err := parseTrustedProxyCIDRs(cidrs) + if err != nil { + t.Fatalf("parse trusted proxies: %v", err) + } + return newClientIPResolver(prefixes) +} + func newRelayHarness(t *testing.T) *relayHarness { t.Helper() tmpDir := t.TempDir() return newRelayHarnessAt(t, tmpDir, filepath.Join(tmpDir, "rooms.json")) } +func newRelayHarnessNoTrust(t *testing.T) *relayHarness { + t.Helper() + tmpDir := t.TempDir() + return newRelayHarnessAtWithResolver( + t, + tmpDir, + filepath.Join(tmpDir, "rooms.json"), + newClientIPResolver(nil), + ) +} + // newRelayHarnessAt lets a test control the stateFile path so two harnesses // can share a snapshot across a simulated restart. func newRelayHarnessAt(t *testing.T, logDir, stateFile string) *relayHarness { t.Helper() - srv := newServer(logDir, stateFile, filepath.Join(t.TempDir(), "posters")) + return newRelayHarnessAtWithResolver(t, logDir, stateFile, mustClientIPResolver(t, "127.0.0.0/8")) +} +func newRelayHarnessAtWithResolver( + t *testing.T, + logDir, stateFile string, + clientIPs clientIPResolver, +) *relayHarness { + t.Helper() + srv := newServer(logDir, stateFile, filepath.Join(t.TempDir(), "posters"), clientIPs) + return newRelayHarnessWithServer(t, srv, true) +} + +func newStorageHarness(t *testing.T, logs *logStore, posters *posterStore) *relayHarness { + t.Helper() + srv := &Server{ + rooms: make(map[string]*Room), + logs: logs, + posters: posters, + posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, time.Now()), + logLookups: make(chan struct{}, maxConcurrentLogLookups), + conns: newConnTracker(), + clientIPs: mustClientIPResolver(t, "127.0.0.0/8"), + } + return newRelayHarnessWithServer(t, srv, false) +} + +func newRelayHarnessWithServer(t *testing.T, srv *Server, stopSnapshot bool) *relayHarness { + t.Helper() mux := http.NewServeMux() mux.HandleFunc("/relay", srv.handleWS) mux.HandleFunc("/logs", srv.handlePostLogs) @@ -336,13 +658,116 @@ func newRelayHarnessAt(t *testing.T, logDir, stateFile string) *relayHarness { httpSrv := httptest.NewServer(mux) t.Cleanup(func() { httpSrv.Close() - _ = srv.snap.flushAndStop(time.Second) + if stopSnapshot { + _ = srv.snap.flushAndStop(time.Second) + } }) u, _ := url.Parse(httpSrv.URL) wsURL := "ws://" + u.Host + "/relay" return &relayHarness{srv: srv, httpSrv: httpSrv, wsURL: wsURL, baseURL: httpSrv.URL} } +func createModernRoomWithGuest( + t *testing.T, + h *relayHarness, + sessionID, hostIP, guestIP string, +) (host, guest *testConn, hostToken, guestToken string) { + t.Helper() + hostToken, _ = mustReconnectToken(t) + host = h.dial(t, hostIP) + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: sessionID, + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ = mustReconnectToken(t) + guest = h.dial(t, guestIP) + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: sessionID, + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + return host, guest, hostToken, guestToken +} + +func injectSnapshotPersistenceFailure(t *testing.T, sn *snapshotter, injected error) { + t.Helper() + if err := sn.write(); err != nil { + t.Fatalf("persist pre-failure baseline: %v", err) + } + sn.writeMu.Lock() + original := sn.persist + sn.persist = func([]byte) error { return injected } + sn.writeMu.Unlock() + t.Cleanup(func() { + sn.writeMu.Lock() + sn.persist = original + sn.writeMu.Unlock() + }) +} + +func copySnapshotForRestart(t *testing.T, source string) string { + t.Helper() + data, err := os.ReadFile(source) + if err != nil { + t.Fatalf("read crash-window snapshot: %v", err) + } + path := filepath.Join(t.TempDir(), "rooms.json") + if err := os.WriteFile(path, data, 0644); err != nil { + t.Fatalf("copy crash-window snapshot: %v", err) + } + return path +} + +type deterministicRemover struct { + mu sync.Mutex + failures map[string]error + calls map[string]int +} + +func newDeterministicRemover() *deterministicRemover { + return &deterministicRemover{ + failures: make(map[string]error), + calls: make(map[string]int), + } +} + +func (r *deterministicRemover) remove(path string) error { + r.mu.Lock() + r.calls[path]++ + err := r.failures[path] + r.mu.Unlock() + if err != nil { + return err + } + return os.Remove(path) +} + +func (r *deterministicRemover) fail(path string, err error) { + r.mu.Lock() + defer r.mu.Unlock() + r.failures[path] = err +} + +func (r *deterministicRemover) recover(path string) { + r.mu.Lock() + defer r.mu.Unlock() + delete(r.failures, path) +} + +func (r *deterministicRemover) callCount(path string) int { + r.mu.Lock() + defer r.mu.Unlock() + return r.calls[path] +} func (h *relayHarness) dial(t *testing.T, ip string) *testConn { t.Helper() @@ -368,6 +793,56 @@ func (h *relayHarness) dialRaw(ip string) (*websocket.Conn, error) { return conn, err } +func newWebSocketPair(t *testing.T) (*websocket.Conn, *websocket.Conn) { + t.Helper() + + serverConnCh := make(chan *websocket.Conn, 1) + upgradeErrCh := make(chan error, 1) + httpServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + upgradeErrCh <- err + return + } + serverConnCh <- conn + })) + + u, err := url.Parse(httpServer.URL) + if err != nil { + httpServer.Close() + t.Fatalf("parse websocket pair URL: %v", err) + } + peerConn, _, err := websocket.DefaultDialer.Dial("ws://"+u.Host, nil) + if err != nil { + httpServer.Close() + t.Fatalf("dial websocket pair: %v", err) + } + + var serverConn *websocket.Conn + select { + case serverConn = <-serverConnCh: + case err := <-upgradeErrCh: + peerConn.Close() + httpServer.Close() + t.Fatalf("upgrade websocket pair: %v", err) + case <-time.After(2 * time.Second): + peerConn.Close() + httpServer.Close() + t.Fatal("timed out waiting for websocket pair upgrade") + } + + t.Cleanup(func() { + serverConn.Close() + peerConn.Close() + httpServer.Close() + }) + return serverConn, peerConn +} + +func (h *relayHarness) dialWithHeaders(headers http.Header) (*websocket.Conn, *http.Response, error) { + return websocket.DefaultDialer.Dial(h.wsURL, headers) +} + func (h *relayHarness) waitRoomPeers(t *testing.T, sessionID string, want int) { t.Helper() deadline := time.Now().Add(2 * time.Second) @@ -458,6 +933,18 @@ func (c *testConn) expectError(code string) serverMsg { return m } +func (c *testConn) expectAuthority(typ, hostPeerID string) serverMsg { + c.t.Helper() + message := c.expect(typ) + if message.HostPeerID != hostPeerID { + c.t.Fatalf("%s hostPeerId=%q, want %q", typ, message.HostPeerID, hostPeerID) + } + if _, ok := reconnectVerifierFromToken(message.ReconnectToken); !ok { + c.t.Fatalf("%s reconnectToken has invalid shape", typ) + } + return message +} + // recvNothing asserts no message arrives within the given window. Used to // verify silent paths (sender not receiving own broadcast, stale-peer skip). func (c *testConn) recvNothing(within time.Duration) { @@ -500,6 +987,144 @@ func (c *testConn) recvUntilClosed(within time.Duration) ([]serverMsg, error) { } } +func requireClientClosed(t *testing.T, client *Client) { + t.Helper() + select { + case <-client.done: + case <-time.After(2 * time.Second): + t.Fatal("client did not close") + } +} + +func requirePeerClosed(t *testing.T, peer *websocket.Conn) { + t.Helper() + testPeer := &testConn{t: t, conn: peer} + if messages, err := testPeer.recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("peer remained open after client failure (messages=%v): %v", messages, err) + } +} + +func TestClientQueueOverflowClosesConnection(t *testing.T) { + serverConn, peerConn := newWebSocketPair(t) + client := &Client{ + conn: serverConn, + send: make(chan outboundFrame, 1), + done: make(chan struct{}), + } + + if !client.enqueue([]byte(`{"sequence":1}`)) { + t.Fatal("first frame was not accepted") + } + if client.enqueue([]byte(`{"sequence":2}`)) { + t.Fatal("overflowing frame was accepted") + } + + requireClientClosed(t, client) + requirePeerClosed(t, peerConn) + if client.enqueue([]byte(`{"sequence":3}`)) { + t.Fatal("frame was accepted after terminal close") + } + + client.close() +} + +func TestClientQueueOverflowBroadcastKeepsHealthyRecipient(t *testing.T) { + slowServerConn, slowPeerConn := newWebSocketPair(t) + slow := &Client{ + conn: slowServerConn, + send: make(chan outboundFrame, 1), + done: make(chan struct{}), + } + if !slow.enqueue([]byte(`{"sequence":1}`)) { + t.Fatal("failed to prime slow client queue") + } + + healthyServerConn, healthyPeerConn := newWebSocketPair(t) + healthy := newClient(healthyServerConn) + t.Cleanup(healthy.close) + + room := &Room{ + Peers: map[string]*Client{ + "slow": slow, + "healthy": healthy, + }, + } + payload := json.RawMessage(`{"sequence":2}`) + room.broadcastExcept("sender", serverMsg{ + Type: relayTypeMessage, + From: "sender", + Payload: payload, + }) + + requireClientClosed(t, slow) + requirePeerClosed(t, slowPeerConn) + + received := (&testConn{t: t, conn: healthyPeerConn}).expect(relayTypeMessage) + if received.From != "sender" { + t.Fatalf("healthy recipient sender=%q, want sender", received.From) + } + if string(received.Payload) != string(payload) { + t.Fatalf("healthy recipient payload=%s, want %s", received.Payload, payload) + } +} + +func TestClientQueueOverflowDirectedTargetStillExists(t *testing.T) { + serverConn, peerConn := newWebSocketPair(t) + target := &Client{ + conn: serverConn, + send: make(chan outboundFrame, 1), + done: make(chan struct{}), + } + if !target.enqueue([]byte(`{"sequence":1}`)) { + t.Fatal("failed to prime directed target queue") + } + + sender := &Client{} + room := &Room{Peers: map[string]*Client{"sender": sender, "target": target}} + if result := room.sendFrom("sender", sender, "target", serverMsg{ + Type: relayTypeMessage, + From: "sender", + Payload: json.RawMessage(`{"sequence":2}`), + }); result != directedTargetFound { + t.Fatalf("full existing target result=%v, want directedTargetFound", result) + } + + requireClientClosed(t, target) + requirePeerClosed(t, peerConn) + if result := room.sendFrom("sender", sender, "missing", serverMsg{Type: relayTypeMessage}); result != directedTargetMissing { + t.Fatalf("missing target result=%v, want directedTargetMissing", result) + } +} + +func TestClientWriteFailureClosesConnection(t *testing.T) { + serverConn, _ := newWebSocketPair(t) + client := &Client{ + conn: serverConn, + send: make(chan outboundFrame, 1), + done: make(chan struct{}), + } + + if err := serverConn.Close(); err != nil { + t.Fatalf("close writer connection: %v", err) + } + client.send <- outboundFrame{data: []byte(`{"type":"queued"}`)} + + exited := make(chan struct{}) + go func() { + client.writePump() + close(exited) + }() + + requireClientClosed(t, client) + select { + case <-exited: + case <-time.After(2 * time.Second): + t.Fatal("write pump did not exit after write failure") + } + + client.close() +} + // ====================================================================== // Unit tests — pure logic // ====================================================================== @@ -695,7 +1320,7 @@ func TestConnTrackerCleanupPreservesEffectiveRateLimits(t *testing.T) { func TestConnTrackerConnectRateLimit(t *testing.T) { ct := newConnTracker() ip := "10.0.0.4" - for i := 0; i < connRateBurst; i++ { + for i := range connRateBurst { if !ct.tryConnect(ip) { t.Fatalf("warmup tryConnect %d: expected true", i) } @@ -708,41 +1333,202 @@ func TestConnTrackerConnectRateLimit(t *testing.T) { } } +func TestPosterUploadLimiterAdmissionPolicy(t *testing.T) { + now := time.Unix(1_700_000_000, 0) + + t.Run("per IP burst and independent clients", func(t *testing.T) { + limiter := newPosterUploadLimiter(2, 1, 10, 1, 10, now) + for range 2 { + if !limiter.tryStart("203.0.113.1", now) { + t.Fatal("per-IP burst rejected early") + } + limiter.finish() + } + if limiter.tryStart("203.0.113.1", now) { + t.Fatal("request beyond per-IP burst succeeded") + } + if !limiter.tryStart("203.0.113.2", now) { + t.Fatal("independent IP was denied") + } + limiter.finish() + }) + + t.Run("global burst spans distinct clients", func(t *testing.T) { + limiter := newPosterUploadLimiter(10, 1, 2, 1, 10, now) + for _, ip := range []string{"203.0.113.1", "203.0.113.2"} { + if !limiter.tryStart(ip, now) { + t.Fatalf("%s rejected before global burst exhausted", ip) + } + limiter.finish() + } + if limiter.tryStart("203.0.113.3", now) { + t.Fatal("request beyond global burst succeeded") + } + if len(limiter.perIP) != 2 { + t.Fatalf("globally denied request allocated per-IP state: %d buckets", len(limiter.perIP)) + } + }) + + t.Run("concurrency denial consumes no tokens", func(t *testing.T) { + limiter := newPosterUploadLimiter(1, 0, 2, 0, 1, now) + if !limiter.tryStart("203.0.113.1", now) { + t.Fatal("first upload denied") + } + if limiter.tryStart("203.0.113.2", now) { + t.Fatal("upload above concurrency limit succeeded") + } + limiter.finish() + if !limiter.tryStart("203.0.113.2", now) { + t.Fatal("concurrency denial consumed admission tokens") + } + limiter.finish() + }) + + t.Run("per IP denial refunds global token", func(t *testing.T) { + limiter := newPosterUploadLimiter(1, 0, 2, 0, 2, now) + if !limiter.tryStart("203.0.113.1", now) { + t.Fatal("first upload denied") + } + limiter.finish() + if limiter.tryStart("203.0.113.1", now) { + t.Fatal("exhausted IP unexpectedly admitted") + } + if !limiter.tryStart("203.0.113.2", now) { + t.Fatal("refunded global token was unavailable to another IP") + } + limiter.finish() + }) + + t.Run("finish restores only concurrency and time restores rate", func(t *testing.T) { + limiter := newPosterUploadLimiter(1, 1, 1, 1, 1, now) + if !limiter.tryStart("203.0.113.1", now) { + t.Fatal("first upload denied") + } + limiter.finish() + if limiter.active != 0 { + t.Fatalf("active=%d, want 0", limiter.active) + } + if limiter.tryStart("203.0.113.1", now) { + t.Fatal("finish incorrectly refunded rate tokens") + } + if !limiter.tryStart("203.0.113.1", now.Add(time.Second)) { + t.Fatal("sustained refill did not restore capacity") + } + limiter.finish() + }) + + t.Run("cleanup retains effective buckets then reclaims full ones", func(t *testing.T) { + limiter := newPosterUploadLimiter(2, 1, 10, 1, 2, now) + if !limiter.tryStart("203.0.113.1", now) { + t.Fatal("first upload denied") + } + limiter.finish() + limiter.cleanup(now) + if _, ok := limiter.perIP["203.0.113.1"]; !ok { + t.Fatal("cleanup removed effective per-IP limiter") + } + limiter.cleanup(now.Add(time.Second)) + if _, ok := limiter.perIP["203.0.113.1"]; ok { + t.Fatal("cleanup retained fully refilled per-IP limiter") + } + }) +} + // ====================================================================== -// clientIP unit tests +// clientIPResolver unit tests // ====================================================================== -func TestClientIPFromRemoteAddr(t *testing.T) { - r := &http.Request{RemoteAddr: "127.0.0.1:12345"} - if got := clientIP(r); got != "127.0.0.1" { - t.Fatalf("got %q, want 127.0.0.1", got) +func TestClientIPResolverTrustChains(t *testing.T) { + tests := []struct { + name string + trusted string + remote string + headers []string + want string + wantErr bool + }{ + {name: "absent forwarding header", remote: "127.0.0.1:12345", want: "127.0.0.1"}, + {name: "untrusted peer ignores spoof", remote: "198.51.100.10:12345", headers: []string{"203.0.113.5"}, want: "198.51.100.10"}, + {name: "one trusted proxy", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5"}, want: "203.0.113.5"}, + {name: "append chain ignores forged leftmost", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"198.51.100.99, 203.0.113.5"}, want: "203.0.113.5"}, + {name: "two trusted proxies", trusted: "10.0.0.0/8, 192.0.2.0/24", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5, 192.0.2.10"}, want: "203.0.113.5"}, + {name: "untrusted intermediate is client boundary", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5, 198.51.100.7"}, want: "198.51.100.7"}, + {name: "repeated header lines preserve chain", trusted: "10.0.0.0/8, 192.0.2.0/24", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5", "192.0.2.10"}, want: "203.0.113.5"}, + {name: "trusted proxy without forwarding header", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", want: "10.0.0.2"}, + {name: "IPv4 mapped peer is unmapped", remote: "[::ffff:192.0.2.4]:8080", want: "192.0.2.4"}, + {name: "IPv4 mapped forwarded address is unmapped", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"::ffff:203.0.113.5"}, want: "203.0.113.5"}, + {name: "native IPv6 client is grouped to 64", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"2001:db8:85a3:12::abcd"}, want: "2001:db8:85a3:12::"}, + {name: "untrusted malformed header is ignored", remote: "198.51.100.10:12345", headers: []string{"bad,,host:123"}, want: "198.51.100.10"}, + {name: "malformed immediate peer", remote: "not-an-address", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resolver := mustClientIPResolver(t, tt.trusted) + req := &http.Request{RemoteAddr: tt.remote, Header: make(http.Header)} + for _, value := range tt.headers { + req.Header.Add("X-Forwarded-For", value) + } + got, err := resolver.resolve(req) + if tt.wantErr { + if err == nil { + t.Fatalf("resolve()=%q, want error", got) + } + return + } + if err != nil { + t.Fatalf("resolve(): %v", err) + } + if got != tt.want { + t.Fatalf("resolve()=%q, want %q", got, tt.want) + } + }) } } -func TestClientIPFromXForwardedFor(t *testing.T) { - r := &http.Request{ - RemoteAddr: "10.0.0.1:8080", - Header: http.Header{"X-Forwarded-For": []string{"203.0.113.5, 10.0.0.1"}}, - } - if got := clientIP(r); got != "203.0.113.5" { - t.Fatalf("got %q, want 203.0.113.5", got) +func TestClientIPResolverRejectsMalformedTrustedChains(t *testing.T) { + resolver := mustClientIPResolver(t, "10.0.0.0/8") + for _, value := range []string{ + "", + "203.0.113.5,", + "203.0.113.5,,192.0.2.1", + "not-an-ip", + "203.0.113.5:1234", + "fe80::1%eth0", + } { + t.Run(fmt.Sprintf("%q", value), func(t *testing.T) { + req := &http.Request{ + RemoteAddr: "10.0.0.2:8080", + Header: http.Header{"X-Forwarded-For": []string{value}}, + } + if got, err := resolver.resolve(req); err == nil { + t.Fatalf("resolve()=%q, want error", got) + } + }) } } -func TestClientIPXFFTrimsWhitespace(t *testing.T) { - r := &http.Request{ - Header: http.Header{"X-Forwarded-For": []string{" 203.0.113.5 "}}, +func TestParseTrustedProxyCIDRs(t *testing.T) { + prefixes, err := parseTrustedProxyCIDRs(" ") + if err != nil || len(prefixes) != 0 { + t.Fatalf("empty config = %v, %v; want no prefixes", prefixes, err) } - if got := clientIP(r); got != "203.0.113.5" { - t.Fatalf("got %q, want 203.0.113.5", got) + prefixes, err = parseTrustedProxyCIDRs(" 10.1.2.3/8, ::ffff:192.0.2.12/120, 2001:db8::1/32 ") + if err != nil { + t.Fatalf("valid config: %v", err) } -} - -func TestClientIPIPv6NormalizesTo64(t *testing.T) { - r := &http.Request{RemoteAddr: "[2001:db8:85a3::8a2e:370:7334]:54321"} - got := clientIP(r) - if got != "2001:db8:85a3::" { - t.Fatalf("got %q, want 2001:db8:85a3::", got) + got := make([]string, len(prefixes)) + for i, prefix := range prefixes { + got[i] = prefix.String() + } + want := []string{"10.0.0.0/8", "192.0.2.0/24", "2001:db8::/32"} + if fmt.Sprint(got) != fmt.Sprint(want) { + t.Fatalf("prefixes=%v, want %v", got, want) + } + for _, value := range []string{"10.0.0.0/8,", "10.0.0.0/8,garbage", "::ffff:192.0.2.1/64"} { + if prefixes, err := parseTrustedProxyCIDRs(value); err == nil || prefixes != nil { + t.Fatalf("parseTrustedProxyCIDRs(%q)=(%v, %v), want nil error result", value, prefixes, err) + } } } @@ -751,6 +1537,9 @@ func TestClientIPIPv6NormalizesTo64(t *testing.T) { // ====================================================================== func TestGenerateLogIDShape(t *testing.T) { + if entropyBits := float64(logIDLength) * math.Log2(float64(len(idChars))); entropyBits < 128 { + t.Fatalf("log capability entropy=%f bits, want at least 128", entropyBits) + } seen := map[string]struct{}{} for i := 0; i < 200; i++ { id := generateLogID() @@ -776,10 +1565,10 @@ func TestGenerateLogIDShape(t *testing.T) { func TestCreateSucceeds(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "1.1.1.1") - c.send(clientMsg{Type: "create", SessionID: "ROOM1", PeerID: "host-a"}) - m := c.expect("created") - if m.SessionID != "ROOM1" { - t.Errorf("SessionID=%q want ROOM1", m.SessionID) + c.send(clientMsg{Type: relayTypeCreate, SessionID: "ROOM1", PeerID: "host-a"}) + message := c.expectAuthority(relayTypeCreated, "host-a") + if message.SessionID != "ROOM1" { + t.Errorf("SessionID=%q want ROOM1", message.SessionID) } h.waitRoomPeers(t, "ROOM1", 1) } @@ -810,32 +1599,214 @@ func TestCreateDuplicateReturnsRoomExists(t *testing.T) { c2.expectError("room_exists") } -func TestCreateReclaimsEmptyStaleRoom(t *testing.T) { +func TestCreateNegotiatesModernProtocolWithClientKnownToken(t *testing.T) { h := newRelayHarness(t) - // Pre-seed an empty stale room — mimics a post-restart reload. - h.srv.mu.Lock() - h.srv.rooms["STALE"] = &Room{ + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "1.1.1.40") + + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "MODERN_CREATE", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion + 1, + }) + mismatch := host.expectError(relayErrorProtocolMismatch) + if mismatch.ProtocolVersion != relayProtocolVersion { + t.Fatalf("protocol mismatch advertised version=%d, want %d", mismatch.ProtocolVersion, relayProtocolVersion) + } + + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "MODERN_CREATE", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: mismatch.ProtocolVersion, + }) + created := host.expectAuthority(relayTypeCreated, "H") + if created.ReconnectToken != hostToken { + t.Fatal("modern create rotated the client-known reconnect token") + } + if created.ProtocolVersion != relayProtocolVersion { + t.Fatalf("created protocolVersion=%d, want %d", created.ProtocolVersion, relayProtocolVersion) + } +} + +func TestModernCreateRetryAfterLostSetupResponseIsIdempotent(t *testing.T) { + h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) + create := clientMsg{ + Type: relayTypeCreate, + SessionID: "CREATE_RETRY", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + } + + first := h.dial(t, "1.1.1.41") + first.send(create) + if err := first.conn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatalf("set discarded response deadline: %v", err) + } + messageType, _, err := first.conn.ReadMessage() + if err != nil { + t.Fatalf("read discarded setup response: %v", err) + } + if messageType != websocket.TextMessage { + t.Fatalf("discarded setup response type=%d, want text", messageType) + } + + retry := h.dial(t, "1.1.1.42") + retry.send(create) + created := retry.expectAuthority(relayTypeCreated, "H") + if created.ReconnectToken != hostToken || created.ProtocolVersion != relayProtocolVersion { + t.Fatalf("retry authority changed: tokenMatch=%v protocol=%d", created.ReconnectToken == hostToken, created.ProtocolVersion) + } + if len(created.Peers) != 0 { + t.Fatalf("retry reported unexpected peers: %v", created.Peers) + } + if messages, err := first.recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("superseded create connection remained open: %v (frames=%v)", err, messages) + } + + h.srv.mu.RLock() + roomCount := len(h.srv.rooms) + room := h.srv.rooms["CREATE_RETRY"] + h.srv.mu.RUnlock() + if roomCount != 1 || room == nil { + t.Fatalf("idempotent retry retained rooms=%d targetPresent=%v", roomCount, room != nil) + } +} + +func TestIdempotentCreateReannouncesPreviouslyAbsentHost(t *testing.T) { + h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) + create := clientMsg{ + Type: relayTypeCreate, + SessionID: "CREATE_REANNOUNCE", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + } + + host := h.dial(t, "1.1.1.43") + host.send(create) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, "1.1.1.44") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "CREATE_REANNOUNCE", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + if err := host.conn.Close(); err != nil { + t.Fatalf("close original host: %v", err) + } + left := guest.expect(relayTypePeerLeft) + if left.PeerID != "H" { + t.Fatalf("disconnected host event peerId=%q, want H", left.PeerID) + } + + returning := h.dial(t, "1.1.1.45") + returning.send(create) + returning.expectAuthority(relayTypeCreated, "H") + reannounced := guest.expect(relayTypePeerJoined) + if reannounced.PeerID != "H" { + t.Fatalf("returning host event peerId=%q, want H", reannounced.PeerID) + } +} + +func TestCreateCannotReclaimReservedEmptyRoom(t *testing.T) { + h := newRelayHarness(t) + hostToken, hostVerifier := mustReconnectToken(t) + original := &Room{ SessionID: "STALE", HostPeerID: "old-host", + hostVerifier: hostVerifier, + peerVerifiers: make(map[string]reconnectVerifier), Peers: map[string]*Client{}, - CreatedAt: time.Now().Add(-time.Hour), - LastActivityAt: time.Now().Add(-time.Hour), + CreatedAt: time.Now().Add(-time.Minute), + LastActivityAt: time.Now(), + } + h.srv.mu.Lock() + h.srv.rooms["STALE"] = original + h.srv.mu.Unlock() + + creator := h.dial(t, "1.1.1.6") + creator.send(clientMsg{Type: relayTypeCreate, SessionID: "STALE", PeerID: "new-host"}) + creator.expectError(relayErrorRoomExists) + + unproved := h.dial(t, "1.1.1.60") + unproved.send(clientMsg{Type: relayTypeJoin, SessionID: "STALE", PeerID: "old-host"}) + unproved.expectError(relayErrorPeerIdUnavailable) + + reconnected := h.dial(t, "1.1.1.61") + reconnected.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "STALE", + PeerID: "old-host", + ReconnectToken: hostToken, + }) + reconnected.expectAuthority(relayTypeJoined, "old-host") + + h.srv.mu.RLock() + current := h.srv.rooms["STALE"] + h.srv.mu.RUnlock() + if current != original { + t.Fatal("reserved room identity was replaced") + } +} + +func TestCreateReclaimsOwnedEmptyRoomWithoutDoubleCharging(t *testing.T) { + h := newRelayHarness(t) + ownerKey := "1.1.1.61" + now := time.Now() + replacementToken, replacementVerifier := mustReconnectToken(t) + h.srv.mu.Lock() + for i := range maxRoomsPerIP { + sessionID := fmt.Sprintf("OWNED%d", i) + h.srv.rooms[sessionID] = &Room{ + SessionID: sessionID, + HostPeerID: "old-host", + hostVerifier: replacementVerifier, + peerVerifiers: make(map[string]reconnectVerifier), + Peers: map[string]*Client{}, + quotaOwnerKey: ownerKey, + CreatedAt: now.Add(-time.Hour), + LastActivityAt: now, + } + if !h.srv.conns.tryCreateRoom(ownerKey) { + h.srv.mu.Unlock() + t.Fatal("failed to seed retained quota") + } } h.srv.mu.Unlock() - c := h.dial(t, "1.1.1.6") - c.send(clientMsg{Type: "create", SessionID: "STALE", PeerID: "new-host"}) - c.expect("created") - - h.waitRoomPeers(t, "STALE", 1) + c := h.dial(t, ownerKey) + c.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "OWNED0", + PeerID: "old-host", + ReconnectToken: replacementToken, + }) + c.expect(relayTypeCreated) + h.srv.conns.mu.Lock() + quota := h.srv.conns.roomsPerIP[ownerKey] + h.srv.conns.mu.Unlock() + if quota != maxRoomsPerIP { + t.Fatalf("replacement quota=%d, want %d", quota, maxRoomsPerIP) + } h.srv.mu.RLock() - room := h.srv.rooms["STALE"] + roomCount := len(h.srv.rooms) h.srv.mu.RUnlock() - room.mu.RLock() - host := room.HostPeerID - room.mu.RUnlock() - if host != "new-host" { - t.Errorf("HostPeerID=%q, reclaim should have reset to new-host", host) + if roomCount != maxRoomsPerIP { + t.Fatalf("replacement room count=%d, want %d", roomCount, maxRoomsPerIP) } } @@ -853,6 +1824,326 @@ func TestCreateHitsRoomsPerIPLimit(t *testing.T) { c.expectError("rate_limited") } +func TestRetainedRoomQuotaSurvivesDisconnectAndReturnsOnRemoval(t *testing.T) { + h := newRelayHarness(t) + ownerKey := "1.1.1.70" + sessionIDs := make([]string, 0, maxRoomsPerIP) + for i := range maxRoomsPerIP { + sessionID := fmt.Sprintf("RETAIN%d", i) + host := h.dial(t, ownerKey) + host.send(clientMsg{Type: relayTypeCreate, SessionID: sessionID, PeerID: "H"}) + host.expect(relayTypeCreated) + host.conn.Close() + h.waitRoomPeers(t, sessionID, 0) + sessionIDs = append(sessionIDs, sessionID) + } + + for i, sessionID := range sessionIDs { + guest := h.dial(t, fmt.Sprintf("1.1.2.%d", i+1)) + guest.send(clientMsg{Type: relayTypeJoin, SessionID: sessionID, PeerID: "G"}) + guest.expect(relayTypeJoined) + guest.conn.Close() + h.waitRoomPeers(t, sessionID, 0) + } + + blocked := h.dial(t, ownerKey) + blocked.send(clientMsg{Type: relayTypeCreate, SessionID: "RETAINX", PeerID: "H"}) + blocked.expectError(relayErrorRateLimited) + + h.srv.conns.mu.Lock() + retainedQuota := h.srv.conns.roomsPerIP[ownerKey] + h.srv.conns.mu.Unlock() + if retainedQuota != maxRoomsPerIP { + t.Fatalf("retained quota=%d, want %d after creator disconnects", retainedQuota, maxRoomsPerIP) + } + + now := time.Now() + h.srv.mu.RLock() + idleRoom := h.srv.rooms[sessionIDs[0]] + h.srv.mu.RUnlock() + idleRoom.mu.Lock() + idleRoom.LastActivityAt = now.Add(-emptyRoomMaxAge - time.Second) + idleRoom.mu.Unlock() + h.srv.runCleanupStep(now) + + h.srv.conns.mu.Lock() + afterIdleRemoval := h.srv.conns.roomsPerIP[ownerKey] + h.srv.conns.mu.Unlock() + if afterIdleRemoval != maxRoomsPerIP-1 { + t.Fatalf("quota after idle removal=%d, want %d", afterIdleRemoval, maxRoomsPerIP-1) + } + blocked.send(clientMsg{Type: relayTypeCreate, SessionID: "RETAIN3", PeerID: "H"}) + blocked.expect(relayTypeCreated) + + occupied := h.dial(t, "1.1.2.99") + occupied.send(clientMsg{Type: relayTypeJoin, SessionID: sessionIDs[1], PeerID: "G"}) + occupied.expect(relayTypeJoined) + h.srv.mu.RLock() + expiringRoom := h.srv.rooms[sessionIDs[1]] + h.srv.mu.RUnlock() + expiringRoom.mu.Lock() + expiringRoom.CreatedAt = now.Add(-roomMaxAge - time.Second) + expiringRoom.mu.Unlock() + h.srv.runCleanupStep(now) + if _, err := occupied.recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("occupied expired room client remained connected: %v", err) + } + + h.srv.conns.mu.Lock() + afterHardExpiry := h.srv.conns.roomsPerIP[ownerKey] + h.srv.conns.mu.Unlock() + if afterHardExpiry != maxRoomsPerIP-1 { + t.Fatalf("quota after hard expiry=%d, want %d", afterHardExpiry, maxRoomsPerIP-1) + } + recovered := h.dial(t, ownerKey) + recovered.send(clientMsg{Type: relayTypeCreate, SessionID: "RETAIN4", PeerID: "H"}) + recovered.expect(relayTypeCreated) + h.srv.conns.mu.Lock() + finalQuota := h.srv.conns.roomsPerIP[ownerKey] + h.srv.conns.mu.Unlock() + if finalQuota != maxRoomsPerIP { + t.Fatalf("final retained quota=%d, want %d", finalQuota, maxRoomsPerIP) + } +} + +func TestGlobalRetainedRoomCapBlocksCreateButPreservesJoin(t *testing.T) { + h := newRelayHarness(t) + now := time.Now() + h.srv.mu.Lock() + for _, persisted := range makeRoomSnapshots(maxRetainedRooms, false, now) { + h.srv.rooms[persisted.SessionID] = &Room{ + SessionID: persisted.SessionID, + HostPeerID: persisted.HostPeerID, + Peers: map[string]*Client{}, + CreatedAt: persisted.CreatedAt, + LastActivityAt: persisted.LastActivityAt, + } + } + h.srv.mu.Unlock() + + ownerKey := "1.1.3.1" + client := h.dial(t, ownerKey) + client.send(clientMsg{Type: relayTypeCreate, SessionID: "OVERGLOBAL", PeerID: "H"}) + client.expectError(relayErrorRateLimited) + h.srv.mu.RLock() + roomCount := len(h.srv.rooms) + h.srv.mu.RUnlock() + snapshotRoomCount := len(h.srv.buildSnapshot().Rooms) + if roomCount != maxRetainedRooms || snapshotRoomCount != maxRetainedRooms { + t.Fatalf("rejected create mutated retained state: rooms=%d snapshot=%d", roomCount, snapshotRoomCount) + } + h.srv.conns.mu.Lock() + reservation := h.srv.conns.roomsPerIP[ownerKey] + h.srv.conns.mu.Unlock() + if reservation != 0 { + t.Fatalf("global rejection reserved per-source quota: %d", reservation) + } + + client.send(clientMsg{Type: relayTypeJoin, SessionID: "S0000", PeerID: "G"}) + client.expect(relayTypeJoined) + client.conn.Close() + h.waitRoomPeers(t, "S0000", 0) + + h.srv.mu.RLock() + expired := h.srv.rooms["S0001"] + h.srv.mu.RUnlock() + expired.mu.Lock() + expired.LastActivityAt = now.Add(-emptyRoomMaxAge - time.Second) + expired.mu.Unlock() + h.srv.runCleanupStep(now) + + creator := h.dial(t, ownerKey) + creator.send(clientMsg{Type: relayTypeCreate, SessionID: "AFTERGLOBAL", PeerID: "H"}) + creator.expect(relayTypeCreated) + h.srv.mu.RLock() + roomCount = len(h.srv.rooms) + h.srv.mu.RUnlock() + if roomCount != maxRetainedRooms { + t.Fatalf("room count after cleanup and create=%d, want %d", roomCount, maxRetainedRooms) + } +} + +func TestConcurrentCreatesCannotExceedGlobalRetainedRoomCap(t *testing.T) { + for iteration := range 8 { + t.Run(fmt.Sprintf("iteration_%d", iteration), func(t *testing.T) { + h := newRelayHarness(t) + now := time.Now() + h.srv.mu.Lock() + for _, persisted := range makeRoomSnapshots(maxRetainedRooms-1, false, now) { + h.srv.rooms[persisted.SessionID] = &Room{ + SessionID: persisted.SessionID, + HostPeerID: persisted.HostPeerID, + Peers: map[string]*Client{}, + CreatedAt: persisted.CreatedAt, + LastActivityAt: persisted.LastActivityAt, + } + } + h.srv.mu.Unlock() + + type createResult struct { + ownerKey string + message serverMsg + err error + } + connections := make([]*websocket.Conn, 2) + for i := range connections { + conn, err := h.dialRaw(fmt.Sprintf("1.1.4.%d", i+1)) + if err != nil { + t.Fatalf("dial create contender %d: %v", i, err) + } + connections[i] = conn + t.Cleanup(func() { conn.Close() }) + } + start := make(chan struct{}) + results := make(chan createResult, len(connections)) + for i, conn := range connections { + ownerKey := fmt.Sprintf("1.1.4.%d", i+1) + go func(conn *websocket.Conn, ownerKey string, index int) { + <-start + err := conn.WriteJSON(clientMsg{ + Type: relayTypeCreate, + SessionID: fmt.Sprintf("RACE%d", index), + PeerID: "H", + }) + var message serverMsg + if err == nil { + err = conn.ReadJSON(&message) + } + results <- createResult{ownerKey: ownerKey, message: message, err: err} + }(conn, ownerKey, i) + } + close(start) + + created, rejected := 0, 0 + acceptedOwner := "" + for range connections { + result := <-results + if result.err != nil { + t.Fatalf("concurrent create failed: %v", result.err) + } + switch { + case result.message.Type == relayTypeCreated: + created++ + acceptedOwner = result.ownerKey + case result.message.Type == relayTypeError && result.message.Code == relayErrorRateLimited: + rejected++ + default: + t.Fatalf("unexpected concurrent result: %+v", result.message) + } + } + if created != 1 || rejected != 1 { + t.Fatalf("created=%d rejected=%d, want one each", created, rejected) + } + h.srv.mu.RLock() + roomCount := len(h.srv.rooms) + h.srv.mu.RUnlock() + if roomCount != maxRetainedRooms { + t.Fatalf("room count=%d, want %d", roomCount, maxRetainedRooms) + } + h.srv.conns.mu.Lock() + acceptedQuota := h.srv.conns.roomsPerIP[acceptedOwner] + totalQuota := 0 + for _, count := range h.srv.conns.roomsPerIP { + totalQuota += count + } + h.srv.conns.mu.Unlock() + if acceptedQuota != 1 || totalQuota != 1 { + t.Fatalf("accepted quota=%d total quota=%d, want 1 and 1", acceptedQuota, totalQuota) + } + }) + } +} + +func TestRelayUntrustedXFFCannotRotateConnectionOrRoomIdentity(t *testing.T) { + t.Run("connections", func(t *testing.T) { + h := newRelayHarnessNoTrust(t) + var conns []*websocket.Conn + for i := 0; i < maxConnsPerIP; i++ { + conn, err := h.dialRaw(fmt.Sprintf("203.0.113.%d", i+1)) + if err != nil { + t.Fatalf("dial %d: %v", i, err) + } + conns = append(conns, conn) + } + t.Cleanup(func() { + for _, conn := range conns { + conn.Close() + } + }) + + headers := http.Header{"X-Forwarded-For": []string{"198.51.100.200"}} + conn, resp, err := h.dialWithHeaders(headers) + if conn != nil { + conn.Close() + t.Fatal("connection above direct peer limit unexpectedly succeeded") + } + if err == nil || resp == nil { + t.Fatalf("dial error=%v response=%v, want HTTP 429", err, resp) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusTooManyRequests { + t.Fatalf("status=%d, want 429", resp.StatusCode) + } + }) + + t.Run("rooms", func(t *testing.T) { + h := newRelayHarnessNoTrust(t) + for i := 0; i <= maxRoomsPerIP; i++ { + client := h.dial(t, fmt.Sprintf("203.0.113.%d", i+1)) + client.send(clientMsg{Type: "create", SessionID: fmt.Sprintf("SPOOF%d", i), PeerID: "host"}) + if i < maxRoomsPerIP { + client.expect("created") + } else { + client.expectError("rate_limited") + } + } + }) +} + +func TestRelayTrustedClientsHaveIndependentConnectionBuckets(t *testing.T) { + h := newRelayHarness(t) + var conns []*websocket.Conn + for i := 0; i < maxConnsPerIP; i++ { + conn, err := h.dialRaw("203.0.113.10") + if err != nil { + t.Fatalf("client A dial %d: %v", i, err) + } + conns = append(conns, conn) + } + conn, err := h.dialRaw("203.0.113.11") + if err != nil { + t.Fatalf("client B should have an independent bucket: %v", err) + } + conns = append(conns, conn) + t.Cleanup(func() { + for _, conn := range conns { + conn.Close() + } + }) +} + +func TestRelayMalformedTrustedChainDoesNotMutateAdmission(t *testing.T) { + h := newRelayHarness(t) + headers := http.Header{"X-Forwarded-For": []string{"203.0.113.5,"}} + conn, resp, err := h.dialWithHeaders(headers) + if conn != nil { + conn.Close() + t.Fatal("malformed trusted chain unexpectedly upgraded") + } + if err == nil || resp == nil { + t.Fatalf("dial error=%v response=%v, want HTTP 400", err, resp) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusBadRequest { + t.Fatalf("status=%d, want 400", resp.StatusCode) + } + h.srv.conns.mu.Lock() + defer h.srv.conns.mu.Unlock() + if h.srv.conns.globalCount != 0 || len(h.srv.conns.perIP) != 0 || len(h.srv.conns.ipRate) != 0 { + t.Fatalf("malformed chain mutated connection admission: %+v", h.srv.conns) + } +} + func TestConnectionCannotRetainMultipleRoomMemberships(t *testing.T) { h := newRelayHarness(t) ip := "1.1.1.8" @@ -893,12 +2184,12 @@ func TestConnectionCannotRetainMultipleRoomMemberships(t *testing.T) { func TestJoinSucceedsAndBroadcastsPeerJoined(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "2.0.0.1") - host.send(clientMsg{Type: "create", SessionID: "J1", PeerID: "H"}) - host.expect("created") + host.send(clientMsg{Type: relayTypeCreate, SessionID: "J1", PeerID: "H"}) + host.expectAuthority(relayTypeCreated, "H") guest := h.dial(t, "2.0.0.2") - guest.send(clientMsg{Type: "join", SessionID: "J1", PeerID: "G"}) - joined := guest.expect("joined") + guest.send(clientMsg{Type: relayTypeJoin, SessionID: "J1", PeerID: "G"}) + joined := guest.expectAuthority(relayTypeJoined, "H") if joined.SessionID != "J1" { t.Errorf("SessionID=%q want J1", joined.SessionID) } @@ -906,13 +2197,59 @@ func TestJoinSucceedsAndBroadcastsPeerJoined(t *testing.T) { t.Errorf("Peers=%v, want [H]", joined.Peers) } - // Host is broadcast a peerJoined for the new guest. - peerJoined := host.expect("peerJoined") + peerJoined := host.expect(relayTypePeerJoined) if peerJoined.PeerID != "G" { t.Errorf("peerJoined.PeerID=%q want G", peerJoined.PeerID) } } +func TestHostIdentityClaimsRequireReconnectCapability(t *testing.T) { + h := newRelayHarness(t) + host := h.dial(t, "2.0.1.1") + host.send(clientMsg{Type: relayTypeCreate, SessionID: "AUTH", PeerID: "HOST"}) + created := host.expectAuthority(relayTypeCreated, "HOST") + + guest := h.dial(t, "2.0.1.2") + guest.send(clientMsg{Type: relayTypeJoin, SessionID: "AUTH", PeerID: "GUEST"}) + guest.expectAuthority(relayTypeJoined, "HOST") + host.expect(relayTypePeerJoined) + + attacker := h.dial(t, "2.0.1.3") + attacker.send(clientMsg{Type: relayTypeJoin, SessionID: "AUTH", PeerID: "HOST"}) + attacker.expectError(relayErrorPeerIdUnavailable) + wrongToken, _ := mustReconnectToken(t) + attacker.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "AUTH", + PeerID: "HOST", + ReconnectToken: wrongToken, + }) + attacker.expectError(relayErrorPeerIdUnavailable) + attacker.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"forged":true}`)}) + attacker.expectError(relayErrorNotInRoom) + + host.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"real":true}`)}) + message := guest.expect(relayTypeMessage) + if message.From != "HOST" { + t.Fatalf("message sender=%q, want HOST", message.From) + } + + hostVerifier, ok := reconnectVerifierFromToken(created.ReconnectToken) + if !ok { + t.Fatal("created reconnect token became invalid") + } + h.srv.mu.RLock() + room := h.srv.rooms["AUTH"] + h.srv.mu.RUnlock() + room.mu.RLock() + currentVerifier := room.hostVerifier + currentHost := room.Peers["HOST"] + room.mu.RUnlock() + if !reconnectVerifierMatches(hostVerifier, currentVerifier) || currentHost == nil { + t.Fatal("failed claims mutated host authority") + } +} + func TestJoinMissingFieldsRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "2.0.0.3") @@ -943,22 +2280,464 @@ func TestJoinUnknownRoomFails(t *testing.T) { c.expectError("room_not_found") } -func TestJoinFullRoomRejected(t *testing.T) { +func TestReleasedGuestProbesDoNotConsumeDurableCapacity(t *testing.T) { h := newRelayHarness(t) - host := h.dial(t, "2.1.0.1") - host.send(clientMsg{Type: "create", SessionID: "FULL", PeerID: "H"}) - host.expect("created") + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "2.0.0.40") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "PROBE_CAPACITY", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") - // Fill up to maxRoomSize (host is #1), each from a distinct IP to avoid per-IP conn cap. - for i := 1; i < maxRoomSize; i++ { - guest := h.dial(t, fmt.Sprintf("2.1.0.%d", 100+i)) - guest.send(clientMsg{Type: "join", SessionID: "FULL", PeerID: fmt.Sprintf("G%d", i)}) - guest.expect("joined") + for i := range maxRoomSize * 2 { + peerID := fmt.Sprintf("P%d", i) + probeToken, _ := mustReconnectToken(t) + probe := h.dial(t, fmt.Sprintf("2.0.1.%d", i+1)) + probe.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "PROBE_CAPACITY", + PeerID: peerID, + ReconnectToken: probeToken, + ProtocolVersion: relayProtocolVersion, + }) + probe.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + probe.send(clientMsg{ + Type: relayTypeLeave, + ReconnectToken: probeToken, + ProtocolVersion: relayProtocolVersion, + }) + probe.expect(relayTypeLeft) + left := host.expect(relayTypePeerLeft) + if left.PeerID != peerID { + t.Fatalf("released probe event peerId=%q, want %q", left.PeerID, peerID) + } } + h.srv.mu.RLock() + room := h.srv.rooms["PROBE_CAPACITY"] + h.srv.mu.RUnlock() + room.mu.RLock() + reservations := len(room.peerVerifiers) + room.mu.RUnlock() + if reservations != 0 { + t.Fatalf("released probes retained %d durable reservations", reservations) + } + + for i := range maxRoomSize - 1 { + peerID := fmt.Sprintf("G%d", i) + token, _ := mustReconnectToken(t) + guest := h.dial(t, fmt.Sprintf("2.0.2.%d", i+1)) + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "PROBE_CAPACITY", + PeerID: peerID, + ReconnectToken: token, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + } +} + +func TestFullRoomAllowsOnlyAuthenticatedLiveReplacements(t *testing.T) { + h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "2.1.0.1") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "FULL", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + created := host.expectAuthority(relayTypeCreated, "H") + + guests := make(map[string]*testConn) + guestTokens := make(map[string]string) + for i := 1; i < maxRoomSize; i++ { + peerID := fmt.Sprintf("G%d", i) + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, fmt.Sprintf("2.1.0.%d", 100+i)) + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: peerID, + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + joined := guest.expectAuthority(relayTypeJoined, "H") + if joined.ReconnectToken != guestToken { + t.Fatalf("%s join rotated its client-known token", peerID) + } + guests[peerID] = guest + guestTokens[peerID] = guestToken + } + for range maxRoomSize - 1 { + host.expect(relayTypePeerJoined) + } + for range maxRoomSize - 2 { + guests["G1"].expect(relayTypePeerJoined) + } + + overflowToken, _ := mustReconnectToken(t) overflow := h.dial(t, "2.1.0.250") - overflow.send(clientMsg{Type: "join", SessionID: "FULL", PeerID: "LATE"}) - overflow.expectError("room_full") + overflow.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: "LATE", + ReconnectToken: overflowToken, + ProtocolVersion: relayProtocolVersion, + }) + overflow.expectError(relayErrorRoomFull) + overflow.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{}`)}) + overflow.expectError(relayErrorNotInRoom) + + unprovedHost := h.dial(t, "2.1.0.251") + unprovedHost.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: "H", + ProtocolVersion: relayProtocolVersion, + }) + unprovedHost.expectError(relayErrorPeerIdUnavailable) + unprovedGuest := h.dial(t, "2.1.0.252") + unprovedGuest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: "G1", + ProtocolVersion: relayProtocolVersion, + }) + unprovedGuest.expectError(relayErrorPeerIdUnavailable) + wrongGuestToken, _ := mustReconnectToken(t) + unprovedGuest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: "G1", + ReconnectToken: wrongGuestToken, + ProtocolVersion: relayProtocolVersion, + }) + unprovedGuest.expectError(relayErrorPeerIdUnavailable) + + newHost := h.dial(t, "2.1.0.253") + newHost.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + ProtocolVersion: relayProtocolVersion, + }) + hostJoined := newHost.expectAuthority(relayTypeJoined, "H") + if len(hostJoined.Peers) != maxRoomSize-1 { + t.Fatalf("replacement host peers=%v, want %d peers", hostJoined.Peers, maxRoomSize-1) + } + if messages, err := host.recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("displaced host did not close: %v (frames=%v)", err, messages) + } + hostReturn := guests["G1"].expect(relayTypePeerJoined) + if hostReturn.PeerID != "H" { + t.Fatalf("host replacement event peerId=%q", hostReturn.PeerID) + } + + newGuest := h.dial(t, "2.1.0.254") + newGuest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FULL", + PeerID: "G1", + ReconnectToken: guestTokens["G1"], + ProtocolVersion: relayProtocolVersion, + }) + newGuest.expectAuthority(relayTypeJoined, "H") + if messages, err := guests["G1"].recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("displaced guest did not close: %v (frames=%v)", err, messages) + } + guestReturn := newHost.expect(relayTypePeerJoined) + if guestReturn.PeerID != "G1" { + t.Fatalf("guest replacement event peerId=%q", guestReturn.PeerID) + } + + newHost.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"state":"current"}`)}) + message := newGuest.expect(relayTypeMessage) + if message.From != "H" { + t.Fatalf("replacement sender=%q, want H", message.From) + } +} + +func TestDisconnectedHostKeepsAReservedRoomSlot(t *testing.T) { + h := newRelayHarness(t) + host := h.dial(t, "2.1.1.1") + host.send(clientMsg{Type: relayTypeCreate, SessionID: "HOST_SLOT", PeerID: "H"}) + created := host.expectAuthority(relayTypeCreated, "H") + + for i := 1; i < maxRoomSize; i++ { + guest := h.dial(t, fmt.Sprintf("2.1.1.%d", 100+i)) + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "HOST_SLOT", + PeerID: fmt.Sprintf("G%d", i), + }) + guest.expectAuthority(relayTypeJoined, "H") + } + h.waitRoomPeers(t, "HOST_SLOT", maxRoomSize) + + if err := host.conn.Close(); err != nil { + t.Fatalf("close host: %v", err) + } + h.waitRoomPeers(t, "HOST_SLOT", maxRoomSize-1) + + lateGuest := h.dial(t, "2.1.1.250") + lateGuest.send(clientMsg{Type: relayTypeJoin, SessionID: "HOST_SLOT", PeerID: "LATE"}) + lateGuest.expectError(relayErrorRoomFull) + + returningHost := h.dial(t, "2.1.1.251") + returningHost.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "HOST_SLOT", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + }) + joined := returningHost.expectAuthority(relayTypeJoined, "H") + if len(joined.Peers) != maxRoomSize-1 { + t.Fatalf("returning host peers=%v, want %d peers", joined.Peers, maxRoomSize-1) + } + h.waitRoomPeers(t, "HOST_SLOT", maxRoomSize) +} + +func TestLegacySameSourceHostReconnectAndModernTokenEnforcement(t *testing.T) { + t.Run("legacy same-source tokenless reconnect", func(t *testing.T) { + h := newRelayHarness(t) + source := "2.1.2.1" + host := h.dial(t, source) + host.send(clientMsg{Type: relayTypeCreate, SessionID: "LEGACY_RECONNECT", PeerID: "H"}) + host.expectAuthority(relayTypeCreated, "H") + + guest := h.dial(t, "2.1.2.2") + guest.send(clientMsg{Type: relayTypeJoin, SessionID: "LEGACY_RECONNECT", PeerID: "G"}) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + if err := host.conn.Close(); err != nil { + t.Fatalf("close legacy host: %v", err) + } + left := guest.expect(relayTypePeerLeft) + if left.PeerID != "H" { + t.Fatalf("legacy disconnect peerId=%q, want H", left.PeerID) + } + + returning := h.dial(t, source) + returning.send(clientMsg{Type: relayTypeJoin, SessionID: "LEGACY_RECONNECT", PeerID: "H"}) + joined := returning.expect(relayTypeJoined) + if joined.HostPeerID != "H" || joined.ReconnectToken != "" || joined.ProtocolVersion != legacyRelayProtocolVersion { + t.Fatalf("legacy reconnect authority=%+v", joined) + } + rejoined := guest.expect(relayTypePeerJoined) + if rejoined.PeerID != "H" { + t.Fatalf("legacy reconnect event peerId=%q, want H", rejoined.PeerID) + } + }) + + t.Run("modern same-source reconnect requires token", func(t *testing.T) { + h := newRelayHarness(t) + source := "2.1.3.1" + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, source) + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "MODERN_RECONNECT", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, "2.1.3.2") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "MODERN_RECONNECT", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + if err := host.conn.Close(); err != nil { + t.Fatalf("close modern host: %v", err) + } + left := guest.expect(relayTypePeerLeft) + if left.PeerID != "H" { + t.Fatalf("modern disconnect peerId=%q, want H", left.PeerID) + } + + unproved := h.dial(t, source) + unproved.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "MODERN_RECONNECT", + PeerID: "H", + ProtocolVersion: relayProtocolVersion, + }) + unproved.expectError(relayErrorPeerIdUnavailable) + + returning := h.dial(t, source) + returning.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "MODERN_RECONNECT", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + joined := returning.expectAuthority(relayTypeJoined, "H") + if joined.ReconnectToken != hostToken || joined.ProtocolVersion != relayProtocolVersion { + t.Fatalf("modern reconnect authority changed: %+v", joined) + } + }) +} + +func TestJoinAdmissionIsAtomicWithEmptyRoomCleanup(t *testing.T) { + h := newRelayHarness(t) + _, hostVerifier := mustReconnectToken(t) + now := time.Now() + room := &Room{ + SessionID: "ATOMIC_CLEANUP", + HostPeerID: "H", + hostVerifier: hostVerifier, + peerVerifiers: make(map[string]reconnectVerifier), + Peers: make(map[string]*Client), + CreatedAt: now.Add(-time.Hour), + LastActivityAt: now.Add(-emptyRoomMaxAge - time.Second), + } + h.srv.mu.Lock() + h.srv.rooms[room.SessionID] = room + h.srv.mu.Unlock() + + reached := make(chan struct{}) + release := make(chan struct{}) + var once sync.Once + h.srv.beforeJoinRoomLock = func() { + once.Do(func() { + close(reached) + <-release + }) + } + + joiner := h.dial(t, "2.2.0.1") + joiner.send(clientMsg{Type: relayTypeJoin, SessionID: room.SessionID, PeerID: "G1"}) + <-reached + + cleanupDone := make(chan struct{}) + go func() { + h.srv.runCleanupStep(now) + close(cleanupDone) + }() + select { + case <-cleanupDone: + t.Fatal("cleanup passed a join that still owns the server read lock") + case <-time.After(100 * time.Millisecond): + } + + close(release) + joiner.expectAuthority(relayTypeJoined, "H") + <-cleanupDone + + h.srv.mu.RLock() + current := h.srv.rooms[room.SessionID] + h.srv.mu.RUnlock() + if current != room { + t.Fatal("successful join committed to a detached room") + } + room.mu.RLock() + _, present := room.Peers["G1"] + room.mu.RUnlock() + if !present { + t.Fatal("joining peer missing from authoritative room") + } + + second := h.dial(t, "2.2.0.2") + second.send(clientMsg{Type: relayTypeJoin, SessionID: room.SessionID, PeerID: "G2"}) + second.expectAuthority(relayTypeJoined, "H") + joiner.expect(relayTypePeerJoined) + second.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"atomic":true}`)}) + if message := joiner.expect(relayTypeMessage); message.From != "G2" { + t.Fatalf("message sender=%q, want G2", message.From) + } +} + +func TestJoinAdmissionIsAtomicWithReservedRoomCreate(t *testing.T) { + h := newRelayHarness(t) + _, hostVerifier := mustReconnectToken(t) + now := time.Now() + room := &Room{ + SessionID: "ATOMIC_CREATE", + HostPeerID: "H", + hostVerifier: hostVerifier, + peerVerifiers: make(map[string]reconnectVerifier), + Peers: make(map[string]*Client), + CreatedAt: now, + LastActivityAt: now, + } + h.srv.mu.Lock() + h.srv.rooms[room.SessionID] = room + h.srv.mu.Unlock() + + reached := make(chan struct{}) + release := make(chan struct{}) + var once sync.Once + h.srv.beforeJoinRoomLock = func() { + once.Do(func() { + close(reached) + <-release + }) + } + + joiner := h.dial(t, "2.3.0.1") + joiner.send(clientMsg{Type: relayTypeJoin, SessionID: room.SessionID, PeerID: "G"}) + <-reached + + creator := h.dial(t, "2.3.0.2") + creator.send(clientMsg{Type: relayTypeCreate, SessionID: room.SessionID, PeerID: "OTHER"}) + type readResult struct { + message serverMsg + err error + } + createResult := make(chan readResult, 1) + go func() { + creator.conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, data, err := creator.conn.ReadMessage() + if err != nil { + createResult <- readResult{err: err} + return + } + var message serverMsg + err = json.Unmarshal(data, &message) + createResult <- readResult{message: message, err: err} + }() + + select { + case result := <-createResult: + t.Fatalf("create completed before join admission committed: message=%+v err=%v", result.message, result.err) + case <-time.After(100 * time.Millisecond): + } + + close(release) + joiner.expectAuthority(relayTypeJoined, "H") + result := <-createResult + if result.err != nil { + t.Fatalf("read create result: %v", result.err) + } + if result.message.Type != relayTypeError || result.message.Code != relayErrorRoomExists { + t.Fatalf("create result=%+v, want room_exists", result.message) + } + h.srv.mu.RLock() + current := h.srv.rooms[room.SessionID] + h.srv.mu.RUnlock() + if current != room { + t.Fatal("reserved room was replaced during admission") + } } // ====================================================================== @@ -1179,67 +2958,603 @@ func TestDisconnectBroadcastsPeerLeft(t *testing.T) { func TestStalePeerSkipsCleanupBroadcast(t *testing.T) { h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) host := h.dial(t, "6.1.0.1") - host.send(clientMsg{Type: "create", SessionID: "D2", PeerID: "H"}) - host.expect("created") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "D2", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + guestToken, _ := mustReconnectToken(t) g1 := h.dial(t, "6.1.0.2") - g1.send(clientMsg{Type: "join", SessionID: "D2", PeerID: "G"}) - g1.expect("joined") - host.expect("peerJoined") + g1.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "D2", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + g1.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) - // Second connection with the SAME peerId overwrites room.Peers["G"]. g2 := h.dial(t, "6.1.0.3") - g2.send(clientMsg{Type: "join", SessionID: "D2", PeerID: "G"}) - g2.expect("joined") - host.expect("peerJoined") + g2.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "D2", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + g2.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) - // Now close g1 — its defer should see the stale client and NOT broadcast peerLeft. - g1.conn.Close() - host.recvNothing(300 * time.Millisecond) + if messages, err := g1.recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("displaced guest did not close: %v (frames=%v)", err, messages) + } + host.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"after":"replacement"}`)}) + message := g2.expect(relayTypeMessage) + if message.From != "H" { + t.Fatalf("post-replacement sender=%q, want H", message.From) + } +} + +func TestDisconnectedModernGuestIdentityRejectsTheftAndAcceptsRightfulReconnect(t *testing.T) { + h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "6.1.0.10") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "GUEST_RECONNECT", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, "6.1.0.11") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "GUEST_RECONNECT", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + if err := guest.conn.Close(); err != nil { + t.Fatalf("close guest: %v", err) + } + left := host.expect(relayTypePeerLeft) + if left.PeerID != "G" { + t.Fatalf("disconnected peerId=%q, want G", left.PeerID) + } + + thiefToken, _ := mustReconnectToken(t) + thief := h.dial(t, "6.1.0.12") + thief.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "GUEST_RECONNECT", + PeerID: "G", + ReconnectToken: thiefToken, + ProtocolVersion: relayProtocolVersion, + }) + thief.expectError(relayErrorPeerIdUnavailable) + thief.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"forged":true}`)}) + thief.expectError(relayErrorNotInRoom) + + rightful := h.dial(t, "6.1.0.13") + rightful.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "GUEST_RECONNECT", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + joined := rightful.expectAuthority(relayTypeJoined, "H") + if joined.ReconnectToken != guestToken { + t.Fatal("rightful reconnect rotated the retained guest token") + } + rejoined := host.expect(relayTypePeerJoined) + if rejoined.PeerID != "G" { + t.Fatalf("rightful reconnect event peerId=%q, want G", rejoined.PeerID) + } +} + +func TestTerminalSuccessFramesFollowCrashDurableSnapshot(t *testing.T) { + t.Run("guest leave releases persisted reservation", func(t *testing.T) { + root := t.TempDir() + statePath := filepath.Join(root, "rooms.json") + h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) + terminalReady := make(chan struct{}) + releaseTerminal := make(chan struct{}) + var releaseOnce sync.Once + h.srv.beforeTerminalDelivery = func() { + close(terminalReady) + <-releaseTerminal + } + t.Cleanup(func() { + releaseOnce.Do(func() { close(releaseTerminal) }) + }) + + host, guest, _, guestToken := createModernRoomWithGuest( + t, + h, + "DURABLE_LEAVE", + "6.1.0.20", + "6.1.0.21", + ) + guest.send(clientMsg{ + Type: relayTypeLeave, + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + select { + case <-terminalReady: + case <-time.After(2 * time.Second): + t.Fatal("leave did not reach the post-persistence delivery barrier") + } + + restartPath := copySnapshotForRestart(t, statePath) + restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) + replacementToken, _ := mustReconnectToken(t) + replacement := restarted.dial(t, "6.1.0.22") + replacement.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "DURABLE_LEAVE", + PeerID: "G", + ReconnectToken: replacementToken, + ProtocolVersion: relayProtocolVersion, + }) + replacement.expectAuthority(relayTypeJoined, "H") + + releaseOnce.Do(func() { close(releaseTerminal) }) + guest.expect(relayTypeLeft) + left := host.expect(relayTypePeerLeft) + if left.PeerID != "G" { + t.Fatalf("released peer event peerId=%q, want G", left.PeerID) + } + }) + + t.Run("host end removes persisted room", func(t *testing.T) { + root := t.TempDir() + statePath := filepath.Join(root, "rooms.json") + h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) + terminalReady := make(chan struct{}) + releaseTerminal := make(chan struct{}) + var releaseOnce sync.Once + h.srv.beforeTerminalDelivery = func() { + close(terminalReady) + <-releaseTerminal + } + t.Cleanup(func() { + releaseOnce.Do(func() { close(releaseTerminal) }) + }) + + host, guest, hostToken, _ := createModernRoomWithGuest( + t, + h, + "DURABLE_END", + "6.1.0.23", + "6.1.0.24", + ) + host.send(clientMsg{ + Type: relayTypeEndSession, + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + select { + case <-terminalReady: + case <-time.After(2 * time.Second): + t.Fatal("end did not reach the post-persistence delivery barrier") + } + + restartPath := copySnapshotForRestart(t, statePath) + restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) + probe := restarted.dial(t, "6.1.0.25") + probe.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "DURABLE_END", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + probe.expectError(relayErrorRoomNotFound) + + releaseOnce.Do(func() { close(releaseTerminal) }) + host.expect(relayTypeEnded) + messages, err := guest.recvUntilClosed(2 * time.Second) + if err != nil { + t.Fatalf("guest remained connected after durable end: %v (frames=%v)", err, messages) + } + if len(messages) != 1 || messages[0].Type != relayTypeEnded { + t.Fatalf("guest terminal frames=%+v, want one ended notification", messages) + } + }) +} + +func TestTerminalPersistenceFailureSuppressesSuccess(t *testing.T) { + injectedErr := errors.New("injected snapshot persistence failure") + + t.Run("guest leave", func(t *testing.T) { + root := t.TempDir() + statePath := filepath.Join(root, "rooms.json") + h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) + _, guest, _, guestToken := createModernRoomWithGuest( + t, + h, + "FAILED_LEAVE", + "6.1.0.26", + "6.1.0.27", + ) + injectSnapshotPersistenceFailure(t, h.srv.snap, injectedErr) + + guest.send(clientMsg{ + Type: relayTypeLeave, + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + failure := guest.expectError(relayErrorInvalidMessage) + if !strings.Contains(failure.Message, "persist") { + t.Fatalf("leave persistence error message=%q", failure.Message) + } + + restartPath := copySnapshotForRestart(t, statePath) + restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) + replacementToken, _ := mustReconnectToken(t) + replacement := restarted.dial(t, "6.1.0.28") + replacement.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FAILED_LEAVE", + PeerID: "G", + ReconnectToken: replacementToken, + ProtocolVersion: relayProtocolVersion, + }) + replacement.expectError(relayErrorPeerIdUnavailable) + }) + + t.Run("host end", func(t *testing.T) { + root := t.TempDir() + statePath := filepath.Join(root, "rooms.json") + h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) + host, guest, hostToken, _ := createModernRoomWithGuest( + t, + h, + "FAILED_END", + "6.1.0.29", + "6.1.0.30", + ) + injectSnapshotPersistenceFailure(t, h.srv.snap, injectedErr) + + host.send(clientMsg{ + Type: relayTypeEndSession, + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + failure := host.expectError(relayErrorInvalidMessage) + if !strings.Contains(failure.Message, "persist") { + t.Fatalf("end persistence error message=%q", failure.Message) + } + messages, err := guest.recvUntilClosed(2 * time.Second) + if err != nil { + t.Fatalf("guest remained connected after failed end persistence: %v (frames=%v)", err, messages) + } + if len(messages) != 0 { + t.Fatalf("guest received success frames after persistence failure: %+v", messages) + } + + restartPath := copySnapshotForRestart(t, statePath) + restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) + reconnected := restarted.dial(t, "6.1.0.31") + reconnected.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "FAILED_END", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + reconnected.expectAuthority(relayTypeJoined, "H") + }) +} + +func TestAuthenticatedModernGuestLeaveReleasesIdentity(t *testing.T) { + h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "6.1.0.20") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "GUEST_LEAVE", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, "6.1.0.21") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "GUEST_LEAVE", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + wrongToken, _ := mustReconnectToken(t) + guest.send(clientMsg{ + Type: relayTypeLeave, + ReconnectToken: wrongToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectError(relayErrorPeerIdUnavailable) + guest.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"still":"joined"}`)}) + stillJoined := host.expect(relayTypeMessage) + if stillJoined.From != "G" { + t.Fatalf("message after rejected leave came from %q, want G", stillJoined.From) + } + + guest.send(clientMsg{ + Type: relayTypeLeave, + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + leftAck := guest.expect(relayTypeLeft) + if leftAck.SessionID != "GUEST_LEAVE" || leftAck.PeerID != "G" || leftAck.ProtocolVersion != relayProtocolVersion { + t.Fatalf("left acknowledgement=%+v", leftAck) + } + leftEvent := host.expect(relayTypePeerLeft) + if leftEvent.PeerID != "G" { + t.Fatalf("released peer event peerId=%q, want G", leftEvent.PeerID) + } + + replacementToken, _ := mustReconnectToken(t) + replacement := h.dial(t, "6.1.0.22") + replacement.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "GUEST_LEAVE", + PeerID: "G", + ReconnectToken: replacementToken, + ProtocolVersion: relayProtocolVersion, + }) + joined := replacement.expectAuthority(relayTypeJoined, "H") + if joined.ReconnectToken != replacementToken { + t.Fatal("released guest identity retained its old verifier") + } + rejoined := host.expect(relayTypePeerJoined) + if rejoined.PeerID != "G" { + t.Fatalf("replacement event peerId=%q, want G", rejoined.PeerID) + } +} + +func TestAuthenticatedModernHostEndDeletesRoomAndIsRetrySafe(t *testing.T) { + h := newRelayHarness(t) + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "6.1.0.30") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "HOST_END", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, "6.1.0.31") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "HOST_END", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + wrongToken, _ := mustReconnectToken(t) + host.send(clientMsg{ + Type: relayTypeEndSession, + ReconnectToken: wrongToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectError(relayErrorPeerIdUnavailable) + host.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"room":"live"}`)}) + stillLive := guest.expect(relayTypeMessage) + if stillLive.From != "H" { + t.Fatalf("message after rejected end came from %q, want H", stillLive.From) + } + + host.send(clientMsg{ + Type: relayTypeEndSession, + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + ended := host.expect(relayTypeEnded) + if ended.SessionID != "HOST_END" || ended.ProtocolVersion != relayProtocolVersion { + t.Fatalf("ended acknowledgement=%+v", ended) + } + messages, err := guest.recvUntilClosed(2 * time.Second) + if err != nil { + t.Fatalf("guest remained connected after room end: %v (frames=%v)", err, messages) + } + if len(messages) != 1 || + messages[0].Type != relayTypeEnded || + messages[0].SessionID != "HOST_END" || + messages[0].ProtocolVersion != relayProtocolVersion { + t.Fatalf("guest terminal frames=%+v, want one ended notification", messages) + } + + host.send(clientMsg{ + Type: relayTypeEndSession, + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectError(relayErrorNotInRoom) + + retry := h.dial(t, "6.1.0.32") + retry.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "HOST_END", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + retry.expectError(relayErrorRoomNotFound) +} + +func TestHostEndDeliversEndedAfterConcurrentGuestTraffic(t *testing.T) { + h := newRelayHarness(t) + endDeliveryReady := make(chan struct{}) + releaseEndDelivery := make(chan struct{}) + var releaseOnce sync.Once + h.srv.beforeTerminalDelivery = func() { + close(endDeliveryReady) + <-releaseEndDelivery + } + t.Cleanup(func() { + releaseOnce.Do(func() { close(releaseEndDelivery) }) + }) + + hostToken, _ := mustReconnectToken(t) + host := h.dial(t, "6.1.0.33") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "END_RACE", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := h.dial(t, "6.1.0.34") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "END_RACE", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + host.send(clientMsg{ + Type: relayTypeEndSession, + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + select { + case <-endDeliveryReady: + case <-time.After(2 * time.Second): + t.Fatal("host end did not reach the terminal-delivery barrier") + } + + h.srv.mu.RLock() + _, discoverable := h.srv.rooms["END_RACE"] + h.srv.mu.RUnlock() + if discoverable { + t.Fatal("ending room remained discoverable before terminal delivery") + } + + // WebSocket frames are processed in order. Receiving pong proves the + // preceding membership-sensitive traffic was handled while ended delivery + // was blocked, without closing the guest as a stale client. + guest.send(clientMsg{ + Type: relayTypeBroadcast, + Payload: json.RawMessage(`{"during":"end"}`), + }) + guest.send(clientMsg{ + Type: relayTypeSendTo, + To: "H", + Payload: json.RawMessage(`{"also":"during-end"}`), + }) + guest.send(clientMsg{Type: relayTypePing}) + guest.expect(relayTypePong) + + releaseOnce.Do(func() { close(releaseEndDelivery) }) + endedAck := host.expect(relayTypeEnded) + if endedAck.SessionID != "END_RACE" || endedAck.ProtocolVersion != relayProtocolVersion { + t.Fatalf("host ended acknowledgement=%+v", endedAck) + } + messages, err := guest.recvUntilClosed(2 * time.Second) + if err != nil { + t.Fatalf("guest remained connected after terminal delivery: %v (frames=%v)", err, messages) + } + if len(messages) != 1 || + messages[0].Type != relayTypeEnded || + messages[0].SessionID != "END_RACE" || + messages[0].ProtocolVersion != relayProtocolVersion { + t.Fatalf("guest terminal frames=%+v, want one ended notification", messages) + } } func TestHostReconnectReplacesStaleConnectionWithoutLeaving(t *testing.T) { h := newRelayHarness(t) oldHostIP := "6.1.1.1" oldHost := h.dial(t, oldHostIP) - oldHost.send(clientMsg{Type: "create", SessionID: "REJOIN", PeerID: "H"}) - oldHost.expect("created") + oldHost.send(clientMsg{Type: relayTypeCreate, SessionID: "REJOIN", PeerID: "H"}) + created := oldHost.expectAuthority(relayTypeCreated, "H") guest := h.dial(t, "6.1.1.2") - guest.send(clientMsg{Type: "join", SessionID: "REJOIN", PeerID: "G"}) - guest.expect("joined") - oldHost.expect("peerJoined") + guest.send(clientMsg{Type: relayTypeJoin, SessionID: "REJOIN", PeerID: "G"}) + guest.expectAuthority(relayTypeJoined, "H") + oldHost.expect(relayTypePeerJoined) newHost := h.dial(t, "6.1.1.3") - newHost.send(clientMsg{Type: "join", SessionID: "REJOIN", PeerID: "H"}) - joined := newHost.expect("joined") + newHost.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "REJOIN", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + }) + joined := newHost.expectAuthority(relayTypeJoined, "H") if len(joined.Peers) != 1 || joined.Peers[0] != "G" { t.Fatalf("reconnected host peers=%v, want [G]", joined.Peers) } - guest.expect("peerJoined") + guest.expect(relayTypePeerJoined) - oldHost.conn.Close() + if messages, err := oldHost.recvUntilClosed(2 * time.Second); err != nil { + t.Fatalf("displaced host did not close: %v (frames=%v)", err, messages) + } h.waitIPConnections(t, oldHostIP, 0) - newHost.send(clientMsg{Type: "broadcast", Payload: json.RawMessage(`{"state":"ready"}`)}) - message := guest.expect("message") + newHost.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"state":"ready"}`)}) + message := guest.expect(relayTypeMessage) if message.From != "H" { t.Fatalf("message sender=%q, want H", message.From) } } -func TestEmptyRoomSupportsJoinThenExpiresForReconnectFallback(t *testing.T) { +func TestEmptyRoomRequiresHostProofUntilExpiryThenSupportsFallbackCreate(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "6.1.2.1") - host.send(clientMsg{Type: "create", SessionID: "EMPTY", PeerID: "H"}) - host.expect("created") + host.send(clientMsg{Type: relayTypeCreate, SessionID: "EMPTY", PeerID: "H"}) + created := host.expectAuthority(relayTypeCreated, "H") host.conn.Close() h.waitRoomPeers(t, "EMPTY", 0) + unproved := h.dial(t, "6.1.2.20") + unproved.send(clientMsg{Type: relayTypeJoin, SessionID: "EMPTY", PeerID: "H"}) + unproved.expectError(relayErrorPeerIdUnavailable) + reconnected := h.dial(t, "6.1.2.2") - reconnected.send(clientMsg{Type: "join", SessionID: "EMPTY", PeerID: "H"}) - joined := reconnected.expect("joined") + reconnected.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "EMPTY", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + }) + joined := reconnected.expectAuthority(relayTypeJoined, "H") + if joined.ReconnectToken != created.ReconnectToken { + t.Fatal("host reconnect rotated its capability") + } if len(joined.Peers) != 0 { t.Fatalf("empty-room reconnect peers=%v, want none", joined.Peers) } @@ -1256,8 +3571,23 @@ func TestEmptyRoomSupportsJoinThenExpiresForReconnectFallback(t *testing.T) { h.srv.runCleanupStep(now) fallback := h.dial(t, "6.1.2.3") - fallback.send(clientMsg{Type: "join", SessionID: "EMPTY", PeerID: "H"}) - fallback.expectError("room_not_found") + fallback.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "EMPTY", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + }) + fallback.expectError(relayErrorRoomNotFound) + fallback.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "EMPTY", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + }) + recreated := fallback.expectAuthority(relayTypeCreated, "H") + if recreated.ReconnectToken != created.ReconnectToken { + t.Fatal("fallback create rotated its retained capability") + } } func TestCleanupDisconnectsPeersBeforeRemovingExpiredOccupiedRoom(t *testing.T) { @@ -1288,10 +3618,27 @@ func TestCleanupDisconnectsPeersBeforeRemovingExpiredOccupiedRoom(t *testing.T) t.Fatal("expired room still exists after cleanup") } + room.mu.RLock() + closing := room.closing + remainingPeers := len(room.Peers) + room.mu.RUnlock() + if !closing { + t.Error("expired room was not marked closing") + } + if remainingPeers != 0 { + t.Errorf("expired room retained %d peers after cleanup", remainingPeers) + } + for name, connection := range map[string]*testConn{"host": host, "guest": guest} { - if _, err := connection.recvUntilClosed(2 * time.Second); err != nil { + messages, err := connection.recvUntilClosed(2 * time.Second) + if err != nil { t.Errorf("%s did not reach terminal closure: %v", name, err) } + for _, message := range messages { + if message.Type == relayTypePeerLeft || message.Type == relayTypePeerJoined { + t.Errorf("%s received %s during expired-room teardown", name, message.Type) + } + } } } @@ -1315,6 +3662,22 @@ func postLog(t *testing.T, baseURL, ip string, body []byte) *http.Response { return resp } +func getLog(t *testing.T, baseURL, ip, id string) *http.Response { + t.Helper() + req, err := http.NewRequest(http.MethodGet, baseURL+"/logs/"+id, nil) + if err != nil { + t.Fatalf("new request: %v", err) + } + if ip != "" { + req.Header.Set("X-Forwarded-For", ip) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("get: %v", err) + } + return resp +} + // postLogAndGetID uploads a log and returns the generated id, asserting the // POST succeeded. func postLogAndGetID(t *testing.T, baseURL, ip string, body []byte) string { @@ -1378,6 +3741,132 @@ func postPosterAndDecode(t *testing.T, baseURL, ip string, body []byte) posterUp return out } +type blockingLogResponseWriter struct { + header http.Header + writeStarted chan struct{} + releaseWrite chan struct{} + writeErr error + startOnce sync.Once + mu sync.Mutex + deadlines []time.Time +} + +func newBlockingLogResponseWriter(writeErr error) *blockingLogResponseWriter { + return &blockingLogResponseWriter{ + header: make(http.Header), + writeStarted: make(chan struct{}), + releaseWrite: make(chan struct{}), + writeErr: writeErr, + } +} + +func (w *blockingLogResponseWriter) Header() http.Header { + return w.header +} + +func (w *blockingLogResponseWriter) WriteHeader(int) {} + +func (w *blockingLogResponseWriter) Write(p []byte) (int, error) { + w.startOnce.Do(func() { close(w.writeStarted) }) + <-w.releaseWrite + if w.writeErr != nil { + return 0, w.writeErr + } + return len(p), nil +} + +func (w *blockingLogResponseWriter) SetWriteDeadline(deadline time.Time) error { + w.mu.Lock() + w.deadlines = append(w.deadlines, deadline) + w.mu.Unlock() + return nil +} + +func TestLogResponseTransmissionReleasesLookupSlotAndUsesDeadline(t *testing.T) { + logs := newLogStore(t.TempDir()) + id, _, err := logs.store([]byte("diagnostic"), time.Now()) + if err != nil { + t.Fatalf("store log: %v", err) + } + srv := &Server{ + logs: logs, + logLookups: make(chan struct{}, 1), + clientIPs: newClientIPResolver(nil), + } + writer := newBlockingLogResponseWriter(errors.New("synthetic write failure: " + id)) + request := httptest.NewRequest(http.MethodGet, "/logs/"+id, nil) + + var output bytes.Buffer + previousOutput := log.Writer() + previousFlags := log.Flags() + previousPrefix := log.Prefix() + log.SetOutput(&output) + log.SetFlags(0) + log.SetPrefix("") + t.Cleanup(func() { + log.SetOutput(previousOutput) + log.SetFlags(previousFlags) + log.SetPrefix(previousPrefix) + }) + + done := make(chan struct{}) + go func() { + srv.handleGetLogs(writer, request) + close(done) + }() + select { + case <-writer.writeStarted: + case <-time.After(time.Second): + t.Fatal("log response did not reach blocked write") + } + + if occupied := len(srv.logLookups); occupied != 0 { + t.Fatalf("blocked response retained %d lookup slots", occupied) + } + second := httptest.NewRecorder() + srv.handleGetLogs(second, httptest.NewRequest(http.MethodGet, "/logs/"+id, nil)) + if second.Code != http.StatusOK { + t.Fatalf("lookup while first response blocked status=%d, want 200", second.Code) + } + + writer.mu.Lock() + deadlines := append([]time.Time(nil), writer.deadlines...) + writer.mu.Unlock() + if len(deadlines) == 0 || deadlines[0].IsZero() { + t.Fatalf("response write deadline not applied: %v", deadlines) + } + remaining := time.Until(deadlines[0]) + if remaining <= 0 || remaining > httpResponseWriteTimeout { + t.Fatalf("response write deadline remaining=%v, want within (0, %v]", remaining, httpResponseWriteTimeout) + } + + close(writer.releaseWrite) + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("log handler did not finish after writer release") + } + if !strings.Contains(output.String(), "logs: response write failed") { + t.Fatalf("write failure was not observed: %q", output.String()) + } + if strings.Contains(output.String(), id) { + t.Fatalf("write failure leaked log capability %q", id) + } +} + +func TestHTTPServerWriteTimeoutCoversOAuthResultLongPoll(t *testing.T) { + server := newHTTPServer("127.0.0.1:0", http.NewServeMux()) + if server.WriteTimeout != httpResponseWriteTimeout || server.WriteTimeout <= 0 { + t.Fatalf("WriteTimeout=%v, want bounded timeout %v", server.WriteTimeout, httpResponseWriteTimeout) + } + if server.WriteTimeout <= oauthResultWait { + t.Fatalf("WriteTimeout=%v must exceed oauthResultWait=%v", server.WriteTimeout, oauthResultWait) + } + if margin := server.WriteTimeout - oauthResultWait; margin < httpResponseWriteMargin { + t.Fatalf("WriteTimeout margin=%v, want at least %v", margin, httpResponseWriteMargin) + } +} + func TestLogsRoundTrip(t *testing.T) { h := newRelayHarness(t) payload := []byte("hello log world") @@ -1400,22 +3889,154 @@ func TestLogsRoundTrip(t *testing.T) { } } +func TestLogsUploadDoesNotWriteCapabilityToOperationalLog(t *testing.T) { + h := newRelayHarness(t) + var output bytes.Buffer + previousOutput := log.Writer() + previousFlags := log.Flags() + previousPrefix := log.Prefix() + log.SetOutput(&output) + log.SetFlags(0) + log.SetPrefix("") + t.Cleanup(func() { + log.SetOutput(previousOutput) + log.SetFlags(previousFlags) + log.SetPrefix(previousPrefix) + }) + + id := postLogAndGetID(t, h.baseURL, "203.0.113.40", []byte("safe diagnostic")) + if strings.Contains(output.String(), id) { + t.Fatalf("operational log retained bearer capability %q", id) + } + if !strings.Contains(output.String(), "logs: stored 15 bytes from 203.0.113.40") { + t.Fatalf("successful upload was not observable: %q", output.String()) + } +} + +func TestLogStoreRetiresLegacyCapabilitiesOnStartup(t *testing.T) { + dir := t.TempDir() + now := time.Now().Add(-time.Minute) + legacyID := "abcde" + currentID := strings.Repeat("a", logIDLength) + legacyPath := filepath.Join(dir, legacyID+".log") + currentPath := filepath.Join(dir, currentID+".log") + for path, body := range map[string]string{legacyPath: "legacy", currentPath: "current"} { + if err := os.WriteFile(path, []byte(body), 0o644); err != nil { + t.Fatalf("seed %s: %v", path, err) + } + if err := os.Chtimes(path, now, now); err != nil { + t.Fatalf("chtimes %s: %v", path, err) + } + } + + store := newLogStore(dir) + if _, err := os.Stat(legacyPath); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("legacy capability file still exists: %v", err) + } + if _, ok, err := store.lookup(legacyID, time.Now()); err != nil || ok { + t.Fatalf("legacy capability lookup=(ok=%v, err=%v), want absent", ok, err) + } + if _, ok, err := store.lookup(currentID, time.Now()); err != nil || !ok { + t.Fatalf("current capability lookup=(ok=%v, err=%v), want indexed", ok, err) + } + + restarted := newLogStore(dir) + if _, ok, err := restarted.lookup(currentID, time.Now()); err != nil || !ok { + t.Fatalf("restarted capability lookup=(ok=%v, err=%v), want indexed", ok, err) + } +} + +func TestLogsFailedLookupsAreBoundedButValidCapabilitiesRemainAvailable(t *testing.T) { + h := newRelayHarness(t) + payload := []byte("retrievable") + validID := postLogAndGetID(t, h.baseURL, "203.0.113.1", payload) + source := "203.0.113.50" + + for i := range logLookupRateBurst { + unknownID := strings.Repeat("z", logIDLength-2) + fmt.Sprintf("%02d", i) + resp := getLog(t, h.baseURL, source, unknownID) + resp.Body.Close() + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("failed lookup %d status=%d, want 404", i, resp.StatusCode) + } + if got := resp.Header.Get("Cache-Control"); got != "private, no-store" { + t.Fatalf("failed lookup Cache-Control=%q", got) + } + } + throttled := getLog(t, h.baseURL, source, strings.Repeat("y", logIDLength)) + throttled.Body.Close() + if throttled.StatusCode != http.StatusTooManyRequests { + t.Fatalf("exhausted lookup status=%d, want 429", throttled.StatusCode) + } + if got := throttled.Header.Get("Cache-Control"); got != "private, no-store" { + t.Fatalf("throttled Cache-Control=%q", got) + } + + success := getLog(t, h.baseURL, source, validID) + defer success.Body.Close() + if success.StatusCode != http.StatusOK { + t.Fatalf("valid capability after exhausted failures status=%d", success.StatusCode) + } + got, err := io.ReadAll(success.Body) + if err != nil || !bytes.Equal(got, payload) { + t.Fatalf("valid body=%q err=%v, want %q", got, err, payload) + } + if cache := success.Header.Get("Cache-Control"); cache != "private, no-store" { + t.Fatalf("success Cache-Control=%q", cache) + } + + independent := getLog(t, h.baseURL, "203.0.113.51", strings.Repeat("x", logIDLength)) + independent.Body.Close() + if independent.StatusCode != http.StatusNotFound { + t.Fatalf("independent source status=%d, want 404", independent.StatusCode) + } +} + +func TestLogFailedLookupCleanupIsDeterministic(t *testing.T) { + store := newLogStore(t.TempDir()) + now := time.Unix(1_700_000_000, 0) + id, _, err := store.store([]byte("keep"), now) + if err != nil { + t.Fatalf("store: %v", err) + } + for range logLookupRateBurst { + if !store.allowFailedLookup("203.0.113.1", now) { + t.Fatal("burst rejected early") + } + } + if store.allowFailedLookup("203.0.113.1", now) { + t.Fatal("lookup beyond burst unexpectedly allowed") + } + store.cleanup(now) + if _, ok := store.failedLookupRate["203.0.113.1"]; !ok { + t.Fatal("cleanup removed an effective limiter") + } + store.cleanup(now.Add(time.Duration(logLookupRateBurst) * time.Second)) + if _, ok := store.failedLookupRate["203.0.113.1"]; ok { + t.Fatal("cleanup retained a fully refilled limiter") + } + if _, ok := store.entries[id]; !ok { + t.Fatal("limiter cleanup removed stored log") + } +} + func TestLogStorePersistsAcrossRestartAndAvoidsIDCollisions(t *testing.T) { dir := t.TempDir() now := time.Now().Add(-time.Second) first := newLogStore(dir) - first.generateID = func() string { return "aaaaa" } + first.generateID = func() string { return strings.Repeat("a", logIDLength) } firstID, _, err := first.store([]byte("original"), now) if err != nil { t.Fatalf("store original: %v", err) } restarted := newLogStore(dir) - if _, ok := restarted.lookup(firstID, time.Now()); !ok { + if _, ok, err := restarted.lookup(firstID, time.Now()); err != nil || !ok { t.Fatal("stored log was not restored after restart") } - ids := []string{firstID, "bbbbb"} + secondWant := strings.Repeat("b", logIDLength) + ids := []string{firstID, secondWant} restarted.generateID = func() string { id := ids[0] ids = ids[1:] @@ -1425,8 +4046,8 @@ func TestLogStorePersistsAcrossRestartAndAvoidsIDCollisions(t *testing.T) { if err != nil { t.Fatalf("store after restart: %v", err) } - if secondID != "bbbbb" { - t.Fatalf("collision generated id %q, want bbbbb", secondID) + if secondID != secondWant { + t.Fatalf("collision generated id %q, want %q", secondID, secondWant) } original, err := os.ReadFile(restarted.filePath(firstID)) @@ -1452,6 +4073,62 @@ func TestLogsUploadRateLimitedPerIP(t *testing.T) { } } +func TestLogsUseTrustedCanonicalClientIdentity(t *testing.T) { + t.Run("untrusted spoofing shares direct peer bucket", func(t *testing.T) { + h := newRelayHarnessNoTrust(t) + first := postLog(t, h.baseURL, "203.0.113.1", []byte("first")) + first.Body.Close() + if first.StatusCode != http.StatusOK { + t.Fatalf("first status=%d", first.StatusCode) + } + second := postLog(t, h.baseURL, "203.0.113.2", []byte("second")) + second.Body.Close() + if second.StatusCode != http.StatusTooManyRequests { + t.Fatalf("rotated spoof status=%d, want 429", second.StatusCode) + } + }) + + t.Run("trusted clients have independent buckets", func(t *testing.T) { + h := newRelayHarness(t) + for _, ip := range []string{"203.0.113.1", "203.0.113.2"} { + resp := postLog(t, h.baseURL, ip, []byte(ip)) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("client %s status=%d, want 200", ip, resp.StatusCode) + } + } + }) + + t.Run("malformed trusted chain mutates no log state", func(t *testing.T) { + h := newRelayHarness(t) + post := postLog(t, h.baseURL, "203.0.113.1,", []byte("body")) + post.Body.Close() + if post.StatusCode != http.StatusBadRequest { + t.Fatalf("post status=%d, want 400", post.StatusCode) + } + get := getLog(t, h.baseURL, "203.0.113.1,", strings.Repeat("a", logIDLength)) + get.Body.Close() + if get.StatusCode != http.StatusBadRequest { + t.Fatalf("get status=%d, want 400", get.StatusCode) + } + h.srv.logs.mu.RLock() + defer h.srv.logs.mu.RUnlock() + if len(h.srv.logs.entries) != 0 || len(h.srv.logs.rateLimit) != 0 || len(h.srv.logs.failedLookupRate) != 0 { + t.Fatalf("malformed chain mutated log state: entries=%d uploads=%d failures=%d", + len(h.srv.logs.entries), len(h.srv.logs.rateLimit), len(h.srv.logs.failedLookupRate)) + } + }) + + t.Run("untrusted malformed header is ignored", func(t *testing.T) { + h := newRelayHarnessNoTrust(t) + resp := postLog(t, h.baseURL, "bad,", []byte("body")) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("status=%d, want 200", resp.StatusCode) + } + }) +} + func TestLogsUploadTooLargeRejected(t *testing.T) { h := newRelayHarness(t) body := make([]byte, maxLogSize+1) @@ -1491,7 +4168,7 @@ func TestLogsUploadStoreFull(t *testing.T) { func TestLogsGetUnknownIDIs404(t *testing.T) { h := newRelayHarness(t) - resp, err := http.Get(h.baseURL + "/logs/abcde") + resp, err := http.Get(h.baseURL + "/logs/" + strings.Repeat("c", logIDLength)) if err != nil { t.Fatalf("get: %v", err) } @@ -1503,7 +4180,7 @@ func TestLogsGetUnknownIDIs404(t *testing.T) { func TestLogsGetMalformedIDIs404(t *testing.T) { h := newRelayHarness(t) - for _, id := range []string{"", "abc", "toolongid"} { + for _, id := range []string{"", "abc", "abcde", strings.Repeat("a", logIDLength+1), strings.Repeat("!", logIDLength)} { resp, err := http.Get(h.baseURL + "/logs/" + id) if err != nil { t.Fatalf("get %q: %v", id, err) @@ -1512,6 +4189,9 @@ func TestLogsGetMalformedIDIs404(t *testing.T) { if resp.StatusCode != http.StatusNotFound { t.Errorf("id=%q status=%d want 404", id, resp.StatusCode) } + if got := resp.Header.Get("Cache-Control"); got != "private, no-store" { + t.Errorf("id=%q Cache-Control=%q", id, got) + } } } @@ -1552,6 +4232,410 @@ func TestLogsMethodNotAllowed(t *testing.T) { // Poster endpoints // ====================================================================== +var minimalPNG = []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a, 0x01, 0x02, 0x03} + +type countingReadCloser struct { + reader *bytes.Reader + reads atomic.Int32 +} + +func newCountingReadCloser(data []byte) *countingReadCloser { + return &countingReadCloser{reader: bytes.NewReader(data)} +} + +func (r *countingReadCloser) Read(p []byte) (int, error) { + r.reads.Add(1) + return r.reader.Read(p) +} + +func (r *countingReadCloser) Close() error { return nil } + +type blockingReadCloser struct { + data []byte + offset int + started chan struct{} + release <-chan struct{} + once sync.Once +} + +func (r *blockingReadCloser) Read(p []byte) (int, error) { + r.once.Do(func() { close(r.started) }) + <-r.release + if r.offset == len(r.data) { + return 0, io.EOF + } + n := copy(p, r.data[r.offset:]) + r.offset += n + return n, nil +} + +func (r *blockingReadCloser) Close() error { return nil } + +type deadlineBlockingReadCloser struct { + started chan struct{} + closed chan struct{} + startOnce sync.Once + closeOnce sync.Once +} + +func newDeadlineBlockingReadCloser() *deadlineBlockingReadCloser { + return &deadlineBlockingReadCloser{ + started: make(chan struct{}), + closed: make(chan struct{}), + } +} + +func (r *deadlineBlockingReadCloser) Read([]byte) (int, error) { + r.startOnce.Do(func() { close(r.started) }) + <-r.closed + return 0, errors.New("body closed") +} + +func (r *deadlineBlockingReadCloser) Close() error { + r.closeOnce.Do(func() { close(r.closed) }) + return nil +} + +type failingReadCloser struct{} + +func (failingReadCloser) Read([]byte) (int, error) { return 0, errors.New("read failed") } +func (failingReadCloser) Close() error { return nil } + +func servePosterUpload(s *Server, body io.ReadCloser, xff string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, "/posters", body) + req.RemoteAddr = "198.51.100.10:1234" + if xff != "" { + req.Header.Set("X-Forwarded-For", xff) + } + recorder := httptest.NewRecorder() + s.handlePostPosters(recorder, req) + return recorder +} + +func snapshotPosterStore(t *testing.T, store *posterStore) (int, int64, []string) { + t.Helper() + store.mu.RLock() + entryCount := len(store.entries) + totalBytes := store.totalBytes + store.mu.RUnlock() + files, err := os.ReadDir(store.dir) + if err != nil { + t.Fatalf("read poster dir: %v", err) + } + names := make([]string, len(files)) + for i, file := range files { + names[i] = file.Name() + } + return entryCount, totalBytes, names +} + +func TestPosterHandlerRejectsRateLimitedRequestBeforeReadingOrStoring(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + now := time.Now() + s.posterUploads = newPosterUploadLimiter(1, 0, 10, 0, 2, now) + + first := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "203.0.113.1") + if first.Code != http.StatusOK { + t.Fatalf("first status=%d", first.Code) + } + beforeEntries, beforeBytes, beforeFiles := snapshotPosterStore(t, s.posters) + rejectedBody := newCountingReadCloser(minimalPNG) + rejected := servePosterUpload(s, rejectedBody, "203.0.113.2") + if rejected.Code != http.StatusTooManyRequests { + t.Fatalf("rejected status=%d, want 429", rejected.Code) + } + if rejectedBody.reads.Load() != 0 { + t.Fatalf("rate-limited body read %d times", rejectedBody.reads.Load()) + } + afterEntries, afterBytes, afterFiles := snapshotPosterStore(t, s.posters) + if beforeEntries != afterEntries || beforeBytes != afterBytes || fmt.Sprint(beforeFiles) != fmt.Sprint(afterFiles) { + t.Fatalf("denial mutated poster store: before=(%d,%d,%v) after=(%d,%d,%v)", + beforeEntries, beforeBytes, beforeFiles, afterEntries, afterBytes, afterFiles) + } +} + +func TestPosterHandlerConcurrencyRejectsBeforeReadAndRecovers(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + s.posterUploads = newPosterUploadLimiter(20, 0, 20, 0, 2, time.Now()) + release := make(chan struct{}) + recorders := make(chan *httptest.ResponseRecorder, 2) + + for range 2 { + body := &blockingReadCloser{ + data: minimalPNG, + started: make(chan struct{}), + release: release, + } + go func() { + recorders <- servePosterUpload(s, body, "") + }() + select { + case <-body.started: + case <-time.After(time.Second): + t.Fatal("admitted body was not read") + } + } + + extraBody := newCountingReadCloser(minimalPNG) + extra := servePosterUpload(s, extraBody, "") + if extra.Code != http.StatusTooManyRequests { + t.Fatalf("extra status=%d, want 429", extra.Code) + } + if extraBody.reads.Load() != 0 { + t.Fatalf("concurrency-rejected body read %d times", extraBody.reads.Load()) + } + + close(release) + for range 2 { + select { + case recorder := <-recorders: + if recorder.Code != http.StatusOK { + t.Fatalf("admitted status=%d", recorder.Code) + } + case <-time.After(time.Second): + t.Fatal("admitted upload did not finish") + } + } + recovered := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "") + if recovered.Code != http.StatusOK { + t.Fatalf("post-completion status=%d, want 200", recovered.Code) + } + s.posterUploads.mu.Lock() + active := s.posterUploads.active + s.posterUploads.mu.Unlock() + if active != 0 { + t.Fatalf("active=%d after completion, want 0", active) + } +} + +func TestPosterHandlerDeadlineReleasesStalledUploadSlot(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + s.posterUploads = newPosterUploadLimiter(10, 0, 10, 0, 1, time.Now()) + s.posterBodyReadTimeout = 20 * time.Millisecond + stalled := newDeadlineBlockingReadCloser() + result := make(chan *httptest.ResponseRecorder, 1) + + go func() { + result <- servePosterUpload(s, stalled, "") + }() + select { + case <-stalled.started: + case <-time.After(time.Second): + t.Fatal("stalled body was not read") + } + + var timedOut *httptest.ResponseRecorder + select { + case timedOut = <-result: + case <-time.After(time.Second): + t.Fatal("stalled upload did not honor body deadline") + } + if timedOut.Code != http.StatusRequestTimeout { + t.Fatalf("stalled status=%d, want 408", timedOut.Code) + } + s.posterUploads.mu.Lock() + active := s.posterUploads.active + s.posterUploads.mu.Unlock() + if active != 0 { + t.Fatalf("active=%d after body timeout, want 0", active) + } + + recovered := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "") + if recovered.Code != http.StatusOK { + t.Fatalf("post-timeout status=%d, want 200", recovered.Code) + } +} + +func TestPosterHandlerSlowChunkedBodyDeadline(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + s.posterUploads = newPosterUploadLimiter(10, 0, 10, 0, 1, time.Now()) + s.posterBodyReadTimeout = 30 * time.Millisecond + httpServer := httptest.NewServer(http.HandlerFunc(s.handlePostPosters)) + t.Cleanup(httpServer.Close) + + address := strings.TrimPrefix(httpServer.URL, "http://") + conn, err := net.DialTimeout("tcp", address, time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + if _, err := fmt.Fprintf( + conn, + "POST /posters HTTP/1.1\r\nHost: %s\r\nTransfer-Encoding: chunked\r\n\r\n1\r\nx\r\n", + address, + ); err != nil { + conn.Close() + t.Fatalf("write partial chunked request: %v", err) + } + if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + conn.Close() + t.Fatalf("set response deadline: %v", err) + } + response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodPost}) + if err != nil { + conn.Close() + t.Fatalf("read timeout response: %v", err) + } + response.Body.Close() + conn.Close() + if response.StatusCode != http.StatusRequestTimeout { + t.Fatalf("slow chunked status=%d, want 408", response.StatusCode) + } + + recovered, err := http.Post(httpServer.URL, "image/png", bytes.NewReader(minimalPNG)) + if err != nil { + t.Fatalf("post after timeout: %v", err) + } + recovered.Body.Close() + if recovered.StatusCode != http.StatusOK { + t.Fatalf("post-timeout status=%d, want 200", recovered.StatusCode) + } +} + +func TestPosterHandlerReleasesConcurrencyOnEveryExit(t *testing.T) { + tests := []struct { + name string + body func() io.ReadCloser + wantStatus int + storeFailure bool + }{ + {name: "read error", body: func() io.ReadCloser { return failingReadCloser{} }, wantStatus: http.StatusBadRequest}, + {name: "oversized", body: func() io.ReadCloser { + return io.NopCloser(bytes.NewReader(make([]byte, maxPosterSize+1))) + }, wantStatus: http.StatusRequestEntityTooLarge}, + {name: "empty", body: func() io.ReadCloser { return io.NopCloser(bytes.NewReader(nil)) }, wantStatus: http.StatusBadRequest}, + {name: "unsupported", body: func() io.ReadCloser { + return io.NopCloser(strings.NewReader("not an image")) + }, wantStatus: http.StatusUnsupportedMediaType}, + {name: "store failure", body: func() io.ReadCloser { + return io.NopCloser(bytes.NewReader(minimalPNG)) + }, wantStatus: http.StatusInternalServerError, storeFailure: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + s.posterUploads = newPosterUploadLimiter(10, 0, 10, 0, 1, time.Now()) + originalDir := s.posters.dir + if tt.storeFailure { + s.posters.dir = filepath.Join(t.TempDir(), "missing", "posters") + } + failed := servePosterUpload(s, tt.body(), "") + if failed.Code != tt.wantStatus { + t.Fatalf("status=%d, want %d", failed.Code, tt.wantStatus) + } + s.posters.dir = originalDir + recovery := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "") + if recovery.Code != http.StatusOK { + t.Fatalf("recovery status=%d, want 200", recovery.Code) + } + s.posterUploads.mu.Lock() + active := s.posterUploads.active + s.posterUploads.mu.Unlock() + if active != 0 { + t.Fatalf("active=%d, want 0", active) + } + }) + } +} + +func TestPosterHandlerUsesTrustedCanonicalIdentityAndGlobalBudget(t *testing.T) { + t.Run("untrusted XFF rotation cannot bypass per-IP limit", func(t *testing.T) { + h := newRelayHarnessNoTrust(t) + h.srv.posterUploads = newPosterUploadLimiter(3, 0, 20, 0, 4, time.Now()) + for i := range 3 { + resp := postPoster(t, h.baseURL, fmt.Sprintf("203.0.113.%d", i+1), minimalPNG) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("upload %d status=%d", i, resp.StatusCode) + } + } + beforeEntries, beforeBytes, beforeFiles := snapshotPosterStore(t, h.srv.posters) + denied := postPoster(t, h.baseURL, "203.0.113.99", minimalPNG) + denied.Body.Close() + if denied.StatusCode != http.StatusTooManyRequests { + t.Fatalf("rotated spoof status=%d, want 429", denied.StatusCode) + } + afterEntries, afterBytes, afterFiles := snapshotPosterStore(t, h.srv.posters) + if beforeEntries != afterEntries || beforeBytes != afterBytes || fmt.Sprint(beforeFiles) != fmt.Sprint(afterFiles) { + t.Fatal("per-IP denial mutated poster store") + } + }) + + t.Run("trusted clients are independent but share global budget", func(t *testing.T) { + h := newRelayHarness(t) + h.srv.posterUploads = newPosterUploadLimiter(3, 0, 8, 0, 4, time.Now()) + for range 3 { + resp := postPoster(t, h.baseURL, "203.0.113.1", minimalPNG) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("client A status=%d", resp.StatusCode) + } + } + perIPDenied := postPoster(t, h.baseURL, "203.0.113.1", minimalPNG) + perIPDenied.Body.Close() + if perIPDenied.StatusCode != http.StatusTooManyRequests { + t.Fatalf("client A overflow status=%d, want 429", perIPDenied.StatusCode) + } + for _, ip := range []string{"203.0.113.2", "203.0.113.2", "203.0.113.2", "203.0.113.3", "203.0.113.3"} { + resp := postPoster(t, h.baseURL, ip, minimalPNG) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("client %s status=%d before global exhaustion", ip, resp.StatusCode) + } + } + beforeEntries, beforeBytes, beforeFiles := snapshotPosterStore(t, h.srv.posters) + globalDenied := postPoster(t, h.baseURL, "203.0.113.4", minimalPNG) + globalDenied.Body.Close() + if globalDenied.StatusCode != http.StatusTooManyRequests { + t.Fatalf("global overflow status=%d, want 429", globalDenied.StatusCode) + } + afterEntries, afterBytes, afterFiles := snapshotPosterStore(t, h.srv.posters) + if beforeEntries != afterEntries || beforeBytes != afterBytes || fmt.Sprint(beforeFiles) != fmt.Sprint(afterFiles) { + t.Fatal("global denial mutated poster store") + } + }) +} + +func TestPosterHandlerMalformedTrustedChainMutatesNothing(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + s.clientIPs = mustClientIPResolver(t, "10.0.0.0/8") + body := newCountingReadCloser(minimalPNG) + req := httptest.NewRequest(http.MethodPost, "/posters", body) + req.RemoteAddr = "10.0.0.2:1234" + req.Header.Set("X-Forwarded-For", "203.0.113.1,") + recorder := httptest.NewRecorder() + s.handlePostPosters(recorder, req) + if recorder.Code != http.StatusBadRequest { + t.Fatalf("status=%d, want 400", recorder.Code) + } + if body.reads.Load() != 0 { + t.Fatalf("malformed-chain body read %d times", body.reads.Load()) + } + s.posterUploads.mu.Lock() + active := s.posterUploads.active + perIP := len(s.posterUploads.perIP) + s.posterUploads.global.mu.Lock() + globalTokens := s.posterUploads.global.tokens + s.posterUploads.global.mu.Unlock() + s.posterUploads.mu.Unlock() + if active != 0 || perIP != 0 || globalTokens != posterGlobalRateBurst { + t.Fatalf("malformed chain mutated admission: active=%d perIP=%d global=%v", active, perIP, globalTokens) + } + entries, total, files := snapshotPosterStore(t, s.posters) + if entries != 0 || total != 0 || len(files) != 0 { + t.Fatalf("malformed chain mutated store: entries=%d total=%d files=%v", entries, total, files) + } +} + +func TestPosterHandlerIgnoresMalformedHeaderFromUntrustedPeer(t *testing.T) { + h := newRelayHarnessNoTrust(t) + resp := postPoster(t, h.baseURL, "bad,", minimalPNG) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("status=%d, want 200", resp.StatusCode) + } +} + func TestPostersRoundTrip(t *testing.T) { h := newRelayHarness(t) payload := []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a, 0x01, 0x02, 0x03} @@ -1645,7 +4729,9 @@ func TestPosterStoreCleanupExpiresOldPosters(t *testing.T) { t.Fatalf("store: %v", err) } - ps.cleanup(now) + if err := ps.cleanup(now); err != nil { + t.Fatalf("cleanup: %v", err) + } ps.mu.RLock() _, exists := ps.entries[id] @@ -1662,29 +4748,799 @@ func TestPosterStoreCleanupExpiresOldPosters(t *testing.T) { } } +func regularFileBytes(t *testing.T, dir string) int64 { + t.Helper() + files, err := os.ReadDir(dir) + if err != nil { + t.Fatalf("read directory: %v", err) + } + var total int64 + for _, file := range files { + info, err := file.Info() + if err != nil { + t.Fatalf("stat %s: %v", file.Name(), err) + } + if info.Mode().IsRegular() { + total += info.Size() + } + } + return total +} + +func TestLogStoreRemovalFailureRetainsEntryUntilRetry(t *testing.T) { + dir := t.TempDir() + remover := newDeterministicRemover() + ls := newLogStoreWithRemover(dir, remover.remove) + ls.generateID = func() string { return strings.Repeat("a", logIDLength) } + now := time.Now() + id, _, err := ls.store([]byte("retained"), now) + if err != nil { + t.Fatalf("store: %v", err) + } + path := ls.filePath(id) + ls.mu.Lock() + entry := ls.entries[id] + entry.ExpiresAt = now.Add(-time.Minute) + ls.entries[id] = entry + ls.mu.Unlock() + remover.fail(path, fs.ErrPermission) + + if _, ok, err := ls.lookup(id, now); !errors.Is(err, fs.ErrPermission) || ok { + t.Fatalf("lookup=(ok=%v, err=%v), want unavailable permission error", ok, err) + } + ls.mu.RLock() + _, indexed := ls.entries[id] + artifacts := ls.artifactCountLocked() + ls.mu.RUnlock() + if !indexed || artifacts != 1 { + t.Fatalf("failed removal changed metadata: indexed=%v artifacts=%d", indexed, artifacts) + } + if _, err := os.Stat(path); err != nil { + t.Fatalf("failed removal lost file: %v", err) + } + + remover.recover(path) + if err := ls.cleanup(now); err != nil { + t.Fatalf("retry cleanup: %v", err) + } + if remover.callCount(path) != 2 { + t.Fatalf("remove calls=%d want 2", remover.callCount(path)) + } + if _, err := os.Stat(path); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("file remains after retry: %v", err) + } + if err := ls.cleanup(now); err != nil { + t.Fatalf("idempotent cleanup: %v", err) + } + if remover.callCount(path) != 2 { + t.Fatalf("already committed entry removed again: calls=%d", remover.callCount(path)) + } +} + +func TestRemovalFailureDoesNotBlockUploadsWhileCapacityRemains(t *testing.T) { + t.Run("logs", func(t *testing.T) { + dir := t.TempDir() + remover := newDeterministicRemover() + store := newLogStoreWithRemover(dir, remover.remove) + ids := []string{strings.Repeat("a", logIDLength), strings.Repeat("b", logIDLength)} + nextID := 0 + store.generateID = func() string { + id := ids[nextID] + nextID++ + return id + } + now := time.Now() + firstID, _, err := store.store([]byte("expired"), now) + if err != nil { + t.Fatalf("store expired log: %v", err) + } + store.mu.Lock() + entry := store.entries[firstID] + entry.ExpiresAt = now.Add(-time.Minute) + store.entries[firstID] = entry + store.mu.Unlock() + remover.fail(store.filePath(firstID), fs.ErrPermission) + + secondID, _, err := store.store([]byte("new"), now) + if err != nil { + t.Fatalf("unrelated removal failure blocked log upload: %v", err) + } + store.mu.RLock() + _, firstRetained := store.entries[firstID] + _, secondStored := store.entries[secondID] + store.mu.RUnlock() + if !firstRetained || !secondStored { + t.Fatalf("log accounting lost entries: first=%v second=%v", firstRetained, secondStored) + } + }) + + t.Run("posters", func(t *testing.T) { + dir := t.TempDir() + remover := newDeterministicRemover() + store := newPosterStoreWithRemover(dir, 1024, time.Hour, remover.remove) + now := time.Now() + firstID, first, err := store.store([]byte{1, 2, 3}, "image/png", now.Add(-2*time.Hour)) + if err != nil { + t.Fatalf("store expired poster: %v", err) + } + remover.fail(store.filePath(first.Filename), fs.ErrPermission) + + secondID, _, err := store.store([]byte{4, 5, 6}, "image/png", now) + if err != nil { + t.Fatalf("unrelated removal failure blocked poster upload: %v", err) + } + store.mu.RLock() + _, firstRetained := store.entries[firstID] + _, secondStored := store.entries[secondID] + store.mu.RUnlock() + if !firstRetained || !secondStored { + t.Fatalf("poster accounting lost entries: first=%v second=%v", firstRetained, secondStored) + } + }) +} + +func TestLogStoreErrNotExistCommitsDeletionOnce(t *testing.T) { + remover := newDeterministicRemover() + ls := newLogStoreWithRemover(t.TempDir(), remover.remove) + ls.generateID = func() string { return strings.Repeat("a", logIDLength) } + now := time.Now() + id, _, err := ls.store([]byte("gone"), now) + if err != nil { + t.Fatalf("store: %v", err) + } + path := ls.filePath(id) + if err := os.Remove(path); err != nil { + t.Fatalf("external remove: %v", err) + } + ls.mu.Lock() + entry := ls.entries[id] + entry.ExpiresAt = now.Add(-time.Minute) + ls.entries[id] = entry + ls.mu.Unlock() + + if _, ok, err := ls.lookup(id, now); err != nil || ok { + t.Fatalf("lookup=(ok=%v, err=%v), want clean miss", ok, err) + } + if err := ls.cleanup(now); err != nil { + t.Fatalf("repeat cleanup: %v", err) + } + if remover.callCount(path) != 1 { + t.Fatalf("remove calls=%d want 1", remover.callCount(path)) + } +} + +func TestLogStoreTracksFailedTempCleanup(t *testing.T) { + dir := t.TempDir() + remover := newDeterministicRemover() + ls := newLogStoreWithRemover(dir, remover.remove) + logID := strings.Repeat("a", logIDLength) + ls.generateID = func() string { return logID } + tmpPath := ls.filePath(logID) + ".tmp" + if err := os.Mkdir(tmpPath, 0755); err != nil { + t.Fatalf("seed temp directory: %v", err) + } + cleanupErr := errors.New("synthetic temp removal failure") + remover.fail(tmpPath, cleanupErr) + + if _, _, err := ls.store([]byte("payload"), time.Now()); err == nil { + t.Fatal("store succeeded despite temp write failure") + } + ls.mu.RLock() + _, pending := ls.pendingRemovals[filepath.Base(tmpPath)] + artifacts := ls.artifactCountLocked() + ls.mu.RUnlock() + if !pending || artifacts != 1 { + t.Fatalf("temp cleanup not tracked: pending=%v artifacts=%d", pending, artifacts) + } + + remover.recover(tmpPath) + if err := ls.cleanup(time.Now()); err != nil { + t.Fatalf("retry temp cleanup: %v", err) + } + if _, err := os.Stat(tmpPath); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("temp artifact remains: %v", err) + } +} + +func TestLogStoreStartupReconcilesLiveAndPendingRemovals(t *testing.T) { + dir := t.TempDir() + now := time.Now() + expiredID := strings.Repeat("a", logIDLength) + expiredPath := filepath.Join(dir, expiredID+".log") + tempPath := filepath.Join(dir, "upload.log.tmp") + malformedPath := filepath.Join(dir, "malformed") + for path, data := range map[string][]byte{ + expiredPath: []byte("expired"), + tempPath: []byte("partial"), + malformedPath: []byte("invalid"), + } { + if err := os.WriteFile(path, data, 0644); err != nil { + t.Fatalf("seed %s: %v", filepath.Base(path), err) + } + } + old := now.Add(-logMaxAge - time.Hour) + if err := os.Chtimes(expiredPath, old, old); err != nil { + t.Fatalf("age expired log: %v", err) + } + remover := newDeterministicRemover() + for _, path := range []string{expiredPath, tempPath, malformedPath} { + remover.fail(path, fs.ErrPermission) + } + + ls := newLogStoreWithRemover(dir, remover.remove) + if ls.startupErr == nil { + t.Fatal("startup removal failures were not reported") + } + ls.mu.RLock() + _, live := ls.entries[expiredID] + pending := len(ls.pendingRemovals) + artifacts := ls.artifactCountLocked() + ls.mu.RUnlock() + if !live || pending != 2 || artifacts != 3 { + t.Fatalf("startup accounting: live=%v pending=%d artifacts=%d", live, pending, artifacts) + } + newID, _, err := ls.store([]byte("new"), now) + if err != nil { + t.Fatalf("startup cleanup failure blocked new log: %v", err) + } + + for _, path := range []string{expiredPath, tempPath, malformedPath} { + remover.recover(path) + } + if err := ls.cleanup(now); err != nil { + t.Fatalf("startup retry cleanup: %v", err) + } + restarted := newLogStore(dir) + restarted.mu.RLock() + restartedArtifacts := restarted.artifactCountLocked() + _, newLogRestored := restarted.entries[newID] + restarted.mu.RUnlock() + if restartedArtifacts != 1 || !newLogRestored { + t.Fatalf("restart reconstructed %d artifacts, new log restored=%v", restartedArtifacts, newLogRestored) + } +} + +func TestStoresReconcileConfinedNonEmptyStaleDirectories(t *testing.T) { + t.Run("logs", func(t *testing.T) { + dir := t.TempDir() + staleDir := filepath.Join(dir, "abandoned.log.tmp") + if err := os.MkdirAll(filepath.Join(staleDir, "nested"), 0755); err != nil { + t.Fatalf("seed stale log directory: %v", err) + } + if err := os.WriteFile(filepath.Join(staleDir, "nested", "partial"), []byte("stale"), 0644); err != nil { + t.Fatalf("seed stale log payload: %v", err) + } + + store := newLogStore(dir) + if store.startupErr != nil { + t.Fatalf("startup reconciliation: %v", store.startupErr) + } + if _, err := os.Stat(staleDir); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("stale log directory remains: %v", err) + } + store.generateID = func() string { return strings.Repeat("a", logIDLength) } + if _, _, err := store.store([]byte("new log"), time.Now()); err != nil { + t.Fatalf("store after reconciliation: %v", err) + } + }) + + t.Run("posters", func(t *testing.T) { + dir := t.TempDir() + staleDir := filepath.Join(dir, "abandoned.tmp") + if err := os.MkdirAll(filepath.Join(staleDir, "nested"), 0755); err != nil { + t.Fatalf("seed stale poster directory: %v", err) + } + if err := os.WriteFile(filepath.Join(staleDir, "nested", "partial"), []byte("stale"), 0644); err != nil { + t.Fatalf("seed stale poster payload: %v", err) + } + + store := newPosterStore(dir, 1024, time.Hour) + if store.startupErr != nil { + t.Fatalf("startup reconciliation: %v", store.startupErr) + } + if _, err := os.Stat(staleDir); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("stale poster directory remains: %v", err) + } + if _, _, err := store.store([]byte{1, 2, 3}, "image/png", time.Now()); err != nil { + t.Fatalf("store after reconciliation: %v", err) + } + }) +} + +func TestRecursiveArtifactRemovalRejectsOutsideStore(t *testing.T) { + root := t.TempDir() + outside := t.TempDir() + nested := filepath.Join(outside, "nested") + if err := os.Mkdir(nested, 0755); err != nil { + t.Fatalf("seed outside directory: %v", err) + } + if err := os.WriteFile(filepath.Join(nested, "keep"), []byte("keep"), 0644); err != nil { + t.Fatalf("seed outside payload: %v", err) + } + + err := removeArtifact(os.Remove, root, nested) + if !errors.Is(err, errArtifactOutsideStore) { + t.Fatalf("outside removal error=%v, want confinement error", err) + } + if _, err := os.Stat(filepath.Join(nested, "keep")); err != nil { + t.Fatalf("outside artifact was removed: %v", err) + } +} + +func TestPosterQuotaRemovalFailureDoesNotReclaimAccounting(t *testing.T) { + dir := t.TempDir() + remover := newDeterministicRemover() + ps := newPosterStoreWithRemover(dir, 12, time.Hour, remover.remove) + now := time.Now() + payload := []byte{1, 2, 3, 4, 5, 6, 7} + oldID, oldEntry, err := ps.store(payload, "image/png", now) + if err != nil { + t.Fatalf("store oldest: %v", err) + } + oldPath := ps.filePath(oldEntry.Filename) + remover.fail(oldPath, fs.ErrPermission) + + newID, newEntry, err := ps.store(payload, "image/png", now.Add(time.Minute)) + if !errors.Is(err, fs.ErrPermission) { + t.Fatalf("quota store error=%v want permission error", err) + } + if newID != "" || newEntry != (posterEntry{}) { + t.Fatalf("failed store returned success values: id=%q entry=%+v", newID, newEntry) + } + ps.mu.RLock() + _, retained := ps.entries[oldID] + total := ps.totalBytes + pending := ps.pendingBytes + accounted := ps.accountedBytesLocked() + ps.mu.RUnlock() + if !retained || total != int64(len(payload)) || pending != 0 { + t.Fatalf("failed eviction accounting: retained=%v total=%d pending=%d", retained, total, pending) + } + if physical := regularFileBytes(t, dir); physical != accounted { + t.Fatalf("accounted bytes=%d physical bytes=%d", accounted, physical) + } + + remover.recover(oldPath) + retryID, retryEntry, err := ps.store(payload, "image/png", now.Add(time.Minute)) + if err != nil { + t.Fatalf("retry store: %v", err) + } + if retryID == "" || retryEntry.Size != int64(len(payload)) { + t.Fatalf("retry result: id=%q entry=%+v", retryID, retryEntry) + } + if remover.callCount(oldPath) != 2 { + t.Fatalf("old poster remove calls=%d want 2", remover.callCount(oldPath)) + } + ps.mu.RLock() + accounted = ps.accountedBytesLocked() + total = ps.totalBytes + ps.mu.RUnlock() + if total != int64(len(payload)) || regularFileBytes(t, dir) != accounted { + t.Fatalf("retry accounting: total=%d accounted=%d physical=%d", total, accounted, regularFileBytes(t, dir)) + } +} + +func TestPosterExpiredRemovalFailureAndErrNotExistAreExactOnce(t *testing.T) { + t.Run("failure retains accounting for retry", func(t *testing.T) { + remover := newDeterministicRemover() + ps := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) + now := time.Now() + id, entry, err := ps.store([]byte{1, 2, 3}, "image/png", now) + if err != nil { + t.Fatalf("store: %v", err) + } + path := ps.filePath(entry.Filename) + ps.mu.Lock() + expired := ps.entries[id] + expired.ExpiresAt = now.Add(-time.Minute) + ps.entries[id] = expired + ps.mu.Unlock() + remover.fail(path, fs.ErrPermission) + + if _, ok, err := ps.lookup(entry.Filename, now); !errors.Is(err, fs.ErrPermission) || ok { + t.Fatalf("lookup=(ok=%v, err=%v), want unavailable permission error", ok, err) + } + ps.mu.RLock() + _, retained := ps.entries[id] + total := ps.totalBytes + ps.mu.RUnlock() + if !retained || total != entry.Size { + t.Fatalf("failed expiry accounting: retained=%v total=%d", retained, total) + } + + remover.recover(path) + if err := ps.cleanup(now); err != nil { + t.Fatalf("retry cleanup: %v", err) + } + if err := ps.cleanup(now); err != nil { + t.Fatalf("repeat cleanup: %v", err) + } + if remover.callCount(path) != 2 { + t.Fatalf("remove calls=%d want 2", remover.callCount(path)) + } + ps.mu.RLock() + total = ps.totalBytes + ps.mu.RUnlock() + if total != 0 { + t.Fatalf("totalBytes=%d want 0", total) + } + }) + + t.Run("not exist commits once", func(t *testing.T) { + remover := newDeterministicRemover() + ps := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) + now := time.Now() + id, entry, err := ps.store([]byte{1, 2, 3}, "image/png", now) + if err != nil { + t.Fatalf("store: %v", err) + } + path := ps.filePath(entry.Filename) + if err := os.Remove(path); err != nil { + t.Fatalf("external remove: %v", err) + } + ps.mu.Lock() + expired := ps.entries[id] + expired.ExpiresAt = now.Add(-time.Minute) + ps.entries[id] = expired + ps.mu.Unlock() + + if _, ok, err := ps.lookup(entry.Filename, now); err != nil || ok { + t.Fatalf("lookup=(ok=%v, err=%v), want clean miss", ok, err) + } + if err := ps.cleanup(now); err != nil { + t.Fatalf("repeat cleanup: %v", err) + } + if remover.callCount(path) != 1 { + t.Fatalf("remove calls=%d want 1", remover.callCount(path)) + } + ps.mu.RLock() + total := ps.totalBytes + ps.mu.RUnlock() + if total != 0 { + t.Fatalf("totalBytes=%d want 0", total) + } + }) +} + +func TestPosterStoreKnownCleanupDebtConsumesCapacityAndRetries(t *testing.T) { + dir := t.TempDir() + stalePath := filepath.Join(dir, "poster.tmp") + if err := os.WriteFile(stalePath, []byte("1234"), 0644); err != nil { + t.Fatalf("seed stale poster: %v", err) + } + remover := newDeterministicRemover() + remover.fail(stalePath, fs.ErrPermission) + + ps := newPosterStoreWithRemover(dir, 5, time.Hour, remover.remove) + if ps.startupErr == nil { + t.Fatal("startup removal failure was not reported") + } + if _, _, err := ps.store([]byte{1, 2}, "image/png", time.Now()); err == nil { + t.Fatal("upload exceeded capacity after known stale bytes were accounted") + } + ps.mu.RLock() + pendingBytes := ps.pendingBytes + accountedBytes := ps.accountedBytesLocked() + ps.mu.RUnlock() + if pendingBytes != 4 || accountedBytes != 4 { + t.Fatalf("known debt accounting: pending=%d accounted=%d, want 4", pendingBytes, accountedBytes) + } + if calls := remover.callCount(stalePath); calls != 2 { + t.Fatalf("known debt remove calls=%d, want startup plus upload retry", calls) + } + + remover.recover(stalePath) + if _, entry, err := ps.store([]byte{1, 2}, "image/png", time.Now()); err != nil { + t.Fatalf("store after known debt recovery: %v", err) + } else if entry.Size != 2 { + t.Fatalf("stored entry size=%d, want 2", entry.Size) + } + ps.mu.RLock() + pendingBytes = ps.pendingBytes + ps.mu.RUnlock() + if pendingBytes != 0 { + t.Fatalf("known debt remained after successful retry: %d bytes", pendingBytes) + } +} + +func TestPosterStoreUnknownCleanupDebtDoesNotBlockUploadAndRecovers(t *testing.T) { + dir := t.TempDir() + unknownPath := filepath.Join(dir, "unknown-dir") + if err := os.Mkdir(unknownPath, 0755); err != nil { + t.Fatalf("seed unknown artifact: %v", err) + } + remover := newDeterministicRemover() + remover.fail(unknownPath, fs.ErrPermission) + + ps := newPosterStoreWithRemover(dir, 5, time.Hour, remover.remove) + if ps.startupErr == nil { + t.Fatal("startup removal failure was not reported") + } + if calls := remover.callCount(unknownPath); calls != 1 { + t.Fatalf("startup remove calls=%d, want 1", calls) + } + if _, entry, err := ps.store([]byte{1, 2, 3}, "image/png", time.Now()); err != nil { + t.Fatalf("capacity-safe upload blocked by unknown artifact: %v", err) + } else if entry.Size != 3 { + t.Fatalf("stored entry size=%d, want 3", entry.Size) + } + if calls := remover.callCount(unknownPath); calls != 1 { + t.Fatalf("upload retried permanent unknown debt: calls=%d", calls) + } + + if err := ps.cleanup(time.Now()); !errors.Is(err, fs.ErrPermission) { + t.Fatalf("failed cleanup error=%v, want permission error", err) + } + ps.mu.RLock() + unknownPending := ps.unknownPending + ps.mu.RUnlock() + if unknownPending != 1 { + t.Fatalf("unknown debt count=%d, want 1", unknownPending) + } + + remover.recover(unknownPath) + if err := ps.cleanup(time.Now()); err != nil { + t.Fatalf("cleanup after recovery: %v", err) + } + ps.mu.RLock() + unknownPending = ps.unknownPending + pendingCount := len(ps.pendingRemovals) + ps.mu.RUnlock() + if unknownPending != 0 || pendingCount != 0 { + t.Fatalf("recovered debt remains: unknown=%d pending=%d", unknownPending, pendingCount) + } + if _, err := os.Stat(unknownPath); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("unknown artifact remains after recovery: %v", err) + } +} + +func TestStorageHandlersReturnGenericErrorsForRemovalFailures(t *testing.T) { + remover := newDeterministicRemover() + logs := newLogStoreWithRemover(t.TempDir(), remover.remove) + logs.generateID = func() string { return strings.Repeat("a", logIDLength) } + posters := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) + h := newStorageHarness(t, logs, posters) + now := time.Now() + + logID, _, err := logs.store([]byte("expired log"), now) + if err != nil { + t.Fatalf("store log: %v", err) + } + posterID, poster, err := posters.store([]byte{1, 2, 3}, "image/png", now) + if err != nil { + t.Fatalf("store poster: %v", err) + } + logs.mu.Lock() + logEntry := logs.entries[logID] + logEntry.ExpiresAt = now.Add(-time.Minute) + logs.entries[logID] = logEntry + logs.mu.Unlock() + posters.mu.Lock() + posterEntry := posters.entries[posterID] + posterEntry.ExpiresAt = now.Add(-time.Minute) + posters.entries[posterID] = posterEntry + posters.mu.Unlock() + logPath := logs.filePath(logID) + posterPath := posters.filePath(poster.Filename) + remover.fail(logPath, fs.ErrPermission) + remover.fail(posterPath, fs.ErrPermission) + + for name, target := range map[string]string{ + "log": h.baseURL + "/logs/" + logID, + "poster": h.baseURL + "/posters/" + poster.Filename, + } { + resp, err := http.Get(target) + if err != nil { + t.Fatalf("%s get: %v", name, err) + } + body, readErr := io.ReadAll(resp.Body) + resp.Body.Close() + if readErr != nil { + t.Fatalf("%s response body: %v", name, readErr) + } + if resp.StatusCode != http.StatusInternalServerError { + t.Fatalf("%s status=%d want 500", name, resp.StatusCode) + } + want := "Failed to retrieve " + name + "\n" + if string(body) != want { + t.Fatalf("%s response=%q want %q", name, body, want) + } + } + + remover.recover(logPath) + remover.recover(posterPath) + for name, target := range map[string]string{ + "log": h.baseURL + "/logs/" + logID, + "poster": h.baseURL + "/posters/" + poster.Filename, + } { + resp, err := http.Get(target) + if err != nil { + t.Fatalf("%s recovery get: %v", name, err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("%s recovery status=%d want 404", name, resp.StatusCode) + } + } +} + +func TestPosterHandlerRejectsUploadWhenQuotaRemovalFails(t *testing.T) { + remover := newDeterministicRemover() + logs := newLogStoreWithRemover(t.TempDir(), remover.remove) + payload := []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a, 1, 2, 3} + posters := newPosterStoreWithRemover(t.TempDir(), int64(len(payload)+1), time.Hour, remover.remove) + now := time.Now() + oldID, oldEntry, err := posters.store(payload, "image/png", now) + if err != nil { + t.Fatalf("store old poster: %v", err) + } + oldPath := posters.filePath(oldEntry.Filename) + remover.fail(oldPath, fs.ErrPermission) + h := newStorageHarness(t, logs, posters) + + resp := postPoster(t, h.baseURL, "9.9.9.9", payload) + body, readErr := io.ReadAll(resp.Body) + resp.Body.Close() + if readErr != nil { + t.Fatalf("read failed upload response: %v", readErr) + } + if resp.StatusCode != http.StatusInternalServerError || string(body) != "Failed to store poster\n" { + t.Fatalf("failed upload status=%d body=%q", resp.StatusCode, body) + } + posters.mu.RLock() + _, retained := posters.entries[oldID] + total := posters.totalBytes + posters.mu.RUnlock() + if !retained || total != int64(len(payload)) { + t.Fatalf("failed upload changed old poster: retained=%v total=%d", retained, total) + } + get, err := http.Get(h.baseURL + "/posters/" + oldEntry.Filename) + if err != nil { + t.Fatalf("get retained poster: %v", err) + } + get.Body.Close() + if get.StatusCode != http.StatusOK { + t.Fatalf("retained poster status=%d want 200", get.StatusCode) + } +} + +func TestCleanupStepContinuesAfterRemovalFailureAndThrottlesLogging(t *testing.T) { + remover := newDeterministicRemover() + logs := newLogStoreWithRemover(t.TempDir(), remover.remove) + logs.generateID = func() string { return strings.Repeat("a", logIDLength) } + posters := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) + now := time.Now() + logID, _, err := logs.store([]byte("expired"), now) + if err != nil { + t.Fatalf("store log: %v", err) + } + posterID, poster, err := posters.store([]byte{1, 2, 3}, "image/png", now) + if err != nil { + t.Fatalf("store poster: %v", err) + } + logs.mu.Lock() + logEntry := logs.entries[logID] + logEntry.ExpiresAt = now.Add(-time.Minute) + logs.entries[logID] = logEntry + logs.mu.Unlock() + posters.mu.Lock() + posterEntry := posters.entries[posterID] + posterEntry.ExpiresAt = now.Add(-time.Minute) + posters.entries[posterID] = posterEntry + posters.mu.Unlock() + logPath := logs.filePath(logID) + remover.fail(logPath, fs.ErrPermission) + + srv := &Server{ + rooms: make(map[string]*Room), + logs: logs, + posters: posters, + posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, now), + conns: newConnTracker(), + } + srv.runCleanupStep(now) + logs.mu.RLock() + _, logRetained := logs.entries[logID] + logs.mu.RUnlock() + posters.mu.RLock() + _, posterRetained := posters.entries[posterID] + posters.mu.RUnlock() + if !logRetained || posterRetained { + t.Fatalf("cleanup continuation: log retained=%v poster retained=%v", logRetained, posterRetained) + } + if _, err := os.Stat(posters.filePath(poster.Filename)); !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("poster cleanup did not continue: %v", err) + } + + srv.removalErrors.mu.Lock() + firstLog := srv.removalErrors.lastLog["logs:cleanup"] + srv.removalErrors.mu.Unlock() + if firstLog.IsZero() { + t.Fatal("cleanup removal failure was not made operationally visible") + } + srv.runCleanupStep(now.Add(time.Minute)) + srv.removalErrors.mu.Lock() + secondLog := srv.removalErrors.lastLog["logs:cleanup"] + srv.removalErrors.mu.Unlock() + if !secondLog.Equal(firstLog) { + t.Fatalf("persistent cleanup failure was not throttled: first=%v second=%v", firstLog, secondLog) + } + if remover.callCount(logPath) != 2 { + t.Fatalf("cleanup retry calls=%d want 2", remover.callCount(logPath)) + } +} + +func TestRemovalFailureLogDoesNotExposeCapabilityPath(t *testing.T) { + dir := t.TempDir() + remover := newDeterministicRemover() + store := newLogStoreWithRemover(dir, remover.remove) + id := strings.Repeat("c", logIDLength) + store.generateID = func() string { return id } + now := time.Now() + if _, _, err := store.store([]byte("sensitive"), now); err != nil { + t.Fatalf("store: %v", err) + } + path := store.filePath(id) + store.mu.Lock() + entry := store.entries[id] + entry.ExpiresAt = now.Add(-time.Minute) + store.entries[id] = entry + store.mu.Unlock() + remover.fail(path, &os.PathError{Op: "remove", Path: path, Err: syscall.EACCES}) + removalErr := store.cleanup(now) + if removalErr == nil { + t.Fatal("cleanup unexpectedly succeeded") + } + + var output bytes.Buffer + previousOutput := log.Writer() + previousFlags := log.Flags() + previousPrefix := log.Prefix() + log.SetOutput(&output) + log.SetFlags(0) + log.SetPrefix("") + t.Cleanup(func() { + log.SetOutput(previousOutput) + log.SetFlags(previousFlags) + log.SetPrefix(previousPrefix) + }) + srv := &Server{} + srv.logRemovalError("logs", "cleanup", removalErr) + + message := output.String() + if strings.Contains(message, id) || strings.Contains(message, path) { + t.Fatalf("removal log exposed capability path: %q", message) + } + want := fmt.Sprintf("logs: cleanup removal failed: category=permission errno=%d", syscall.EACCES) + if !strings.Contains(message, want) { + t.Fatalf("removal log=%q, want sanitized context %q", message, want) + } +} + // ====================================================================== // End-to-end: rooms survive a process restart // ====================================================================== -func TestSnapshotSurvivesRestart(t *testing.T) { +func TestSnapshotSurvivesRestartWithHostAuthority(t *testing.T) { stateFile := filepath.Join(t.TempDir(), "rooms.json") hA := newRelayHarnessAt(t, t.TempDir(), stateFile) host := hA.dial(t, "8.0.0.1") - host.send(clientMsg{Type: "create", SessionID: "RESUM", PeerID: "H"}) - host.expect("created") - guest := hA.dial(t, "8.0.0.2") - guest.send(clientMsg{Type: "join", SessionID: "RESUM", PeerID: "G"}) - guest.expect("joined") - host.expect("peerJoined") + host.send(clientMsg{Type: relayTypeCreate, SessionID: "RESUM", PeerID: "H"}) + created := host.expectAuthority(relayTypeCreated, "H") - // Force an early flush so hB can load a populated snapshot. hA's - // t.Cleanup will call flushAndStop again — sync.Once makes it a no-op. if err := hA.srv.snap.flushAndStop(2 * time.Second); err != nil { t.Fatalf("flushAndStop: %v", err) } - if _, err := os.Stat(stateFile); err != nil { - t.Fatalf("snapshot file missing after flush: %v", err) + snapshotBytes, err := os.ReadFile(stateFile) + if err != nil { + t.Fatalf("read snapshot: %v", err) + } + if bytes.Contains(snapshotBytes, []byte(created.ReconnectToken)) { + t.Fatal("snapshot persisted the raw reconnect capability") + } + if !bytes.Contains(snapshotBytes, []byte(`"hostReconnectVerifier"`)) { + t.Fatalf("snapshot omitted host verifier: %s", snapshotBytes) } hB := newRelayHarnessAt(t, t.TempDir(), stateFile) @@ -1695,7 +5551,157 @@ func TestSnapshotSurvivesRestart(t *testing.T) { t.Fatal("room RESUM was not reloaded from snapshot") } - g2 := hB.dial(t, "8.0.0.3") - g2.send(clientMsg{Type: "join", SessionID: "RESUM", PeerID: "G2"}) - g2.expect("joined") + unproved := hB.dial(t, "8.0.0.2") + unproved.send(clientMsg{Type: relayTypeJoin, SessionID: "RESUM", PeerID: "H"}) + unproved.expectError(relayErrorPeerIdUnavailable) + + reconnected := hB.dial(t, "8.0.0.3") + reconnected.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "RESUM", + PeerID: "H", + ReconnectToken: created.ReconnectToken, + }) + joined := reconnected.expectAuthority(relayTypeJoined, "H") + if joined.ReconnectToken != created.ReconnectToken { + t.Fatal("restored host capability changed") + } + + duplicateCreate := hB.dial(t, "8.0.0.4") + duplicateCreate.send(clientMsg{Type: relayTypeCreate, SessionID: "RESUM", PeerID: "OTHER"}) + duplicateCreate.expectError(relayErrorRoomExists) +} + +func TestSnapshotV3RetainsModernHostAndGuestVerifiersAcrossRestart(t *testing.T) { + stateFile := filepath.Join(t.TempDir(), "rooms.json") + hA := newRelayHarnessAt(t, t.TempDir(), stateFile) + + hostToken, _ := mustReconnectToken(t) + host := hA.dial(t, "8.0.1.1") + host.send(clientMsg{ + Type: relayTypeCreate, + SessionID: "V3_RESTART", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + host.expectAuthority(relayTypeCreated, "H") + + guestToken, _ := mustReconnectToken(t) + guest := hA.dial(t, "8.0.1.2") + guest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "V3_RESTART", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guest.expectAuthority(relayTypeJoined, "H") + host.expect(relayTypePeerJoined) + + if err := guest.conn.Close(); err != nil { + t.Fatalf("close guest before snapshot: %v", err) + } + left := host.expect(relayTypePeerLeft) + if left.PeerID != "G" { + t.Fatalf("pre-snapshot disconnect peerId=%q, want G", left.PeerID) + } + if err := hA.srv.snap.flushAndStop(2 * time.Second); err != nil { + t.Fatalf("flush snapshot v3: %v", err) + } + snapshotBytes, err := os.ReadFile(stateFile) + if err != nil { + t.Fatalf("read snapshot v3: %v", err) + } + if !bytes.Contains(snapshotBytes, []byte(`"version":3`)) || + !bytes.Contains(snapshotBytes, []byte(`"peerReconnectVerifiers"`)) { + t.Fatalf("snapshot omitted v3 guest verifier state: %s", snapshotBytes) + } + if bytes.Contains(snapshotBytes, []byte(hostToken)) || bytes.Contains(snapshotBytes, []byte(guestToken)) { + t.Fatal("snapshot persisted a raw reconnect capability") + } + + hB := newRelayHarnessAt(t, t.TempDir(), stateFile) + wrongHostToken, _ := mustReconnectToken(t) + hostThief := hB.dial(t, "8.0.1.3") + hostThief.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "V3_RESTART", + PeerID: "H", + ReconnectToken: wrongHostToken, + ProtocolVersion: relayProtocolVersion, + }) + hostThief.expectError(relayErrorPeerIdUnavailable) + + restartedHost := hB.dial(t, "8.0.1.4") + restartedHost.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "V3_RESTART", + PeerID: "H", + ReconnectToken: hostToken, + ProtocolVersion: relayProtocolVersion, + }) + hostJoined := restartedHost.expectAuthority(relayTypeJoined, "H") + if hostJoined.ReconnectToken != hostToken || hostJoined.ProtocolVersion != relayProtocolVersion { + t.Fatalf("restored host authority changed: %+v", hostJoined) + } + + wrongGuestToken, _ := mustReconnectToken(t) + guestThief := hB.dial(t, "8.0.1.5") + guestThief.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "V3_RESTART", + PeerID: "G", + ReconnectToken: wrongGuestToken, + ProtocolVersion: relayProtocolVersion, + }) + guestThief.expectError(relayErrorPeerIdUnavailable) + + restartedGuest := hB.dial(t, "8.0.1.6") + restartedGuest.send(clientMsg{ + Type: relayTypeJoin, + SessionID: "V3_RESTART", + PeerID: "G", + ReconnectToken: guestToken, + ProtocolVersion: relayProtocolVersion, + }) + guestJoined := restartedGuest.expectAuthority(relayTypeJoined, "H") + if guestJoined.ReconnectToken != guestToken || guestJoined.ProtocolVersion != relayProtocolVersion { + t.Fatalf("restored guest authority changed: %+v", guestJoined) + } + rejoined := restartedHost.expect(relayTypePeerJoined) + if rejoined.PeerID != "G" { + t.Fatalf("restored guest event peerId=%q, want G", rejoined.PeerID) + } +} + +func TestLoadedRoomsConsumeGlobalCapacityWithoutRestoringSourceQuota(t *testing.T) { + stateFile := filepath.Join(t.TempDir(), "rooms.json") + now := time.Now().UTC() + snapshot := stateSnapshot{ + Version: snapshotFormatVersion, + SavedAt: now, + Rooms: makeRoomSnapshots(maxRetainedRooms, false, now), + } + data, err := json.Marshal(snapshot) + if err != nil { + t.Fatalf("marshal full snapshot: %v", err) + } + if err := os.WriteFile(stateFile, data, 0644); err != nil { + t.Fatalf("write full snapshot: %v", err) + } + + h := newRelayHarnessAt(t, t.TempDir(), stateFile) + h.srv.conns.mu.Lock() + restoredQuotaEntries := len(h.srv.conns.roomsPerIP) + h.srv.conns.mu.Unlock() + if restoredQuotaEntries != 0 { + t.Fatalf("restart restored %d process-local quota entries", restoredQuotaEntries) + } + + client := h.dial(t, "8.0.0.4") + client.send(clientMsg{Type: relayTypeCreate, SessionID: "RESTARTOVER", PeerID: "H"}) + client.expectError(relayErrorRateLimited) + client.send(clientMsg{Type: relayTypeJoin, SessionID: "S0000", PeerID: "G"}) + client.expect(relayTypeJoined) } diff --git a/server/oauth.go b/server/oauth.go index 58408cc5..c6ab3761 100644 --- a/server/oauth.go +++ b/server/oauth.go @@ -9,6 +9,7 @@ package main import ( "context" "crypto/rand" + "crypto/sha256" "encoding/base64" "encoding/json" "errors" @@ -30,7 +31,8 @@ const ( oauthMaxSessions = 5000 oauthStartBurst = 3 oauthStartRateSustained = 1 - oauthSessionIDBytes = 18 // 144 bits → 24 base64url chars + oauthBrowserStateBytes = 18 // 144 bits → 24 base64url chars + oauthPollSecretBytes = 18 // Independently generated device capability. oauthPKCEVerifierLen = 64 oauthUpstreamTimeout = 15 * time.Second ) @@ -54,70 +56,88 @@ type oauthTokenResult struct { Error string `json:"error,omitempty"` } -// oauthSession is created by /auth/start and lives until it's consumed by a -// successful /auth/result (which deletes the map entry) or GC'd after -// oauthSessionTTL. The `done` channel is closed exactly once (by complete) and -// unblocks waiters. +// oauthSession is created by /auth/start and lives until its result is claimed +// or it is removed by cleanup. browserState crosses the browser/provider trust +// boundary. Only the SHA-256 digest of the device-only poll secret is retained. +// When both locks are needed, oauthProxy.mu must be acquired before s.mu. type oauthSession struct { - id string + browserState string + pollDigest [sha256.Size]byte service string codeVerifier string // MAL PKCE; empty for AniList. Cleared after token exchange. createdAt time.Time done chan struct{} - mu sync.Mutex - result *oauthTokenResult + mu sync.Mutex + completed bool + result *oauthTokenResult + + // Test seam used to deterministically seat concurrent result waiters. + waitStarted func() } -func (s *oauthSession) complete(r oauthTokenResult) { - s.mu.Lock() - defer s.mu.Unlock() - if s.result != nil { - return +// completeLocked publishes at most one terminal result. The caller holds s.mu. +func (s *oauthSession) completeLocked(r oauthTokenResult) bool { + if s.completed { + return false } + s.completed = true s.result = &r s.codeVerifier = "" // Secret, not needed after exchange. close(s.done) + return true } -// wait blocks until the session is completed or ctx is cancelled. -func (s *oauthSession) wait(ctx context.Context) (*oauthTokenResult, error) { +// wait blocks only until the session is ready or ctx is cancelled. Result +// ownership is transferred separately by oauthProxy.claimResult. +func (s *oauthSession) wait(ctx context.Context) error { + if s.waitStarted != nil { + s.waitStarted() + } select { case <-s.done: - s.mu.Lock() - defer s.mu.Unlock() - return s.result, nil + return nil case <-ctx.Done(): - return nil, ctx.Err() + return ctx.Err() } } -type oauthProxy struct { - baseURL string // e.g. https://ice.plezy.app - services map[string]oauthServiceConfig - client *http.Client +func (s *oauthSession) pkceVerifier() string { + s.mu.Lock() + defer s.mu.Unlock() + return s.codeVerifier +} - mu sync.Mutex - sessions map[string]*oauthSession +type oauthProxy struct { + baseURL string // e.g. https://ice.plezy.app + services map[string]oauthServiceConfig + client *http.Client + clientIPs clientIPResolver + + mu sync.Mutex + browserStates map[string]*oauthSession + pollDigests map[[sha256.Size]byte]*oauthSession ipMu sync.Mutex ipRate map[string]*rateLimiter } -func newOAuthProxy(baseURL string, services map[string]oauthServiceConfig) *oauthProxy { +func newOAuthProxy(baseURL string, services map[string]oauthServiceConfig, clientIPs clientIPResolver) *oauthProxy { return &oauthProxy{ - baseURL: strings.TrimRight(baseURL, "/"), - services: services, - client: &http.Client{Timeout: oauthUpstreamTimeout}, - sessions: make(map[string]*oauthSession), - ipRate: make(map[string]*rateLimiter), + baseURL: strings.TrimRight(baseURL, "/"), + services: services, + client: &http.Client{Timeout: oauthUpstreamTimeout}, + clientIPs: clientIPs, + browserStates: make(map[string]*oauthSession), + pollDigests: make(map[[sha256.Size]byte]*oauthSession), + ipRate: make(map[string]*rateLimiter), } } // oauthConfigFromEnv reads the public base URL and per-service creds from the // environment. Returns (nil, false) if OAUTH_BASE_URL is unset — the caller // wires this as "OAuth disabled, endpoints return 503". -func oauthConfigFromEnv() (*oauthProxy, bool) { +func oauthConfigFromEnv(clientIPs clientIPResolver) (*oauthProxy, bool) { base := os.Getenv("OAUTH_BASE_URL") if base == "" { return nil, false @@ -140,7 +160,7 @@ func oauthConfigFromEnv() (*oauthProxy, bool) { TokenURL: "https://anilist.co/api/v2/oauth/token", } } - return newOAuthProxy(base, services), true + return newOAuthProxy(base, services, clientIPs), true } // registerOAuthRoutes registers all /auth/* handlers. If p is nil (no env @@ -180,13 +200,18 @@ func (p *oauthProxy) handleAuthRoot(w http.ResponseWriter, r *http.Request) { } // POST /auth/start body={"service":"mal"|"anilist"} -// Response: {"session":"...","url":"https://.../auth/:service?session=...","expiresIn":600} +// Response: {"session":"device-only poll capability","url":"https://.../auth/:service?state=...","expiresIn":600} func (p *oauthProxy) handleStart(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "no-store, private") if r.Method != http.MethodPost { http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) return } - ip := clientIP(r) + ip, err := p.clientIPs.resolve(r) + if err != nil { + http.Error(w, "Invalid client address", http.StatusBadRequest) + return + } if !p.ipAllow(ip) { http.Error(w, "Rate limited", http.StatusTooManyRequests) return @@ -205,36 +230,47 @@ func (p *oauthProxy) handleStart(w http.ResponseWriter, r *http.Request) { return } - // Generate tokens outside the map lock — crypto/rand syscalls would - // otherwise serialize concurrent /auth/start calls. - sess := &oauthSession{ - id: randToken(oauthSessionIDBytes), - service: body.Service, - createdAt: time.Now(), - done: make(chan struct{}), - } - if cfg.UsePKCE { - sess.codeVerifier = randPKCEVerifier() - } + // Generate independent trust-domain values outside the map lock — + // crypto/rand syscalls must not serialize concurrent /auth/start calls. + var pollSecret string + var sess *oauthSession + for { + pollSecret = randToken(oauthPollSecretBytes) + sess = &oauthSession{ + browserState: randToken(oauthBrowserStateBytes), + pollDigest: digestPollSecret(pollSecret), + service: body.Service, + createdAt: time.Now(), + done: make(chan struct{}), + } + if cfg.UsePKCE { + sess.codeVerifier = randPKCEVerifier() + } - p.mu.Lock() - if len(p.sessions) >= oauthMaxSessions { + p.mu.Lock() + if len(p.browserStates) >= oauthMaxSessions { + p.mu.Unlock() + http.Error(w, "Server busy", http.StatusServiceUnavailable) + return + } + if p.browserStates[sess.browserState] != nil || p.pollDigests[sess.pollDigest] != nil { + p.mu.Unlock() + continue + } + p.addSessionLocked(sess) p.mu.Unlock() - http.Error(w, "Server busy", http.StatusServiceUnavailable) - return + break } - p.sessions[sess.id] = sess - p.mu.Unlock() resp := map[string]any{ - "session": sess.id, - "url": fmt.Sprintf("%s/auth/%s?session=%s", p.baseURL, url.PathEscape(body.Service), url.QueryEscape(sess.id)), + "session": pollSecret, + "url": fmt.Sprintf("%s/auth/%s?state=%s", p.baseURL, url.PathEscape(body.Service), url.QueryEscape(sess.browserState)), "expiresIn": int(oauthSessionTTL.Seconds()), } writeJSON(w, http.StatusOK, resp) } -// GET /auth/:service?session=X → 302 upstream authorize URL +// GET /auth/:service?state=X → 302 upstream authorize URL func (p *oauthProxy) handleAuthorize(w http.ResponseWriter, r *http.Request, service string) { if r.Method != http.MethodGet { http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) @@ -245,9 +281,9 @@ func (p *oauthProxy) handleAuthorize(w http.ResponseWriter, r *http.Request, ser http.NotFound(w, r) return } - sessionID := r.URL.Query().Get("session") + browserState := r.URL.Query().Get("state") p.mu.Lock() - sess := p.sessions[sessionID] + sess := p.browserStates[browserState] p.mu.Unlock() if sess == nil || sess.service != service { renderErrorPage(w, http.StatusNotFound, "This sign-in link is no longer valid. Start again from Plezy.") @@ -258,13 +294,13 @@ func (p *oauthProxy) handleAuthorize(w http.ResponseWriter, r *http.Request, ser "response_type": {"code"}, "client_id": {cfg.ClientID}, "redirect_uri": {p.redirectURI(service)}, - "state": {sess.id}, + "state": {sess.browserState}, } if cfg.Scopes != "" { q.Set("scope", cfg.Scopes) } if cfg.UsePKCE { - q.Set("code_challenge", sess.codeVerifier) // plain method ⇒ challenge == verifier + q.Set("code_challenge", sess.pkceVerifier()) // plain method ⇒ challenge == verifier q.Set("code_challenge_method", cfg.PKCEMethod) } http.Redirect(w, r, cfg.AuthorizeURL+"?"+q.Encode(), http.StatusFound) @@ -284,7 +320,7 @@ func (p *oauthProxy) handleCallback(w http.ResponseWriter, r *http.Request, serv q := r.URL.Query() state := q.Get("state") p.mu.Lock() - sess := p.sessions[state] + sess := p.browserStates[state] p.mu.Unlock() if sess == nil || sess.service != service { renderErrorPage(w, http.StatusNotFound, "This sign-in link is no longer valid. Start again from Plezy.") @@ -292,13 +328,25 @@ func (p *oauthProxy) handleCallback(w http.ResponseWriter, r *http.Request, serv } if upstreamErr := q.Get("error"); upstreamErr != "" { - sess.complete(oauthTokenResult{Error: upstreamErr}) - renderErrorPage(w, http.StatusOK, "Sign-in was cancelled.") + publicError := "authorization_failed" + message := "Sign-in failed. Please try again." + if upstreamErr == "access_denied" { + publicError = "access_denied" + message = "Sign-in was cancelled." + } + if !p.completeSession(sess, oauthTokenResult{Error: publicError}) { + renderErrorPage(w, http.StatusNotFound, "This sign-in link is no longer valid. Start again from Plezy.") + return + } + renderErrorPage(w, http.StatusOK, message) return } code := q.Get("code") if code == "" { - sess.complete(oauthTokenResult{Error: "missing_code"}) + if !p.completeSession(sess, oauthTokenResult{Error: "missing_code"}) { + renderErrorPage(w, http.StatusNotFound, "This sign-in link is no longer valid. Start again from Plezy.") + return + } renderErrorPage(w, http.StatusBadRequest, "Sign-in response was incomplete. Please try again.") return } @@ -306,23 +354,30 @@ func (p *oauthProxy) handleCallback(w http.ResponseWriter, r *http.Request, serv tok, err := p.exchangeCode(r.Context(), cfg, service, sess, code) if err != nil { log.Printf("oauth: %s token exchange failed: %v", service, err) - sess.complete(oauthTokenResult{Error: "exchange_failed"}) + if !p.completeSession(sess, oauthTokenResult{Error: "exchange_failed"}) { + renderErrorPage(w, http.StatusNotFound, "This sign-in link is no longer valid. Start again from Plezy.") + return + } renderErrorPage(w, http.StatusBadGateway, "Couldn't complete sign-in. Please try again.") return } - sess.complete(tok) + if !p.completeSession(sess, tok) { + renderErrorPage(w, http.StatusNotFound, "This sign-in link is no longer valid. Start again from Plezy.") + return + } renderSuccessPage(w) } -// GET /auth/result?session=X → long-poll, returns tokens on success +// GET /auth/result?session=X → long-poll, returns one terminal result. func (p *oauthProxy) handleResult(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "no-store, private") if r.Method != http.MethodGet { http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) return } - sessionID := r.URL.Query().Get("session") + pollDigest := digestPollSecret(r.URL.Query().Get("session")) p.mu.Lock() - sess := p.sessions[sessionID] + sess := p.pollDigests[pollDigest] p.mu.Unlock() if sess == nil { http.Error(w, "Session not found", http.StatusGone) @@ -331,17 +386,16 @@ func (p *oauthProxy) handleResult(w http.ResponseWriter, r *http.Request) { ctx, cancel := context.WithTimeout(r.Context(), oauthResultWait) defer cancel() - result, err := sess.wait(ctx) - if err != nil { + if err := sess.wait(ctx); err != nil { // Client should retry — session may still receive its callback. w.WriteHeader(http.StatusNoContent) return } - // Session consumed — delete so a retry sees 410 instead of racing another wait. - p.mu.Lock() - delete(p.sessions, sess.id) - p.mu.Unlock() - + result, ok := p.claimResult(pollDigest, sess) + if !ok { + http.Error(w, "Session not found", http.StatusGone) + return + } if result.Error != "" { writeJSON(w, http.StatusOK, map[string]any{"error": result.Error}) return @@ -369,7 +423,7 @@ func (p *oauthProxy) exchangeCode(ctx context.Context, cfg oauthServiceConfig, s form.Set("client_secret", cfg.ClientSecret) } if cfg.UsePKCE { - form.Set("code_verifier", sess.codeVerifier) + form.Set("code_verifier", sess.pkceVerifier()) } req, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.TokenURL, strings.NewReader(form.Encode())) @@ -410,13 +464,59 @@ func (p *oauthProxy) redirectURI(service string) string { return fmt.Sprintf("%s/auth/%s/callback", p.baseURL, service) } +// addSessionLocked installs both independently generated keys as one logical +// session. The caller has already verified that neither key is live. +func (p *oauthProxy) addSessionLocked(sess *oauthSession) { + p.browserStates[sess.browserState] = sess + p.pollDigests[sess.pollDigest] = sess +} + +// removeSessionLocked removes only entries still owned by sess, so a stale +// callback or waiter cannot remove a replacement. +func (p *oauthProxy) removeSessionLocked(sess *oauthSession) { + if p.browserStates[sess.browserState] == sess { + delete(p.browserStates, sess.browserState) + } + if p.pollDigests[sess.pollDigest] == sess { + delete(p.pollDigests, sess.pollDigest) + } +} + +func (p *oauthProxy) completeSession(sess *oauthSession, result oauthTokenResult) bool { + p.mu.Lock() + defer p.mu.Unlock() + if p.browserStates[sess.browserState] != sess || p.pollDigests[sess.pollDigest] != sess { + return false + } + sess.mu.Lock() + defer sess.mu.Unlock() + return sess.completeLocked(result) +} + +func (p *oauthProxy) claimResult(digest [sha256.Size]byte, sess *oauthSession) (oauthTokenResult, bool) { + p.mu.Lock() + defer p.mu.Unlock() + if p.pollDigests[digest] != sess || p.browserStates[sess.browserState] != sess { + return oauthTokenResult{}, false + } + sess.mu.Lock() + defer sess.mu.Unlock() + if sess.result == nil { + return oauthTokenResult{}, false + } + result := *sess.result + sess.result = nil + p.removeSessionLocked(sess) + return result, true +} + // cleanup drops sessions past oauthSessionTTL. Called by the main cleanup loop. func (p *oauthProxy) cleanup() { now := time.Now() p.mu.Lock() - for id, sess := range p.sessions { + for _, sess := range p.browserStates { if now.Sub(sess.createdAt) > oauthSessionTTL { - delete(p.sessions, id) + p.removeSessionLocked(sess) } } p.mu.Unlock() @@ -463,6 +563,10 @@ func writeJSON(w http.ResponseWriter, status int, v any) { _ = json.NewEncoder(w).Encode(v) } +func digestPollSecret(secret string) [sha256.Size]byte { + return sha256.Sum256([]byte(secret)) +} + func randToken(numBytes int) string { b := make([]byte, numBytes) if _, err := rand.Read(b); err != nil { diff --git a/server/oauth_test.go b/server/oauth_test.go index 23da5dad..7c0797c4 100644 --- a/server/oauth_test.go +++ b/server/oauth_test.go @@ -104,13 +104,18 @@ func (m *mockUpstream) form() url.Values { // newOAuthHarness boots a relay-less httptest server that mounts /auth/* only, // with `mal` and `anilist` services pointed at a shared mock upstream. type oauthHarness struct { - proxy *oauthProxy - srv *httptest.Server - base string + proxy *oauthProxy + srv *httptest.Server + base string upstream *mockUpstream } func newOAuthHarness(t *testing.T) *oauthHarness { + t.Helper() + return newOAuthHarnessWithResolver(t, mustClientIPResolver(t, "127.0.0.0/8")) +} + +func newOAuthHarnessWithResolver(t *testing.T, clientIPs clientIPResolver) *oauthHarness { t.Helper() up := newMockUpstream(t) proxy := newOAuthProxy("http://placeholder", map[string]oauthServiceConfig{ @@ -127,7 +132,7 @@ func newOAuthHarness(t *testing.T) *oauthHarness { AuthorizeURL: up.srv.URL + "/oauth/authorize", TokenURL: up.srv.URL + "/oauth/token", }, - }) + }, clientIPs) mux := http.NewServeMux() registerOAuthRoutes(mux, proxy) srv := httptest.NewServer(mux) @@ -137,7 +142,7 @@ func newOAuthHarness(t *testing.T) *oauthHarness { return &oauthHarness{proxy: proxy, srv: srv, base: srv.URL, upstream: up} } -func (h *oauthHarness) startSession(t *testing.T, service, ip string) (sessionID, qrURL string) { +func (h *oauthHarness) startSession(t *testing.T, service, ip string) (pollSecret, browserState, authorizeURL string) { t.Helper() body, _ := json.Marshal(map[string]string{"service": service}) req, _ := http.NewRequest(http.MethodPost, h.base+"/auth/start", bytes.NewReader(body)) @@ -153,6 +158,9 @@ func (h *oauthHarness) startSession(t *testing.T, service, ip string) (sessionID if resp.StatusCode != http.StatusOK { t.Fatalf("start status=%d", resp.StatusCode) } + if got := resp.Header.Get("Cache-Control"); got != "no-store, private" { + t.Fatalf("start Cache-Control=%q", got) + } var out struct { Session string `json:"session"` URL string `json:"url"` @@ -161,22 +169,58 @@ func (h *oauthHarness) startSession(t *testing.T, service, ip string) (sessionID if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { t.Fatalf("decode: %v", err) } - if out.Session == "" || out.URL == "" { - t.Fatalf("empty session/url: %+v", out) + parsed, err := url.Parse(out.URL) + if err != nil { + t.Fatalf("parse authorize URL: %v", err) } - return out.Session, out.URL + state := parsed.Query().Get("state") + if out.Session == "" || out.URL == "" || state == "" { + t.Fatalf("empty poll secret/url/browser state: %+v", out) + } + return out.Session, state, out.URL +} + +func postOAuthStart(t *testing.T, h *oauthHarness, service, xff string) *http.Response { + t.Helper() + body, err := json.Marshal(map[string]string{"service": service}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + req, err := http.NewRequest(http.MethodPost, h.base+"/auth/start", bytes.NewReader(body)) + if err != nil { + t.Fatalf("new request: %v", err) + } + req.Header.Set("Content-Type", "application/json") + if xff != "" { + req.Header.Set("X-Forwarded-For", xff) + } + return httpDo(t, req) } // ====== /auth/start ====== -func TestOAuthStartReturnsSessionAndURL(t *testing.T) { +func TestOAuthStartSeparatesDeviceCapabilityFromBrowserState(t *testing.T) { h := newOAuthHarness(t) - session, qr := h.startSession(t, "mal", "1.2.3.4") - if !strings.HasPrefix(qr, h.base+"/auth/mal?session=") { - t.Fatalf("url=%q doesn't look like the authorize start URL", qr) + pollSecret, browserState, authorizeURL := h.startSession(t, "mal", "1.2.3.4") + if pollSecret == browserState { + t.Fatal("device poll capability must differ from browser state") } - if !strings.Contains(qr, url.QueryEscape(session)) { - t.Fatalf("url=%q missing session token", qr) + if !strings.HasPrefix(authorizeURL, h.base+"/auth/mal?state=") { + t.Fatalf("url=%q doesn't look like the authorize start URL", authorizeURL) + } + if strings.Contains(authorizeURL, pollSecret) { + t.Fatalf("authorize URL disclosed device poll capability") + } + if got := h.proxy.pollDigests[digestPollSecret(pollSecret)]; got == nil { + t.Fatal("poll capability digest was not indexed") + } + if got := h.proxy.browserStates[browserState]; got == nil { + t.Fatal("browser state was not indexed") + } + pollAsBrowser := httpGet(t, h.base+"/auth/mal?state="+url.QueryEscape(pollSecret)) + pollAsBrowser.Body.Close() + if pollAsBrowser.StatusCode != http.StatusNotFound { + t.Fatalf("poll capability authorized browser path: status=%d", pollAsBrowser.StatusCode) } } @@ -202,8 +246,8 @@ func TestOAuthStartRejectsInvalidJSON(t *testing.T) { func TestOAuthStartRateLimitedPerIP(t *testing.T) { h := newOAuthHarness(t) ip := "5.5.5.5" - for i := 0; i < oauthStartBurst; i++ { - h.startSession(t, "mal", ip) // should all succeed + for range oauthStartBurst { + _, _, _ = h.startSession(t, "mal", ip) // should all succeed } body, _ := json.Marshal(map[string]string{"service": "mal"}) req, err := http.NewRequest(http.MethodPost, h.base+"/auth/start", bytes.NewReader(body)) @@ -219,6 +263,98 @@ func TestOAuthStartRateLimitedPerIP(t *testing.T) { } } +func TestOAuthSessionLimitCountsLogicalSessions(t *testing.T) { + h := newOAuthHarness(t) + h.proxy.mu.Lock() + for i := range oauthMaxSessions - 1 { + sess := &oauthSession{ + browserState: fmt.Sprintf("state-%d", i), + pollDigest: digestPollSecret(fmt.Sprintf("poll-%d", i)), + } + h.proxy.addSessionLocked(sess) + } + h.proxy.mu.Unlock() + + body, _ := json.Marshal(map[string]string{"service": "mal"}) + req := httptest.NewRequest(http.MethodPost, "/auth/start", bytes.NewReader(body)) + rec := httptest.NewRecorder() + h.proxy.handleStart(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("start with %d logical sessions: status=%d want 200", oauthMaxSessions-1, rec.Code) + } + h.proxy.mu.Lock() + browserCount := len(h.proxy.browserStates) + pollCount := len(h.proxy.pollDigests) + h.proxy.mu.Unlock() + if browserCount != oauthMaxSessions || pollCount != oauthMaxSessions { + t.Fatalf("index counts browser=%d poll=%d want %d each", browserCount, pollCount, oauthMaxSessions) + } +} + +func TestOAuthStartUsesTrustedCanonicalClientIdentity(t *testing.T) { + t.Run("untrusted spoof rotation shares direct peer bucket", func(t *testing.T) { + h := newOAuthHarnessWithResolver(t, newClientIPResolver(nil)) + for i := range oauthStartBurst { + resp := postOAuthStart(t, h, "mal", fmt.Sprintf("203.0.113.%d", i+1)) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("start %d status=%d", i, resp.StatusCode) + } + } + denied := postOAuthStart(t, h, "mal", "198.51.100.10") + denied.Body.Close() + if denied.StatusCode != http.StatusTooManyRequests { + t.Fatalf("rotated spoof status=%d, want 429", denied.StatusCode) + } + }) + + t.Run("validated clients have independent buckets", func(t *testing.T) { + h := newOAuthHarness(t) + for _, ip := range []string{"203.0.113.1", "203.0.113.2"} { + for range oauthStartBurst { + resp := postOAuthStart(t, h, "mal", ip) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("client %s status=%d", ip, resp.StatusCode) + } + } + } + }) + + t.Run("malformed trusted chain mutates no state", func(t *testing.T) { + h := newOAuthHarness(t) + resp := postOAuthStart(t, h, "mal", "203.0.113.1,") + resp.Body.Close() + if resp.StatusCode != http.StatusBadRequest { + t.Fatalf("status=%d, want 400", resp.StatusCode) + } + h.proxy.ipMu.Lock() + rateCount := len(h.proxy.ipRate) + h.proxy.ipMu.Unlock() + h.proxy.mu.Lock() + browserCount := len(h.proxy.browserStates) + pollCount := len(h.proxy.pollDigests) + h.proxy.mu.Unlock() + if rateCount != 0 || browserCount != 0 || pollCount != 0 { + t.Fatalf( + "malformed chain mutated OAuth state: rates=%d browser=%d poll=%d", + rateCount, + browserCount, + pollCount, + ) + } + }) + + t.Run("untrusted malformed header is ignored", func(t *testing.T) { + h := newOAuthHarnessWithResolver(t, newClientIPResolver(nil)) + resp := postOAuthStart(t, h, "mal", "bad,") + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("status=%d, want 200", resp.StatusCode) + } + }) +} + func TestOAuthStartMethodNotAllowed(t *testing.T) { h := newOAuthHarness(t) resp := httpGet(t, h.base+"/auth/start") @@ -232,10 +368,10 @@ func TestOAuthStartMethodNotAllowed(t *testing.T) { func TestOAuthAuthorizeMALRedirectIncludesPKCE(t *testing.T) { h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "1.1.1.1") + pollSecret, browserState, _ := h.startSession(t, "mal", "1.1.1.1") client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }} - resp, err := client.Get(h.base + "/auth/mal?session=" + url.QueryEscape(sess)) + resp, err := client.Get(h.base + "/auth/mal?state=" + url.QueryEscape(browserState)) if err != nil { t.Fatalf("get: %v", err) } @@ -254,8 +390,11 @@ func TestOAuthAuthorizeMALRedirectIncludesPKCE(t *testing.T) { if q.Get("response_type") != "code" { t.Errorf("response_type=%q", q.Get("response_type")) } - if q.Get("state") != sess { - t.Errorf("state=%q, want session %q", q.Get("state"), sess) + if q.Get("state") != browserState { + t.Errorf("state=%q, want browser state %q", q.Get("state"), browserState) + } + if q.Get("state") == pollSecret { + t.Error("provider state disclosed device poll capability") } if q.Get("code_challenge_method") != "plain" { t.Errorf("code_challenge_method=%q, want plain", q.Get("code_challenge_method")) @@ -270,9 +409,9 @@ func TestOAuthAuthorizeMALRedirectIncludesPKCE(t *testing.T) { func TestOAuthAuthorizeAnilistRedirectOmitsPKCE(t *testing.T) { h := newOAuthHarness(t) - sess, _ := h.startSession(t, "anilist", "1.1.1.2") + _, browserState, _ := h.startSession(t, "anilist", "1.1.1.2") client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }} - resp, err := client.Get(h.base + "/auth/anilist?session=" + url.QueryEscape(sess)) + resp, err := client.Get(h.base + "/auth/anilist?state=" + url.QueryEscape(browserState)) if err != nil { t.Fatalf("get: %v", err) } @@ -286,7 +425,7 @@ func TestOAuthAuthorizeAnilistRedirectOmitsPKCE(t *testing.T) { func TestOAuthAuthorizeUnknownSessionRendersError(t *testing.T) { h := newOAuthHarness(t) - resp := httpGet(t, h.base+"/auth/mal?session=bogus") + resp := httpGet(t, h.base+"/auth/mal?state=bogus") defer resp.Body.Close() if resp.StatusCode != http.StatusNotFound { t.Fatalf("status=%d want 404", resp.StatusCode) @@ -299,9 +438,9 @@ func TestOAuthAuthorizeUnknownSessionRendersError(t *testing.T) { func TestOAuthAuthorizeWrongServiceRejected(t *testing.T) { h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "1.1.1.3") - // Try to use the MAL session against the AniList authorize endpoint. - resp := httpGet(t, h.base+"/auth/anilist?session="+url.QueryEscape(sess)) + _, browserState, _ := h.startSession(t, "mal", "1.1.1.3") + // Try to use the MAL browser state against the AniList authorize endpoint. + resp := httpGet(t, h.base+"/auth/anilist?state="+url.QueryEscape(browserState)) defer resp.Body.Close() if resp.StatusCode != http.StatusNotFound { t.Fatalf("status=%d want 404", resp.StatusCode) @@ -310,117 +449,238 @@ func TestOAuthAuthorizeWrongServiceRejected(t *testing.T) { // ====== /auth/:service/callback + /auth/result ====== -func TestOAuthCallbackExchangesCodeAndResultReturnsTokens(t *testing.T) { - h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "2.2.2.1") +type oauthResultResponse struct { + status int + cacheControl string + body map[string]any + err error +} - resultCh := make(chan map[string]any, 1) - go func() { - resp, err := http.Get(h.base + "/auth/result?session=" + url.QueryEscape(sess)) - if err != nil { - resultCh <- map[string]any{"_err": err.Error()} - return - } - defer resp.Body.Close() - var m map[string]any - _ = json.NewDecoder(resp.Body).Decode(&m) - resultCh <- m - }() - - // Hit the callback as the upstream browser would. - cbURL := fmt.Sprintf("%s/auth/mal/callback?code=CODE123&state=%s", h.base, url.QueryEscape(sess)) - resp, err := http.Get(cbURL) +func requestOAuthResult(rawURL string) oauthResultResponse { + resp, err := http.Get(rawURL) if err != nil { - t.Fatalf("callback: %v", err) + return oauthResultResponse{err: err} } - resp.Body.Close() - if resp.StatusCode != http.StatusOK { - t.Fatalf("callback status=%d", resp.StatusCode) + defer resp.Body.Close() + out := oauthResultResponse{ + status: resp.StatusCode, + cacheControl: resp.Header.Get("Cache-Control"), } + if resp.StatusCode == http.StatusOK { + out.body = make(map[string]any) + out.err = json.NewDecoder(resp.Body).Decode(&out.body) + } + return out +} - // Upstream should have been called with PKCE + code. - form := h.upstream.form() - if form.Get("code") != "CODE123" { - t.Errorf("upstream code=%q", form.Get("code")) - } - if form.Get("code_verifier") == "" { - t.Error("upstream missing code_verifier (PKCE)") - } - if form.Get("grant_type") != "authorization_code" { - t.Errorf("grant_type=%q", form.Get("grant_type")) - } +func TestOAuthBrowserStateCannotClaimResult(t *testing.T) { + for _, service := range []string{"mal", "anilist"} { + t.Run(service, func(t *testing.T) { + h := newOAuthHarness(t) + pollSecret, browserState, _ := h.startSession(t, service, "2.2.2.1") - select { - case got := <-resultCh: - if got["accessToken"] != "tok-abc" { - t.Errorf("accessToken=%v want tok-abc", got["accessToken"]) - } - if got["refreshToken"] != "ref-xyz" { - t.Errorf("refreshToken=%v want ref-xyz", got["refreshToken"]) - } - case <-time.After(3 * time.Second): - t.Fatal("result never returned") + browserClaim := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(browserState)) + if browserClaim.err != nil { + t.Fatalf("browser-state result request: %v", browserClaim.err) + } + if browserClaim.status != http.StatusGone { + t.Fatalf("browser-state result status=%d want 410", browserClaim.status) + } + + callback := fmt.Sprintf("%s/auth/%s/callback?code=CODE123&state=%s", h.base, service, url.QueryEscape(browserState)) + resp := httpGet(t, callback) + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("callback status=%d", resp.StatusCode) + } + + result := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(pollSecret)) + if result.err != nil { + t.Fatalf("device result request: %v", result.err) + } + if result.status != http.StatusOK { + t.Fatalf("device result status=%d want 200", result.status) + } + if result.cacheControl != "no-store, private" { + t.Fatalf("result Cache-Control=%q", result.cacheControl) + } + if result.body["accessToken"] != "tok-abc" || result.body["refreshToken"] != "ref-xyz" { + t.Fatalf("unexpected result: %v", result.body) + } + + second := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(pollSecret)) + if second.err != nil { + t.Fatalf("second result request: %v", second.err) + } + if second.status != http.StatusGone { + t.Fatalf("second result status=%d want 410", second.status) + } + + form := h.upstream.form() + if form.Get("code") != "CODE123" { + t.Errorf("upstream code=%q", form.Get("code")) + } + if service == "mal" && form.Get("code_verifier") == "" { + t.Error("upstream missing MAL code verifier") + } + if service == "anilist" && form.Get("code_verifier") != "" { + t.Error("AniList exchange unexpectedly included a code verifier") + } + }) } } -func TestOAuthCallbackUpstreamError(t *testing.T) { +func TestOAuthConcurrentResultClaimIsOneShot(t *testing.T) { + for _, tc := range []struct { + name string + result oauthTokenResult + }{ + {name: "token", result: oauthTokenResult{AccessToken: "tok"}}, + {name: "provider error", result: oauthTokenResult{Error: "authorization_failed"}}, + } { + t.Run(tc.name, func(t *testing.T) { + h := newOAuthHarness(t) + pollSecret, browserState, _ := h.startSession(t, "mal", "2.2.2.2") + digest := digestPollSecret(pollSecret) + + h.proxy.mu.Lock() + sess := h.proxy.pollDigests[digest] + h.proxy.mu.Unlock() + if sess == nil { + t.Fatal("session missing from poll index") + } + + seated := make(chan struct{}, 2) + sess.waitStarted = func() { seated <- struct{}{} } + results := make(chan oauthResultResponse, 2) + resultURL := h.base + "/auth/result?session=" + url.QueryEscape(pollSecret) + go func() { results <- requestOAuthResult(resultURL) }() + go func() { results <- requestOAuthResult(resultURL) }() + + for range 2 { + select { + case <-seated: + case <-time.After(3 * time.Second): + t.Fatal("result waiter did not reach readiness boundary") + } + } + if !h.proxy.completeSession(sess, tc.result) { + t.Fatal("could not complete current session") + } + + statuses := map[int]int{} + var winningBody map[string]any + for range 2 { + select { + case got := <-results: + if got.err != nil { + t.Fatalf("result request: %v", got.err) + } + if got.cacheControl != "no-store, private" { + t.Errorf("result Cache-Control=%q", got.cacheControl) + } + statuses[got.status]++ + if got.status == http.StatusOK { + winningBody = got.body + } + case <-time.After(3 * time.Second): + t.Fatal("result request did not return") + } + } + if statuses[http.StatusOK] != 1 || statuses[http.StatusGone] != 1 { + t.Fatalf("statuses=%v want one 200 and one 410", statuses) + } + if tc.result.Error != "" && winningBody["error"] != tc.result.Error { + t.Fatalf("winning error result=%v", winningBody) + } + if tc.result.AccessToken != "" && winningBody["accessToken"] != tc.result.AccessToken { + t.Fatalf("winning token result=%v", winningBody) + } + + h.proxy.mu.Lock() + _, hasBrowserState := h.proxy.browserStates[browserState] + _, hasPollDigest := h.proxy.pollDigests[digest] + h.proxy.mu.Unlock() + sess.mu.Lock() + storedResult := sess.result + sess.mu.Unlock() + if hasBrowserState || hasPollDigest || storedResult != nil { + t.Fatalf("claim left state: browser=%v poll=%v result=%v", hasBrowserState, hasPollDigest, storedResult) + } + }) + } +} + +func TestOAuthCallbackErrorsAreGenericAndOneShot(t *testing.T) { + for _, tc := range []struct { + name string + callback func(base, state string) string + wantError string + wantCBState int + }{ + { + name: "provider detail", + callback: func(base, state string) string { + return fmt.Sprintf("%s/auth/mal/callback?error=provider_detail_canary&state=%s", base, url.QueryEscape(state)) + }, + wantError: "authorization_failed", + wantCBState: http.StatusOK, + }, + { + name: "user cancelled", + callback: func(base, state string) string { + return fmt.Sprintf("%s/auth/mal/callback?error=access_denied&state=%s", base, url.QueryEscape(state)) + }, + wantError: "access_denied", + wantCBState: http.StatusOK, + }, + } { + t.Run(tc.name, func(t *testing.T) { + h := newOAuthHarness(t) + pollSecret, browserState, _ := h.startSession(t, "mal", "2.2.2.3") + resp := httpGet(t, tc.callback(h.base, browserState)) + resp.Body.Close() + if resp.StatusCode != tc.wantCBState { + t.Fatalf("callback status=%d want %d", resp.StatusCode, tc.wantCBState) + } + result := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(pollSecret)) + if result.err != nil { + t.Fatalf("result request: %v", result.err) + } + if result.status != http.StatusOK || result.body["error"] != tc.wantError { + t.Fatalf("result status=%d body=%v", result.status, result.body) + } + if result.cacheControl != "no-store, private" { + t.Fatalf("result Cache-Control=%q", result.cacheControl) + } + second := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(pollSecret)) + if second.status != http.StatusGone { + t.Fatalf("second result status=%d want 410", second.status) + } + }) + } +} + +func TestOAuthCallbackExchangeFailureIsOneShot(t *testing.T) { h := newOAuthHarness(t) h.upstream.setReply(http.StatusBadRequest, `{"error":"invalid_grant"}`) - sess, _ := h.startSession(t, "mal", "2.2.2.2") + pollSecret, browserState, _ := h.startSession(t, "mal", "2.2.2.4") - resultCh := make(chan map[string]any, 1) - go func() { - resp, err := http.Get(h.base + "/auth/result?session=" + url.QueryEscape(sess)) - if err != nil { - resultCh <- map[string]any{"_err": err.Error()} - return - } - defer resp.Body.Close() - var m map[string]any - _ = json.NewDecoder(resp.Body).Decode(&m) - resultCh <- m - }() - - resp := httpGet(t, fmt.Sprintf("%s/auth/mal/callback?code=CODE&state=%s", h.base, url.QueryEscape(sess))) + resp := httpGet(t, fmt.Sprintf("%s/auth/mal/callback?code=CODE&state=%s", h.base, url.QueryEscape(browserState))) resp.Body.Close() - - select { - case got := <-resultCh: - if got["error"] != "exchange_failed" { - t.Errorf("expected error=exchange_failed, got %v", got) - } - case <-time.After(3 * time.Second): - t.Fatal("result never returned") + if resp.StatusCode != http.StatusBadGateway { + t.Fatalf("callback status=%d want 502", resp.StatusCode) } -} - -func TestOAuthCallbackUserCancelled(t *testing.T) { - h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "2.2.2.3") - - resultCh := make(chan map[string]any, 1) - go func() { - resp, err := http.Get(h.base + "/auth/result?session=" + url.QueryEscape(sess)) - if err != nil { - resultCh <- map[string]any{"_err": err.Error()} - return - } - defer resp.Body.Close() - var m map[string]any - _ = json.NewDecoder(resp.Body).Decode(&m) - resultCh <- m - }() - - resp := httpGet(t, fmt.Sprintf("%s/auth/mal/callback?error=access_denied&state=%s", h.base, url.QueryEscape(sess))) - resp.Body.Close() - - select { - case got := <-resultCh: - if got["error"] != "access_denied" { - t.Errorf("expected error=access_denied, got %v", got) - } - case <-time.After(3 * time.Second): - t.Fatal("result never returned") + result := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(pollSecret)) + if result.status != http.StatusOK || result.body["error"] != "exchange_failed" { + t.Fatalf("result status=%d body=%v", result.status, result.body) + } + if result.cacheControl != "no-store, private" { + t.Fatalf("result Cache-Control=%q", result.cacheControl) + } + second := requestOAuthResult(h.base + "/auth/result?session=" + url.QueryEscape(pollSecret)) + if second.status != http.StatusGone { + t.Fatalf("second result status=%d want 410", second.status) } } @@ -433,45 +693,35 @@ func TestOAuthCallbackUnknownSessionIgnored(t *testing.T) { } } -func TestOAuthResultConsumedSecondCallIsGone(t *testing.T) { - h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "3.3.3.1") - - // Pre-seat the result so the first /auth/result returns immediately. - h.proxy.mu.Lock() - h.proxy.sessions[sess].complete(oauthTokenResult{AccessToken: "tok"}) - h.proxy.mu.Unlock() - - r1 := httpGet(t, h.base+"/auth/result?session="+url.QueryEscape(sess)) - r1.Body.Close() - if r1.StatusCode != http.StatusOK { - t.Fatalf("first result status=%d", r1.StatusCode) - } - // After consumption the session is deleted; second call sees unknown session. - r2 := httpGet(t, h.base+"/auth/result?session="+url.QueryEscape(sess)) - r2.Body.Close() - if r2.StatusCode != http.StatusGone { - t.Fatalf("second result status=%d want 410", r2.StatusCode) - } -} - // ====== Cleanup ====== -func TestOAuthCleanupExpiresOldSessions(t *testing.T) { +func TestOAuthCleanupRemovesBothIndexesAndSuppressesStaleCompletion(t *testing.T) { h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "4.4.4.1") + pollSecret, browserState, _ := h.startSession(t, "mal", "4.4.4.1") + digest := digestPollSecret(pollSecret) h.proxy.mu.Lock() - h.proxy.sessions[sess].createdAt = time.Now().Add(-2 * oauthSessionTTL) + sess := h.proxy.browserStates[browserState] + sess.createdAt = time.Now().Add(-2 * oauthSessionTTL) h.proxy.mu.Unlock() h.proxy.cleanup() h.proxy.mu.Lock() - _, exists := h.proxy.sessions[sess] + _, hasBrowserState := h.proxy.browserStates[browserState] + _, hasPollDigest := h.proxy.pollDigests[digest] h.proxy.mu.Unlock() - if exists { - t.Fatal("expired session should have been cleaned up") + if hasBrowserState || hasPollDigest { + t.Fatalf("expired session indexes remain: browser=%v poll=%v", hasBrowserState, hasPollDigest) + } + if h.proxy.completeSession(sess, oauthTokenResult{AccessToken: "stale-token"}) { + t.Fatal("stale callback completed a removed session") + } + sess.mu.Lock() + storedResult := sess.result + sess.mu.Unlock() + if storedResult != nil { + t.Fatal("stale callback retained an orphaned token result") } } @@ -525,7 +775,7 @@ func TestOAuthResultBlocksUntilCancel(t *testing.T) { // that /auth/result blocks until the session completes or the client // cancels. The 204-after-server-timeout path takes 50s so isn't asserted. h := newOAuthHarness(t) - sess, _ := h.startSession(t, "mal", "5.5.5.1") + sess, _, _ := h.startSession(t, "mal", "5.5.5.1") ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) defer cancel() diff --git a/server/rate_limit.go b/server/rate_limit.go index 195573df..b7fddbeb 100644 --- a/server/rate_limit.go +++ b/server/rate_limit.go @@ -16,27 +16,27 @@ type rateLimiter struct { } func newRateLimiter(burst, sustained int) *rateLimiter { + return newRateLimiterAt(burst, sustained, time.Now()) +} + +func newRateLimiterAt(burst, sustained int, now time.Time) *rateLimiter { return &rateLimiter{ tokens: float64(burst), maxTokens: float64(burst), refillRate: float64(sustained), - lastTime: time.Now(), + lastTime: now, } } func (rl *rateLimiter) allow() bool { + return rl.allowAt(time.Now()) +} + +func (rl *rateLimiter) allowAt(now time.Time) bool { rl.mu.Lock() defer rl.mu.Unlock() - now := time.Now() - elapsed := now.Sub(rl.lastTime).Seconds() - rl.lastTime = now - - rl.tokens += elapsed * rl.refillRate - if rl.tokens > rl.maxTokens { - rl.tokens = rl.maxTokens - } - + rl.refillAtLocked(now) if rl.tokens < 1 { return false } @@ -44,6 +44,27 @@ func (rl *rateLimiter) allow() bool { return true } +func (rl *rateLimiter) refund() { + rl.mu.Lock() + defer rl.mu.Unlock() + rl.tokens++ + if rl.tokens > rl.maxTokens { + rl.tokens = rl.maxTokens + } +} + +func (rl *rateLimiter) refillAtLocked(now time.Time) { + if now.Before(rl.lastTime) { + return + } + elapsed := now.Sub(rl.lastTime).Seconds() + rl.lastTime = now + rl.tokens += elapsed * rl.refillRate + if rl.tokens > rl.maxTokens { + rl.tokens = rl.maxTokens + } +} + // reclaimable reports whether discarding this limiter would preserve its // behavior: enough idle time has passed for the bucket to be full again. func (rl *rateLimiter) reclaimable(now time.Time) bool { @@ -69,6 +90,68 @@ func cleanupRateWindows(windows map[string]time.Time, now time.Time, duration ti } } +// --- Poster upload admission (global, per-IP, and concurrency) --- + +type posterUploadLimiter struct { + mu sync.Mutex + global *rateLimiter + perIP map[string]*rateLimiter + active int + maxConcurrent int + perIPBurst int + perIPSustained int +} + +func newPosterUploadLimiter( + perIPBurst, perIPSustained, globalBurst, globalSustained, maxConcurrent int, + now time.Time, +) *posterUploadLimiter { + return &posterUploadLimiter{ + global: newRateLimiterAt(globalBurst, globalSustained, now), + perIP: make(map[string]*rateLimiter), + maxConcurrent: maxConcurrent, + perIPBurst: perIPBurst, + perIPSustained: perIPSustained, + } +} + +func (pl *posterUploadLimiter) tryStart(ip string, now time.Time) bool { + pl.mu.Lock() + defer pl.mu.Unlock() + + if pl.active >= pl.maxConcurrent { + return false + } + if !pl.global.allowAt(now) { + return false + } + limiter := pl.perIP[ip] + if limiter == nil { + limiter = newRateLimiterAt(pl.perIPBurst, pl.perIPSustained, now) + pl.perIP[ip] = limiter + } + if !limiter.allowAt(now) { + pl.global.refund() + return false + } + pl.active++ + return true +} + +func (pl *posterUploadLimiter) finish() { + pl.mu.Lock() + defer pl.mu.Unlock() + if pl.active > 0 { + pl.active-- + } +} + +func (pl *posterUploadLimiter) cleanup(now time.Time) { + pl.mu.Lock() + defer pl.mu.Unlock() + cleanupRateLimiters(pl.perIP, now, nil) +} + // --- Connection tracker (per-IP limits) --- type connTracker struct { @@ -127,6 +210,9 @@ func (ct *connTracker) disconnect(ip string) { } } +// tryCreateRoom reserves capacity for a retained room created in this process. +// The reservation survives creator disconnect and is released only when the +// authoritative room is removed from Server.rooms. func (ct *connTracker) tryCreateRoom(ip string) bool { ct.mu.Lock() defer ct.mu.Unlock() @@ -137,6 +223,23 @@ func (ct *connTracker) tryCreateRoom(ip string) bool { return true } +// tryCreateRoomReplacing reserves a room while accounting for the reservation +// that removeRoomLocked will immediately release from an empty same-ID room. +// Server.mu serializes this paired reservation/removal transaction. +func (ct *connTracker) tryCreateRoomReplacing(ip, replacedOwnerKey string) bool { + ct.mu.Lock() + defer ct.mu.Unlock() + projected := ct.roomsPerIP[ip] + if replacedOwnerKey == ip { + projected-- + } + if projected >= maxRoomsPerIP { + return false + } + ct.roomsPerIP[ip]++ + return true +} + func (ct *connTracker) releaseRoom(ip string) { ct.mu.Lock() defer ct.mu.Unlock() diff --git a/server/relay_protocol_gen.go b/server/relay_protocol_gen.go index 2a23b003..b41836f0 100644 --- a/server/relay_protocol_gen.go +++ b/server/relay_protocol_gen.go @@ -3,25 +3,34 @@ package main const ( - relayTypeCreate = "create" - relayTypeJoin = "join" - relayTypeBroadcast = "broadcast" - relayTypeSendTo = "sendTo" - relayTypePing = "ping" - relayTypeCreated = "created" - relayTypeJoined = "joined" - relayTypePeerJoined = "peerJoined" - relayTypePeerLeft = "peerLeft" - relayTypeMessage = "message" - relayTypeError = "error" - relayTypePong = "pong" - relayErrorRateLimited = "rate_limited" - relayErrorInvalidMessage = "invalid_message" - relayErrorRoomExists = "room_exists" - relayErrorRoomNotFound = "room_not_found" - relayErrorRoomFull = "room_full" - relayErrorNotInRoom = "not_in_room" - relayErrorAlreadyInRoom = "already_in_room" + relayProtocolVersion = 2 + legacyRelayProtocolVersion = 0 + + relayTypeCreate = "create" + relayTypeJoin = "join" + relayTypeBroadcast = "broadcast" + relayTypeSendTo = "sendTo" + relayTypePing = "ping" + relayTypeLeave = "leave" + relayTypeEndSession = "endSession" + relayTypeCreated = "created" + relayTypeJoined = "joined" + relayTypePeerJoined = "peerJoined" + relayTypePeerLeft = "peerLeft" + relayTypeMessage = "message" + relayTypeError = "error" + relayTypePong = "pong" + relayTypeLeft = "left" + relayTypeEnded = "ended" + relayErrorRateLimited = "rate_limited" + relayErrorInvalidMessage = "invalid_message" + relayErrorRoomExists = "room_exists" + relayErrorRoomNotFound = "room_not_found" + relayErrorRoomFull = "room_full" + relayErrorNotInRoom = "not_in_room" + relayErrorAlreadyInRoom = "already_in_room" + relayErrorPeerIdUnavailable = "peer_id_unavailable" + relayErrorProtocolMismatch = "protocol_mismatch" maxRoomSize = 8 maxMessageSize = 65536 diff --git a/test/models/json_model_round_trip_test.dart b/test/models/json_model_round_trip_test.dart index bcefc5b5..541aab02 100644 --- a/test/models/json_model_round_trip_test.dart +++ b/test/models/json_model_round_trip_test.dart @@ -34,12 +34,25 @@ void main() { test('RecentRoom preserves epoch timestamp and control mode index', () { final lastUsed = DateTime.fromMillisecondsSinceEpoch(1700000000000); - final room = RecentRoom(code: 'ABCD', name: 'Movie night', lastUsed: lastUsed, controlMode: ControlMode.anyone); + final room = RecentRoom( + code: 'ABCD', + relayScope: 'scope-digest', + name: 'Movie night', + lastUsed: lastUsed, + controlMode: ControlMode.anyone, + ); - expect(room.toJson(), {'code': 'ABCD', 'name': 'Movie night', 'lastUsed': 1700000000000, 'controlMode': 1}); + expect(room.toJson(), { + 'code': 'ABCD', + 'relayScope': 'scope-digest', + 'name': 'Movie night', + 'lastUsed': 1700000000000, + 'controlMode': 1, + }); final decoded = RecentRoom.fromJson(room.toJson()); expect(decoded.code, room.code); + expect(decoded.relayScope, room.relayScope); expect(decoded.name, room.name); expect(decoded.lastUsed, lastUsed); expect(decoded.controlMode, ControlMode.anyone); @@ -47,9 +60,13 @@ void main() { test('RecentRoom omits nullable fields when absent', () { final lastUsed = DateTime.fromMillisecondsSinceEpoch(1700000000000); - final room = RecentRoom(code: 'ABCD', lastUsed: lastUsed); + final room = RecentRoom(code: 'ABCD', relayScope: 'scope-digest', lastUsed: lastUsed); - expect(room.toJson(), {'code': 'ABCD', 'lastUsed': 1700000000000}); + expect(room.toJson(), {'code': 'ABCD', 'relayScope': 'scope-digest', 'lastUsed': 1700000000000}); + }); + + test('RecentRoom requires relay scope when decoding', () { + expect(() => RecentRoom.fromJson({'code': 'ABCD', 'lastUsed': 1700000000000}), throwsA(anything)); }); test('RemoteCommand keeps compact protocol keys and unknown fallback', () { diff --git a/test/navigation/profile_session_screen_test.dart b/test/navigation/profile_session_screen_test.dart index 8063e757..559edc7e 100644 --- a/test/navigation/profile_session_screen_test.dart +++ b/test/navigation/profile_session_screen_test.dart @@ -11,6 +11,7 @@ import 'package:plezy/profiles/plex_home_service.dart'; import 'package:plezy/profiles/profile.dart'; import 'package:plezy/profiles/profile_connection_registry.dart'; import 'package:plezy/profiles/profile_registry.dart'; +import 'package:plezy/providers/companion_remote_provider.dart'; import 'package:plezy/providers/discover_provider.dart'; import 'package:plezy/providers/hidden_libraries_provider.dart'; import 'package:plezy/providers/multi_server_provider.dart'; @@ -57,11 +58,12 @@ void main() { final discoverProviders = []; final hiddenProviders = []; final trackerProviders = []; + final companionProviders = []; final disposedActiveIds = []; addTearDown(() async { await tester.pumpWidget(const SizedBox.shrink()); - await tester.pump(); + await tester.pumpAndSettle(); await activeProfile.resetForTesting(); activeProfile.dispose(); multiServer.dispose(); @@ -85,6 +87,9 @@ void main() { providers: [ Provider.value(value: storage), Provider.value(value: db), + Provider.value(value: connectionRegistry), + Provider.value(value: profileConnectionRegistry), + Provider.value(value: plexHome), ChangeNotifierProvider.value(value: activeProfile), ChangeNotifierProvider.value(value: multiServer), ChangeNotifierProvider.value(value: offlineWatch), @@ -96,6 +101,7 @@ void main() { discoverProviders: discoverProviders, hiddenProviders: hiddenProviders, trackerProviders: trackerProviders, + companionProviders: companionProviders, disposedActiveIds: disposedActiveIds, ), ), @@ -109,10 +115,12 @@ void main() { expect(discoverProviders.single.profileId, owner.id); expect(discoverProviders, hasLength(1)); expect(hiddenProviders, hasLength(1)); + expect(companionProviders, hasLength(1)); final ownerNavigator = profileNavigationRegistry.navigator; final ownerDiscover = discoverProviders.single; final ownerHidden = hiddenProviders.single; final ownerTrackers = trackerProviders.single; + final ownerCompanion = companionProviders.single; await ownerHidden.ensureInitialized(); expect(ownerHidden.profileId, owner.id); expect(ownerHidden.hiddenLibraryKeys, {'srv:owner'}); @@ -134,6 +142,9 @@ void main() { expect(trackerProviders, hasLength(2)); expect(trackerProviders.last, isNot(same(ownerTrackers))); expect(ownerTrackers.isDisposed, isTrue); + expect(companionProviders, hasLength(2)); + expect(companionProviders.last, isNot(same(ownerCompanion))); + expect(ownerCompanion.isDisposed, isTrue); await hiddenProviders.last.ensureInitialized(); expect(hiddenProviders.last.profileId, kids.id); expect(hiddenProviders.last.hiddenLibraryKeys, {'srv:kids'}); @@ -145,17 +156,27 @@ void main() { await tester.pumpAndSettle(); expect(SystemShelfService().debugActiveOwner, isNull); expect(discoverProviders.last.profileId, isNull); + + // Companion provider disposal cancels Drift-backed profile watches. Give + // their asynchronous cancellation timers a frame before test invariants + // are checked. + await tester.pumpWidget(const SizedBox.shrink()); + for (var i = 0; i < 3; i++) { + await tester.pump(const Duration(milliseconds: 1)); + } }); } class _ProfileProbeShell extends StatefulWidget { const _ProfileProbeShell({ required this.discoverProviders, + required this.companionProviders, required this.hiddenProviders, required this.disposedActiveIds, required this.trackerProviders, }); + final List companionProviders; final List discoverProviders; final List hiddenProviders; final List trackerProviders; @@ -166,6 +187,7 @@ class _ProfileProbeShell extends StatefulWidget { } class _ProfileProbeShellState extends State<_ProfileProbeShell> { + CompanionRemoteProvider? _companionProvider; DiscoverProvider? _discoverProvider; HiddenLibrariesProvider? _hiddenProvider; TrackersProvider? _trackersProvider; @@ -174,6 +196,7 @@ class _ProfileProbeShellState extends State<_ProfileProbeShell> { @override void didChangeDependencies() { super.didChangeDependencies(); + _companionProvider = context.read(); _discoverProvider = context.read(); _hiddenProvider = context.read(); _trackersProvider = context.read(); @@ -187,6 +210,9 @@ class _ProfileProbeShellState extends State<_ProfileProbeShell> { if (widget.trackerProviders.isEmpty || !identical(widget.trackerProviders.last, _trackersProvider)) { widget.trackerProviders.add(_trackersProvider!); } + if (widget.companionProviders.isEmpty || !identical(widget.companionProviders.last, _companionProvider)) { + widget.companionProviders.add(_companionProvider!); + } } @override diff --git a/test/providers/companion_remote_provider_test.dart b/test/providers/companion_remote_provider_test.dart index dafb20e1..44db7c3a 100644 --- a/test/providers/companion_remote_provider_test.dart +++ b/test/providers/companion_remote_provider_test.dart @@ -1,3 +1,5 @@ +import 'dart:async'; + import 'package:drift/native.dart'; import 'package:flutter_test/flutter_test.dart'; import 'package:plezy/connection/connection.dart'; @@ -15,7 +17,9 @@ import 'package:plezy/profiles/profile_connection.dart'; import 'package:plezy/profiles/profile_connection_registry.dart'; import 'package:plezy/profiles/profile_registry.dart'; import 'package:plezy/providers/companion_remote_provider.dart'; -import 'package:plezy/services/base_peer_service.dart'; +import 'package:plezy/services/companion_remote/companion_remote_peer_service.dart'; +import 'package:plezy/services/companion_remote/lan_discovery_service.dart'; +import 'package:plezy/services/companion_remote/remote_auth_context.dart'; import 'package:plezy/services/companion_remote/remote_auth_service.dart'; import 'package:plezy/services/storage_service.dart'; @@ -121,6 +125,325 @@ void main() { }); }); + group('CompanionRemoteProvider — reconnect ownership', () { + test('newest overlapping connect owns the session and disposes the stale candidate once', () async { + final firstJoinGate = Completer(); + final firstDisposalGate = Completer(); + final first = _FakeCompanionRemotePeerService(joinGate: firstJoinGate, disconnectGate: firstDisposalGate); + final second = _FakeCompanionRemotePeerService(); + final factory = _FakePeerFactory([first, second]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + + final olderConnect = harness.provider.connectToManualHost('192.0.2.20:48634'); + await first.joinStarted.future; + final newerConnect = harness.provider.connectToManualHost('192.0.2.21:48634'); + await first.disconnectStarted.future; + firstJoinGate.complete(); + await Future.delayed(Duration.zero); + expect(first.disposeCalls, 1); + + firstDisposalGate.complete(); + await second.joinStarted.future; + await newerConnect; + + harness.provider.sendCommand(RemoteCommandType.playPause); + expect(harness.provider.status, RemoteSessionStatus.connected); + expect(second.sentCommands.single.type, RemoteCommandType.playPause); + + await olderConnect; + + expect(harness.provider.status, RemoteSessionStatus.connected); + expect(first.disposeCalls, 1); + expect(first.disconnectCalls, 1); + expect(first.hasListeners, isFalse); + expect(second.disposeCalls, 0); + expect(factory.created, 2); + }); + + test('replacement disconnect reconnects while intentional predecessor teardown is blocked', () async { + final predecessorDisposalGate = Completer(); + final predecessor = _FakeCompanionRemotePeerService(disconnectGate: predecessorDisposalGate); + final replacement = _FakeCompanionRemotePeerService(); + final factory = _FakePeerFactory([predecessor, replacement]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + + Future? predecessorTeardown; + addTearDown(() async { + if (!predecessorDisposalGate.isCompleted) { + predecessorDisposalGate.complete(); + } + final teardown = predecessorTeardown; + if (teardown != null) await teardown; + }); + + await harness.provider.connectToManualHost('192.0.2.22:48634'); + + var teardownCompleted = false; + predecessorTeardown = harness.provider.leaveSession().whenComplete(() { + teardownCompleted = true; + }); + await predecessor.disconnectStarted.future; + expect(teardownCompleted, isFalse); + + await harness.provider.connectToManualHost('192.0.2.23:48634'); + expect(harness.provider.status, RemoteSessionStatus.connected); + expect(replacement.hasListeners, isTrue); + + final publishedStatuses = []; + void captureStatus() => publishedStatuses.add(harness.provider.status); + harness.provider.addListener(captureStatus); + + predecessor.emitError( + RemotePeerError(type: RemotePeerErrorType.connectionFailed, message: 'stale predecessor error'), + ); + predecessor.emitDeviceDisconnected(); + + expect(publishedStatuses, isEmpty); + expect(harness.provider.status, RemoteSessionStatus.connected); + expect(harness.provider.session?.errorMessage, isNull); + expect(harness.provider.reconnectAttempts, 0); + + replacement.emitDeviceDisconnected(); + + expect(publishedStatuses, [RemoteSessionStatus.reconnecting]); + expect(harness.provider.status, RemoteSessionStatus.reconnecting); + expect(harness.provider.reconnectAttempts, 1); + expect(teardownCompleted, isFalse); + + predecessorDisposalGate.complete(); + await predecessorTeardown; + + expect(publishedStatuses, [RemoteSessionStatus.reconnecting]); + expect(harness.provider.status, RemoteSessionStatus.reconnecting); + expect(harness.provider.reconnectAttempts, 1); + + harness.provider.removeListener(captureStatus); + await harness.provider.cancelReconnect(); + }); + + test('cancelReconnect fully tears down a running host and discovery', () async { + final hostDisposalGate = Completer(); + final host = _FakeCompanionRemotePeerService(disconnectGate: hostDisposalGate); + final discovery = _FakeLanDiscoveryService(); + final factory = _FakePeerFactory([host]); + final harness = await _RemoteHarness.create(factory.call, discoveryServiceFactory: () => discovery); + addTearDown(harness.close); + + await harness.provider.startHostServer(); + expect(harness.provider.isHostServerRunning, isTrue); + expect(harness.provider.debugIsDiscoveryBroadcasting, isTrue); + expect(host.hasListeners, isTrue); + + final cancellation = harness.provider.cancelReconnect(); + await host.disconnectStarted.future; + expect(harness.provider.debugIsDiscoveryBroadcasting, isFalse); + expect(host.hasListeners, isFalse); + hostDisposalGate.complete(); + await cancellation; + + expect(harness.provider.session, isNull); + expect(harness.provider.isHostServerRunning, isFalse); + expect(harness.provider.debugIsDiscoveryBroadcasting, isFalse); + expect(harness.provider.debugIsDiscoveryListening, isFalse); + expect(discovery.stopBroadcastingCalls, 1); + expect(discovery.stopListeningCalls, 1); + expect(host.disposeCalls, 1); + expect(host.disconnectCalls, 1); + expect(host.hasListeners, isFalse); + }); + + test('cancel during join keeps disconnected state and disposes the candidate once', () async { + final initial = _FakeCompanionRemotePeerService(); + final joinGate = Completer(); + final candidate = _FakeCompanionRemotePeerService(joinGate: joinGate); + final factory = _FakePeerFactory([initial, candidate]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.10:48634'); + + final statuses = []; + harness.provider.addListener(() => statuses.add(harness.provider.status)); + initial.emitDeviceDisconnected(); + expect(harness.provider.status, RemoteSessionStatus.reconnecting); + + final retry = harness.provider.retryReconnectNow(); + await candidate.joinStarted.future; + expect(candidate.hasListeners, isTrue); + + final cancellation = harness.provider.cancelReconnect(); + expect(harness.provider.status, RemoteSessionStatus.disconnected); + await cancellation; + final statusCountAfterCancel = statuses.length; + + joinGate.complete(); + await retry; + await Future.delayed(Duration.zero); + + expect(harness.provider.status, RemoteSessionStatus.disconnected); + expect(statuses.skip(statusCountAfterCancel), isNot(contains(RemoteSessionStatus.connected))); + expect(candidate.disposeCalls, 1); + expect(candidate.disconnectCalls, 1); + expect(candidate.hasListeners, isFalse); + expect(factory.created, 2); + expect(harness.provider.reconnectAttempts, 0); + }); + + test('leave during join clears the session and rejects late commands', () async { + final initial = _FakeCompanionRemotePeerService(); + final joinGate = Completer(); + final candidate = _FakeCompanionRemotePeerService(joinGate: joinGate); + final factory = _FakePeerFactory([initial, candidate]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.11:48634'); + + var deliveredCommands = 0; + harness.provider.onCommandReceived = (_) => deliveredCommands++; + initial.emitDeviceDisconnected(); + final retry = harness.provider.retryReconnectNow(); + await candidate.joinStarted.future; + + await harness.provider.leaveSession(); + expect(harness.provider.session, isNull); + expect(candidate.hasListeners, isFalse); + candidate.emitCommand(const RemoteCommand(type: RemoteCommandType.playPause)); + joinGate.complete(); + await retry; + + expect(harness.provider.session, isNull); + expect(deliveredCommands, 0); + expect(candidate.disposeCalls, 1); + expect(factory.created, 2); + expect(harness.provider.reconnectAttempts, 0); + }); + + for (final action in ['cancel', 'leave']) { + test('$action before candidate creation prevents replacement allocation', () async { + final disconnectGate = Completer(); + final initial = _FakeCompanionRemotePeerService(disconnectGate: disconnectGate); + final replacement = _FakeCompanionRemotePeerService(); + final factory = _FakePeerFactory([initial, replacement]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.12:48634'); + + initial.emitDeviceDisconnected(); + final retry = harness.provider.retryReconnectNow(); + await initial.disconnectStarted.future; + + if (action == 'cancel') { + await harness.provider.cancelReconnect(); + } else { + await harness.provider.leaveSession(); + } + disconnectGate.complete(); + await retry; + + expect(factory.created, 1); + expect(replacement.joinStarted.isCompleted, isFalse); + expect(initial.disposeCalls, 1); + }); + } + + test('dispose during join detaches listeners and prevents a late commit', () async { + final initial = _FakeCompanionRemotePeerService(); + final joinGate = Completer(); + final candidate = _FakeCompanionRemotePeerService(joinGate: joinGate); + final factory = _FakePeerFactory([initial, candidate]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.13:48634'); + + var deliveredCommands = 0; + harness.provider.onCommandReceived = (_) => deliveredCommands++; + initial.emitDeviceDisconnected(); + final retry = harness.provider.retryReconnectNow(); + await candidate.joinStarted.future; + + harness.provider.dispose(); + expect(harness.provider.isDisposed, isTrue); + expect(candidate.hasListeners, isFalse); + candidate.emitCommand(const RemoteCommand(type: RemoteCommandType.playPause)); + joinGate.complete(); + await retry; + await Future.delayed(Duration.zero); + + expect(harness.provider.status, isNot(RemoteSessionStatus.connected)); + expect(deliveredCommands, 0); + expect(candidate.disposeCalls, 1); + expect(factory.created, 2); + }); + + test('logout invalidates a held reconnect before clearing identity', () async { + final initial = _FakeCompanionRemotePeerService(); + final joinGate = Completer(); + final candidate = _FakeCompanionRemotePeerService(joinGate: joinGate); + final factory = _FakePeerFactory([initial, candidate]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.14:48634'); + + initial.emitDeviceDisconnected(); + final retry = harness.provider.retryReconnectNow(); + await candidate.joinStarted.future; + + await harness.provider.resetForLogout(); + expect(harness.provider.session, isNull); + expect(harness.provider.isCryptoReady, isFalse); + expect(harness.provider.debugCryptoConnectionId, isNull); + expect(candidate.hasListeners, isFalse); + + joinGate.complete(); + await retry; + + expect(harness.provider.session, isNull); + expect(harness.provider.isCryptoReady, isFalse); + expect(candidate.disposeCalls, 1); + expect(factory.created, 2); + expect(harness.provider.reconnectAttempts, 0); + }); + + test('current reconnect failure is contained and schedules one retry', () async { + final initial = _FakeCompanionRemotePeerService(); + final candidate = _FakeCompanionRemotePeerService(joinError: StateError('synthetic reconnect failure')); + final factory = _FakePeerFactory([initial, candidate]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.15:48634'); + + initial.emitDeviceDisconnected(); + await harness.provider.retryReconnectNow(); + + expect(harness.provider.status, RemoteSessionStatus.reconnecting); + expect(harness.provider.reconnectAttempts, 1); + expect(candidate.disposeCalls, 1); + expect(candidate.hasListeners, isFalse); + + await harness.provider.cancelReconnect(); + }); + + test('successful current reconnect preserves connected behavior', () async { + final initial = _FakeCompanionRemotePeerService(); + final candidate = _FakeCompanionRemotePeerService(); + final factory = _FakePeerFactory([initial, candidate]); + final harness = await _RemoteHarness.create(factory.call); + addTearDown(harness.close); + await harness.provider.connectToManualHost('192.0.2.16:48634'); + + initial.emitDeviceDisconnected(); + await harness.provider.retryReconnectNow(); + harness.provider.sendCommand(RemoteCommandType.playPause); + + expect(harness.provider.status, RemoteSessionStatus.connected); + expect(harness.provider.reconnectAttempts, 0); + expect(candidate.sentCommands.single.type, RemoteCommandType.playPause); + expect(initial.disposeCalls, 1); + expect(candidate.disposeCalls, 0); + }); + }); + group('CompanionRemoteProvider — public API safety', () { test('connectToDiscoveredHost reports localized auth failure when crypto is not ready', () async { final p = CompanionRemoteProvider(); @@ -546,3 +869,264 @@ PlexHomeUser _homeUser(String uuid, {required bool admin}) { protected: false, ); } + +class _FakePeerFactory { + _FakePeerFactory(this.peers); + + final List<_FakeCompanionRemotePeerService> peers; + int created = 0; + + CompanionRemotePeerService call() { + if (created >= peers.length) { + throw StateError('Unexpected peer allocation'); + } + return peers[created++]; + } +} + +class _FakeCompanionRemotePeerService extends CompanionRemotePeerService { + _FakeCompanionRemotePeerService({this.joinGate, this.disconnectGate, this.joinError}); + + final Completer? joinGate; + final Completer? disconnectGate; + final Object? joinError; + final Completer joinStarted = Completer(); + final Completer disconnectStarted = Completer(); + final List sentCommands = []; + + final StreamController _commands = StreamController.broadcast(sync: true); + final StreamController _connected = StreamController.broadcast(sync: true); + final StreamController _disconnected = StreamController.broadcast(sync: true); + final StreamController _errors = StreamController.broadcast(sync: true); + final StreamController _statuses = StreamController.broadcast(sync: true); + + int disconnectCalls = 0; + int disposeCalls = 0; + bool _streamsClosed = false; + bool _serverRunning = false; + + @override + bool get isServerRunning => _serverRunning; + + bool get hasListeners => + _commands.hasListener || + _connected.hasListener || + _disconnected.hasListener || + _errors.hasListener || + _statuses.hasListener; + + @override + Stream get onCommandReceived => _commands.stream; + + @override + Stream get onDeviceConnected => _connected.stream; + + @override + Stream get onDeviceDisconnected => _disconnected.stream; + + @override + Stream get onError => _errors.stream; + + @override + Stream get onConnectionStateChanged => _statuses.stream; + + @override + String? get selectedAuthContextId => 'auth-context'; + + @override + String? get selectedHostClientId => 'host-client'; + + @override + Future<({List addresses, int port})> createSessionForContexts( + String deviceName, + String platform, + List authContexts, + ) async { + _serverRunning = true; + return (addresses: const ['127.0.0.1:48634'], port: 48634); + } + + @override + Future joinSessionWithContexts( + String deviceName, + String platform, + String hostAddress, + List authContexts, { + String? authContextId, + String expectedHostClientId = '', + }) async { + if (!joinStarted.isCompleted) joinStarted.complete(); + final gate = joinGate; + if (gate != null) await gate.future; + final error = joinError; + if (error != null) throw error; + } + + @override + Future joinSessionRacingWithContexts( + String deviceName, + String platform, + List hostAddresses, + List authContexts, { + String? authContextId, + String expectedHostClientId = '', + }) async { + await joinSessionWithContexts( + deviceName, + platform, + hostAddresses.first, + authContexts, + authContextId: authContextId, + expectedHostClientId: expectedHostClientId, + ); + return hostAddresses.first; + } + + @override + void sendCommand(RemoteCommand command) { + sentCommands.add(command); + } + + void emitDeviceDisconnected() { + if (!_streamsClosed) _disconnected.add(null); + } + + void emitCommand(RemoteCommand command) { + if (!_streamsClosed) _commands.add(command); + } + + void emitError(RemotePeerError error) { + if (!_streamsClosed) _errors.add(error); + } + + @override + Future disconnect() async { + disconnectCalls++; + _serverRunning = false; + if (!disconnectStarted.isCompleted) disconnectStarted.complete(); + final gate = disconnectGate; + if (gate != null) await gate.future; + } + + @override + Future dispose() async { + disposeCalls++; + await disconnect(); + if (_streamsClosed) return; + _streamsClosed = true; + await Future.wait([ + _commands.close(), + _connected.close(), + _disconnected.close(), + _errors.close(), + _statuses.close(), + ]); + } +} + +class _FakeLanDiscoveryService extends LanDiscoveryService { + bool _broadcasting = false; + bool _listening = false; + int stopBroadcastingCalls = 0; + int stopListeningCalls = 0; + + @override + bool get isBroadcasting => _broadcasting; + + @override + bool get isListening => _listening; + + @override + Future startBroadcastingForContexts({ + required List contexts, + required String deviceName, + required String platform, + required int wsPort, + required List ips, + }) async { + _broadcasting = true; + } + + @override + Future stopBroadcasting() async { + stopBroadcastingCalls++; + _broadcasting = false; + } + + @override + void stopListening() { + stopListeningCalls++; + _listening = false; + } +} + +class _RemoteHarness { + _RemoteHarness({required this.provider, required this.database, required this.activeProfile, required this.plexHome}); + + final CompanionRemoteProvider provider; + final AppDatabase database; + final ActiveProfileProvider activeProfile; + final PlexHomeService plexHome; + bool _closed = false; + + static Future<_RemoteHarness> create( + CompanionRemotePeerServiceFactory peerServiceFactory, { + LanDiscoveryServiceFactory discoveryServiceFactory = LanDiscoveryService.new, + }) async { + final database = AppDatabase.forTesting(NativeDatabase.memory()); + final connections = ConnectionRegistry(database); + final profileConnections = ProfileConnectionRegistry(database); + final profiles = ProfileRegistry(database); + final storage = await StorageService.getInstance(); + final plexHome = PlexHomeService( + connections: connections, + profileConnections: profileConnections, + storage: storage, + plexHomeUserFetcher: (_) async => const [], + ); + final activeProfile = ActiveProfileProvider( + registry: profiles, + plexHome: plexHome, + connections: connections, + storage: storage, + ); + + final account = _plexAccount('remote-account', 'remote-client'); + final profile = _localProfile('remote-profile'); + await connections.upsert(account); + await profiles.upsert(profile); + await profileConnections.upsert( + ProfileConnection(profileId: profile.id, connectionId: account.id, userIdentifier: 'remote-admin'), + makeDefault: true, + ); + await storage.setActiveProfileId(profile.id); + await activeProfile.initialize(); + + final provider = CompanionRemoteProvider.forTesting( + peerServiceFactory: peerServiceFactory, + discoveryServiceFactory: discoveryServiceFactory, + ); + final ready = await provider.ensureCryptoReady( + _home('remote-admin'), + connections: connections, + activeProfile: activeProfile, + profileConnections: profileConnections, + account: account, + ); + if (!ready) { + throw StateError('Remote test harness failed to initialize crypto'); + } + + return _RemoteHarness(provider: provider, database: database, activeProfile: activeProfile, plexHome: plexHome); + } + + Future close() async { + if (_closed) return; + _closed = true; + if (!provider.isDisposed) provider.dispose(); + await activeProfile.resetForTesting(); + activeProfile.dispose(); + await plexHome.dispose(); + await database.close(); + } +} diff --git a/test/screens/settings/logs_screen_test.dart b/test/screens/settings/logs_screen_test.dart index 1f6e15a3..3b07d9f4 100644 --- a/test/screens/settings/logs_screen_test.dart +++ b/test/screens/settings/logs_screen_test.dart @@ -1,9 +1,29 @@ import 'dart:convert'; +import 'package:flutter/material.dart'; +import 'package:flutter/services.dart'; import 'package:flutter_test/flutter_test.dart'; +import 'package:http/http.dart' as http; +import 'package:http/testing.dart'; +import 'package:package_info_plus/package_info_plus.dart'; +import 'package:plezy/focus/input_mode_tracker.dart'; +import 'package:plezy/i18n/strings.g.dart'; import 'package:plezy/screens/settings/logs_screen.dart'; +import 'package:plezy/utils/app_logger.dart'; +import 'package:plezy/utils/media_server_http_client.dart'; void main() { + setUpAll(() { + LocaleSettings.setLocaleSync(AppLocale.en); + }); + + setUp(() { + MemoryLogOutput.clearLogs(); + }); + + tearDown(() { + MemoryLogOutput.clearLogs(); + }); test('log upload payload preserves the header and newest complete lines', () { const header = 'Plezy test device\n---\n'; const logs = 'oldest line that should be removed\nmiddle line that should be removed\nnewest 🚀 line'; @@ -22,4 +42,108 @@ void main() { expect(constrainLogUploadPayload(header: header, logs: logs, maxBytes: 128), '$header$logs'); }); + + testWidgets('long upload capability displays and copies at narrow width', (tester) async { + const capability = 'abcdefghijklmnopqrstuvwxy'; + String? uploadedBody; + http.Request? uploadedRequest; + String? clipboardText; + + PackageInfo.setMockInitialValues( + appName: 'Plezy', + packageName: 'com.plezy.test', + version: '1.2.3', + buildNumber: '45', + buildSignature: '', + ); + const deviceInfoChannel = MethodChannel('dev.fluttercommunity.plus/device_info'); + final messenger = TestDefaultBinaryMessengerBinding.instance.defaultBinaryMessenger; + messenger.setMockMethodCallHandler(deviceInfoChannel, (call) async { + expect(call.method, 'getDeviceInfo'); + return { + 'computerName': 'test-mac', + 'hostName': 'test-mac.local', + 'arch': 'arm64', + 'model': 'Mac15,3', + 'modelName': 'Mac', + 'kernelVersion': 'test', + 'osRelease': '15.0', + 'majorVersion': 15, + 'minorVersion': 0, + 'patchVersion': 0, + 'activeCPUs': 8, + 'memorySize': 16 * 1024 * 1024 * 1024, + 'cpuFrequency': 0, + 'systemGUID': 'test-guid', + }; + }); + messenger.setMockMethodCallHandler(SystemChannels.platform, (call) async { + if (call.method == 'Clipboard.setData') { + clipboardText = (call.arguments as Map)['text'] as String?; + } + return null; + }); + addTearDown(() { + messenger.setMockMethodCallHandler(deviceInfoChannel, null); + messenger.setMockMethodCallHandler(SystemChannels.platform, null); + }); + + final client = MediaServerHttpClient( + client: MockClient((request) async { + uploadedRequest = request; + uploadedBody = request.body; + return http.Response( + jsonEncode({'id': capability}), + 200, + headers: {'content-type': 'application/json'}, + ); + }), + ); + addTearDown(client.close); + appLogger.i('widget upload seed'); + + tester.view.physicalSize = const Size(320, 640); + tester.view.devicePixelRatio = 1; + addTearDown(tester.view.resetPhysicalSize); + addTearDown(tester.view.resetDevicePixelRatio); + await tester.pumpWidget( + TranslationProvider( + child: InputModeTracker( + child: MaterialApp(home: LogsScreen(httpClient: client)), + ), + ), + ); + await tester.pumpAndSettle(); + + await tester.tap(find.byTooltip(t.logs.uploadLogs)); + await tester.pumpAndSettle(); + + expect(tester.takeException(), isNull); + expect(find.text(capability, findRichText: true), findsOneWidget); + expect(utf8.encode(uploadedBody!).length, lessThanOrEqualTo(maxLogUploadBytes)); + expect(uploadedRequest!.method, 'POST'); + expect(uploadedRequest!.url, Uri.parse('https://ice.plezy.app/logs')); + final contentType = uploadedRequest!.headers.entries + .singleWhere((entry) => entry.key.toLowerCase() == 'content-type') + .value; + expect(contentType, startsWith('text/plain')); + expect(uploadedBody, contains('widget upload seed')); + + final dialog = find.byType(AlertDialog); + final copyButton = find.descendant(of: dialog, matching: find.byType(IconButton)); + expect(copyButton, findsOneWidget); + await tester.tap(copyButton); + await tester.pump(); + expect(clipboardText, capability); + + clipboardText = null; + await tester.sendKeyDownEvent(LogicalKeyboardKey.shiftLeft); + await tester.sendKeyEvent(LogicalKeyboardKey.tab); + await tester.sendKeyUpEvent(LogicalKeyboardKey.shiftLeft); + await tester.pump(); + await tester.sendKeyEvent(LogicalKeyboardKey.enter); + await tester.pump(); + expect(clipboardText, capability); + expect(tester.takeException(), isNull); + }); } diff --git a/test/screens/settings/settings_screen_test.dart b/test/screens/settings/settings_screen_test.dart index dc0f1c32..d34e7d3a 100644 --- a/test/screens/settings/settings_screen_test.dart +++ b/test/screens/settings/settings_screen_test.dart @@ -226,6 +226,39 @@ void main() { expect(materialUpdateTile.focusNode, isNotNull); }); + testWidgets('relay dialog rejects invalid bases without persisting them', (tester) async { + final harness = await _pumpSettingsScreen(tester); + addTearDown(() => harness.dispose(tester)); + + await tester.tap(find.text(t.settings.watchTogetherRelay)); + await _pumpUi(tester); + await tester.enterText(find.byType(TextField), 'ws://opaque-user@relay.example.test/path?wrong=route'); + await tester.tap(find.text(t.common.save)); + await _pumpUi(tester); + + expect(find.byType(AlertDialog), findsOneWidget); + expect(find.text(t.settings.watchTogetherRelayInvalid), findsOneWidget); + expect(SettingsService.instance.read(SettingsService.customRelayUrl), isNull); + }); + + testWidgets('relay dialog saves and reopens the canonical base', (tester) async { + final harness = await _pumpSettingsScreen(tester); + addTearDown(() => harness.dispose(tester)); + + await tester.tap(find.text(t.settings.watchTogetherRelay)); + await _pumpUi(tester); + await tester.enterText(find.byType(TextField), ' HTTP://Relay.Example.Test:8080/prefix/// '); + await tester.tap(find.text(t.common.save)); + await _pumpUi(tester); + + expect(find.byType(AlertDialog), findsNothing); + expect(SettingsService.instance.read(SettingsService.customRelayUrl), 'http://relay.example.test:8080/prefix'); + + await tester.tap(find.text(t.settings.watchTogetherRelay)); + await _pumpUi(tester); + expect(find.widgetWithText(TextField, 'http://relay.example.test:8080/prefix'), findsOneWidget); + }); + testWidgets('folder replacement uses the provider coordinator', (tester) async { final selectedDirectory = Directory('${temporaryDirectory.path}/selected-downloads'); directoryPicker.directoryPath = selectedDirectory.path; diff --git a/test/services/companion_remote_peer_service_test.dart b/test/services/companion_remote_peer_service_test.dart index 28d6f680..856fe7fe 100644 --- a/test/services/companion_remote_peer_service_test.dart +++ b/test/services/companion_remote_peer_service_test.dart @@ -1,8 +1,15 @@ +import 'dart:async'; +import 'dart:convert'; +import 'dart:io'; + import 'package:flutter_test/flutter_test.dart'; import 'package:plezy/i18n/strings.g.dart'; import 'package:plezy/models/companion_remote/remote_command.dart'; import 'package:plezy/services/companion_remote/companion_remote_peer_service.dart'; import 'package:plezy/services/companion_remote/remote_auth_context.dart'; +import 'package:plezy/services/companion_remote/remote_auth_service.dart'; + +const _ioTimeout = Duration(seconds: 5); void main() { test('auth precondition error uses the active locale', () async { @@ -17,6 +24,428 @@ void main() { ); }); + test('per-source admission rejects before upgrade and recovers exactly one slot', () async { + final host = CompanionRemotePeerService.forTesting(maxTotalHostConnections: 4, maxHostConnectionsPerSource: 2); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + final port = await _startHost(host); + + clients.add(await _RawWebSocketClient.connect(port)); + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + expect(clients, everyElement(predicate<_RawWebSocketClient>((client) => client.isOpen))); + + await clients.removeAt(0).dispose(); + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + }); + + test('global admission rejects independently and releases one reservation once', () async { + final host = CompanionRemotePeerService.forTesting(maxTotalHostConnections: 2, maxHostConnectionsPerSource: 4); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + final port = await _startHost(host); + + clients.add(await _RawWebSocketClient.connect(port)); + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + + await clients.removeAt(0).dispose(); + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + }); + + test('authentication timeout releases capacity for a legitimate encrypted session', () async { + final host = CompanionRemotePeerService.forTesting( + maxTotalHostConnections: 1, + maxHostConnectionsPerSource: 1, + authTimeout: const Duration(milliseconds: 300), + ); + final remote = CompanionRemotePeerService(); + final clients = <_RawWebSocketClient>[]; + addTearDown(() async { + await remote.dispose(); + await _disposeHostAndClients(host, clients); + }); + final port = await _startHost(host); + + final silent = await _RawWebSocketClient.connect(port); + clients.add(silent); + expect(await silent.closed.timeout(_ioTimeout), 4001); + clients.remove(silent); + + await remote.joinSessionWithContexts( + 'Test Remote', + 'ios', + '127.0.0.1:$port', + [_authContext], + authContextId: _authContext.id, + expectedHostClientId: _authContext.clientIdentifier, + ); + final command = host.onCommandReceived.firstWhere((event) => event.type == RemoteCommandType.play); + remote.sendCommand(const RemoteCommand(type: RemoteCommandType.play)); + expect((await command.timeout(_ioTimeout)).type, RemoteCommandType.play); + }); + + test('disconnect closes pending sockets and restart restores full capacity', () async { + final host = CompanionRemotePeerService.forTesting( + maxTotalHostConnections: 2, + maxHostConnectionsPerSource: 2, + authTimeout: const Duration(seconds: 1), + ); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + var port = await _startHost(host); + + final firstGeneration = [await _RawWebSocketClient.connect(port), await _RawWebSocketClient.connect(port)]; + clients.addAll(firstGeneration); + await host.disconnect(); + expect( + await Future.wait(firstGeneration.map((client) => client.closed.timeout(_ioTimeout))), + everyElement(isNot(4001)), + ); + clients.removeWhere(firstGeneration.contains); + + port = await _startHost(host); + clients.add(await _RawWebSocketClient.connect(port)); + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + }); + + test('failed HTTP upgrade and abrupt socket loss both release the only slot', () async { + final host = CompanionRemotePeerService.forTesting(maxTotalHostConnections: 1, maxHostConnectionsPerSource: 1); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + final port = await _startHost(host); + + final httpClient = HttpClient(); + addTearDown(() => httpClient.close(force: true)); + final request = await httpClient.getUrl(Uri.parse('http://127.0.0.1:$port/ws')); + final response = await request.close(); + expect(response.statusCode, HttpStatus.badRequest); + await response.drain(); + + final afterFailedUpgrade = await _RawWebSocketClient.connect(port); + clients.add(afterFailedUpgrade); + await afterFailedUpgrade.dispose(); + clients.remove(afterFailedUpgrade); + + final abrupt = await _RawUpgradedSocket.connect(port); + addTearDown(abrupt.dispose); + await abrupt.dispose(); + await _flushEventQueue(); + + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + }); + + test('setup failure after upgrade closes the socket and releases the only slot', () async { + var rejectSetup = true; + final host = CompanionRemotePeerService.forTesting( + maxTotalHostConnections: 1, + maxHostConnectionsPerSource: 1, + afterHostUpgrade: () { + if (!rejectSetup) return; + rejectSetup = false; + throw StateError('injected admission setup failure'); + }, + ); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + final port = await _startHost(host); + + final rejected = await WebSocket.connect('ws://127.0.0.1:$port/ws'); + try { + await rejected.drain().timeout(_ioTimeout); + expect(rejected.readyState, WebSocket.closed); + } finally { + await rejected.close(); + } + + clients.add(await _RawWebSocketClient.connect(port)); + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + }); + + test('oversized and malformed first messages close generically and always restore capacity', () async { + final host = CompanionRemotePeerService.forTesting( + maxTotalHostConnections: 1, + maxHostConnectionsPerSource: 1, + maxPreAuthMessageBytes: 256, + maxFailedAuthAttempts: 20, + ); + final clients = <_RawWebSocketClient>[]; + final connected = []; + final commands = []; + final connectedSubscription = host.onDeviceConnected.listen(connected.add); + final commandSubscription = host.onCommandReceived.listen(commands.add); + addTearDown(() async { + await connectedSubscription.cancel(); + await commandSubscription.cancel(); + await _disposeHostAndClients(host, clients); + }); + final port = await _startHost(host); + + final cases = <({String name, dynamic payload, int closeCode})>[ + (name: 'oversized text', payload: 'x' * 257, closeCode: 4003), + (name: 'oversized UTF-8 text', payload: 'é' * 129, closeCode: 4003), + (name: 'binary input', payload: [1, 2, 3], closeCode: 4003), + (name: 'invalid JSON', payload: '{', closeCode: 4003), + (name: 'non-map JSON', payload: '[]', closeCode: 4003), + (name: 'invalid Base64', payload: _minimalAuth(clientNonce: '%%%'), closeCode: 4003), + (name: 'wrong nonce length', payload: _minimalAuth(clientNonce: base64Encode([1, 2])), closeCode: 4003), + (name: 'missing fields', payload: jsonEncode({'type': 'auth'}), closeCode: 4003), + ( + name: 'wrong field type', + payload: jsonEncode({ + 'type': 'auth', + 'authTag': 1, + 'clientNonce': base64Encode(List.filled(32, 1)), + 'userUUID': 'user-1', + 'clientIdentifier': 'remote-client', + 'deviceName': 'Remote', + 'platform': 'test', + }), + closeCode: 4003, + ), + (name: 'unexpected type', payload: jsonEncode({'type': 'ping'}), closeCode: 4002), + ]; + + for (final malformedCase in cases) { + final client = await _RawWebSocketClient.connect(port); + clients.add(client); + client.add(malformedCase.payload); + await client.authFailed.timeout(_ioTimeout); + expect(await client.closed.timeout(_ioTimeout), malformedCase.closeCode, reason: malformedCase.name); + clients.remove(client); + } + + final capacityProbe = await _RawWebSocketClient.connect(port); + clients.add(capacityProbe); + expect(capacityProbe.isOpen, isTrue); + expect(connected, isEmpty); + expect(commands, isEmpty); + }); + + test('malformed first message counts once, ignores a queued valid auth, and locks out before upgrade', () async { + final host = CompanionRemotePeerService.forTesting( + maxTotalHostConnections: 2, + maxHostConnectionsPerSource: 2, + maxFailedAuthAttempts: 2, + ); + final clients = <_RawWebSocketClient>[]; + final connected = []; + final subscription = host.onDeviceConnected.listen(connected.add); + addTearDown(() async { + await subscription.cancel(); + await _disposeHostAndClients(host, clients); + }); + final port = await _startHost(host); + + final first = await _RawWebSocketClient.connect(port); + clients.add(first); + final validAuth = _validAuthMessage(first.challenge, _authContext); + first.add('{'); + first.add(validAuth); + await first.authFailed.timeout(_ioTimeout); + expect(await first.closed.timeout(_ioTimeout), 4003); + clients.remove(first); + + final second = await _RawWebSocketClient.connect(port); + clients.add(second); + second.add([9]); + await second.authFailed.timeout(_ioTimeout); + expect(await second.closed.timeout(_ioTimeout), 4003); + clients.remove(second); + + await _expectWebSocketHandshakeStatus(port, HttpStatus.tooManyRequests); + expect(connected, isEmpty); + }); + + test('unknown context, disallowed user, and bad HMAC share the terminal auth failure path', () async { + final host = CompanionRemotePeerService.forTesting(maxFailedAuthAttempts: 10); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + final port = await _startHost(host); + + final attempts = [ + (client) => _validAuthMessage(client.challenge, _authContext, authContextId: 'missing-context'), + (client) => _validAuthMessage(client.challenge, _authContext, userUuid: 'disallowed-user'), + (client) => _validAuthMessage(client.challenge, _authContext, authTag: 'invalid-tag'), + ]; + + for (final message in attempts) { + final client = await _RawWebSocketClient.connect(port); + clients.add(client); + client.add(message(client)); + await client.authFailed.timeout(_ioTimeout); + expect(await client.closed.timeout(_ioTimeout), 4003); + clients.remove(client); + } + }); + + test('single-context auth remains compatible without authContextId', () async { + final host = CompanionRemotePeerService.forTesting(); + final clients = <_RawWebSocketClient>[]; + addTearDown(() => _disposeHostAndClients(host, clients)); + final port = await _startHost(host); + + final client = await _RawWebSocketClient.connect(port); + clients.add(client); + final connected = host.onDeviceConnected.first; + client.add(_validAuthMessage(client.challenge, _authContext, includeAuthContextId: false)); + + await connected.timeout(_ioTimeout); + expect(host.isConnected, isTrue); + expect(host.selectedAuthContextId, _authContext.id); + }); + + test('closed authentication race cannot commit and replacement late close cannot clear the winner', () async { + final secondDerivationEntered = Completer(); + final releaseSecondDerivation = Completer(); + final secondDerivationReturned = Completer(); + var derivationCount = 0; + final host = CompanionRemotePeerService.forTesting( + deriveSessionEncKey: (homeSecret, hostNonce, clientNonce) async { + derivationCount++; + if (derivationCount == 2) { + secondDerivationEntered.complete(); + await releaseSecondDerivation.future; + } + final key = await RemoteAuthService.instance.deriveSessionEncKey(homeSecret, hostNonce, clientNonce); + if (derivationCount == 2 && !secondDerivationReturned.isCompleted) { + secondDerivationReturned.complete(); + } + return key; + }, + ); + final clients = <_RawWebSocketClient>[]; + final connected = []; + final subscription = host.onDeviceConnected.listen(connected.add); + addTearDown(() async { + if (!releaseSecondDerivation.isCompleted) releaseSecondDerivation.complete(); + await subscription.cancel(); + await _disposeHostAndClients(host, clients); + }); + final port = await _startHost(host); + + final original = await _RawWebSocketClient.connect(port); + clients.add(original); + final originalConnected = host.onDeviceConnected.first; + original.add(_validAuthMessage(original.challenge, _authContext, deviceName: 'Original')); + await originalConnected.timeout(_ioTimeout); + expect(host.isConnected, isTrue); + + final racing = await _RawWebSocketClient.connect(port); + clients.add(racing); + racing.add(_validAuthMessage(racing.challenge, _authContext, deviceName: 'Closed racer')); + await secondDerivationEntered.future.timeout(_ioTimeout); + await racing.dispose(); + clients.remove(racing); + releaseSecondDerivation.complete(); + await secondDerivationReturned.future.timeout(_ioTimeout); + await _flushEventQueue(); + + expect(host.isConnected, isTrue); + expect(original.isOpen, isTrue); + expect(connected, hasLength(1)); + + final replacement = await _RawWebSocketClient.connect(port); + clients.add(replacement); + final replacementConnected = host.onDeviceConnected.first; + replacement.add(_validAuthMessage(replacement.challenge, _authContext, deviceName: 'Replacement')); + await replacementConnected.timeout(_ioTimeout); + expect(await original.closed.timeout(_ioTimeout), 4004); + clients.remove(original); + await _flushEventQueue(); + + expect(host.isConnected, isTrue); + expect(connected, hasLength(2)); + }); + + test('address probes release every source slot before the managed connection', () async { + final host = CompanionRemotePeerService.forTesting(maxTotalHostConnections: 2, maxHostConnectionsPerSource: 2); + final remote = CompanionRemotePeerService(); + addTearDown(() async { + await remote.dispose(); + await host.dispose(); + }); + final port = await _startHost(host); + final address = '127.0.0.1:$port'; + + final winner = await remote.joinSessionRacingWithContexts( + 'Test Remote', + 'ios', + [address, address], + [_authContext], + authContextId: _authContext.id, + expectedHostClientId: _authContext.clientIdentifier, + ); + + expect(winner, address); + final command = host.onCommandReceived.firstWhere((event) => event.type == RemoteCommandType.play); + remote.sendCommand(const RemoteCommand(type: RemoteCommandType.play)); + expect((await command.timeout(_ioTimeout)).type, RemoteCommandType.play); + }); + + test('managed join waits for every probe terminal event and listener cleanup', () async { + final closeRequested = List.generate(2, (_) => Completer()); + final listenerCleanedUp = List.generate(2, (_) => Completer()); + final controllers = List.generate(2, (_) => StreamController(sync: true)); + var nextProbe = 0; + final managedJoinStarted = Completer(); + late final _ManagedJoinObservingPeer remote; + remote = _ManagedJoinObservingPeer( + raceProbeFactory: (_) { + final index = nextProbe++; + return ( + close: () async { + if (!closeRequested[index].isCompleted) { + closeRequested[index].complete(); + } + }, + ready: Future.value(), + stream: _CancelTrackingStream(controllers[index].stream, () { + if (!listenerCleanedUp[index].isCompleted) { + listenerCleanedUp[index].complete(); + } + }), + ); + }, + onManagedJoin: () async { + expect(listenerCleanedUp, everyElement(predicate>((item) => item.isCompleted))); + managedJoinStarted.complete(); + }, + ); + addTearDown(() async { + for (final controller in controllers) { + if (!controller.isClosed) await controller.close(); + } + await remote.dispose(); + }); + + final join = remote.joinSessionRacingWithContexts( + 'Test Remote', + 'ios', + const ['probe-a', 'probe-b'], + [_authContext], + ); + expect(nextProbe, 2); + + controllers.first.add(jsonEncode({'type': 'challenge'})); + await closeRequested.first.future.timeout(_ioTimeout); + expect(managedJoinStarted.isCompleted, isFalse); + + await controllers.first.close(); + await closeRequested.last.future.timeout(_ioTimeout); + expect(listenerCleanedUp.first.isCompleted, isTrue); + expect(managedJoinStarted.isCompleted, isFalse); + + await controllers.last.close(); + expect(await join.timeout(_ioTimeout), 'probe-a'); + expect(listenerCleanedUp, everyElement(predicate>((item) => item.isCompleted))); + expect(managedJoinStarted.isCompleted, isTrue); + }); + test('host and remote dispatch encrypted commands through the same contract', () async { final host = CompanionRemotePeerService(); final remote = CompanionRemotePeerService(); @@ -25,39 +454,320 @@ void main() { await host.dispose(); }); - final context = RemoteAuthContext( - id: 'context-1', - backend: 'plex', - connectionId: 'connection-1', - homeSecret: List.generate(32, (index) => index), - discoveryKey: List.generate(32, (index) => 255 - index), - clientIdentifier: 'host-client', - userUuid: 'user-1', - allowedUserUuids: const ['user-1'], - ); - - final session = await host.createSessionForContexts('Test Host', 'macos', [context]); + final session = await host.createSessionForContexts('Test Host', 'macos', [_authContext]); await remote.joinSessionWithContexts( 'Test Remote', 'ios', '127.0.0.1:${session.port}', - [context], - authContextId: context.id, - expectedHostClientId: context.clientIdentifier, + [_authContext], + authContextId: _authContext.id, + expectedHostClientId: _authContext.clientIdentifier, ); final hostCommand = host.onCommandReceived.firstWhere((command) => command.type == RemoteCommandType.play); remote.sendCommand(const RemoteCommand(type: RemoteCommandType.play, data: {'source': 'remote'})); expect( - await hostCommand.timeout(const Duration(seconds: 5)), + await hostCommand.timeout(_ioTimeout), const RemoteCommand(type: RemoteCommandType.play, data: {'source': 'remote'}), ); final remoteCommand = remote.onCommandReceived.firstWhere((command) => command.type == RemoteCommandType.pause); host.sendCommand(const RemoteCommand(type: RemoteCommandType.pause, data: {'source': 'host'})); expect( - await remoteCommand.timeout(const Duration(seconds: 5)), + await remoteCommand.timeout(_ioTimeout), const RemoteCommand(type: RemoteCommandType.pause, data: {'source': 'host'}), ); }); } + +typedef _TestRaceProbeConnection = ({Future Function() close, Future ready, Stream stream}); + +class _ManagedJoinObservingPeer extends CompanionRemotePeerService { + _ManagedJoinObservingPeer({ + required _TestRaceProbeConnection Function(Uri uri) raceProbeFactory, + required this._onManagedJoin, + }) : super.forTesting(raceProbeFactory: raceProbeFactory); + + final Future Function() _onManagedJoin; + + @override + Future joinSessionWithContexts( + String deviceName, + String platform, + String hostAddress, + List authContexts, { + String? authContextId, + String expectedHostClientId = '', + }) { + return _onManagedJoin(); + } +} + +class _CancelTrackingStream extends Stream { + _CancelTrackingStream(this._delegate, this._onCancel); + + final Stream _delegate; + final void Function() _onCancel; + + @override + StreamSubscription listen( + void Function(T event)? onData, { + Function? onError, + void Function()? onDone, + bool? cancelOnError, + }) { + return _CancelTrackingSubscription( + _delegate.listen(onData, onError: onError, onDone: onDone, cancelOnError: cancelOnError), + _onCancel, + ); + } +} + +class _CancelTrackingSubscription implements StreamSubscription { + _CancelTrackingSubscription(this._delegate, this._onCancel); + + final StreamSubscription _delegate; + final void Function() _onCancel; + + @override + Future cancel() { + _onCancel(); + return _delegate.cancel(); + } + + @override + void onData(void Function(T data)? handleData) => _delegate.onData(handleData); + + @override + void onError(Function? handleError) => _delegate.onError(handleError); + + @override + void onDone(void Function()? handleDone) => _delegate.onDone(handleDone); + + @override + void pause([Future? resumeSignal]) => _delegate.pause(resumeSignal); + + @override + void resume() => _delegate.resume(); + + @override + bool get isPaused => _delegate.isPaused; + + @override + Future asFuture([E? futureValue]) => _delegate.asFuture(futureValue); +} + +final _authContext = RemoteAuthContext( + id: 'context-1', + backend: 'plex', + connectionId: 'connection-1', + homeSecret: List.generate(32, (index) => index), + discoveryKey: List.generate(32, (index) => 255 - index), + clientIdentifier: 'host-client', + userUuid: 'user-1', + allowedUserUuids: const ['user-1'], +); + +Future _startHost(CompanionRemotePeerService host) async { + final session = await host.createSessionForContexts('Test Host', 'macos', [_authContext]); + return session.port; +} + +Future _disposeHostAndClients(CompanionRemotePeerService host, List<_RawWebSocketClient> clients) async { + for (final client in List<_RawWebSocketClient>.of(clients)) { + await client.dispose(); + } + clients.clear(); + await host.dispose(); +} + +Future _expectWebSocketHandshakeStatus(int port, int statusCode) async { + await expectLater( + WebSocket.connect('ws://127.0.0.1:$port/ws'), + throwsA( + isA().having((error) => error.toString(), 'message', contains('status code: $statusCode')), + ), + ); +} + +String _minimalAuth({required String clientNonce}) { + return jsonEncode({ + 'type': 'auth', + 'authTag': 'tag', + 'clientNonce': clientNonce, + 'userUUID': 'user-1', + 'clientIdentifier': 'remote-client', + 'deviceName': 'Remote', + 'platform': 'test', + }); +} + +String _validAuthMessage( + Map challenge, + RemoteAuthContext context, { + String? authContextId, + String? userUuid, + String? authTag, + String deviceName = 'Raw Remote', + bool includeAuthContextId = true, +}) { + final hostNonce = base64Decode(challenge['nonce'] as String); + final clientNonce = RemoteAuthService.instance.generateNonce(); + final selectedUserUuid = userUuid ?? context.userUuid; + final computedAuthTag = RemoteAuthService.instance.computeAuthTag( + homeSecret: context.homeSecret, + hostNonce: hostNonce, + clientNonce: clientNonce, + hostClientId: context.clientIdentifier, + userUUID: selectedUserUuid, + clientIdentifier: 'raw-remote-client', + deviceName: deviceName, + platform: 'test', + ); + return jsonEncode({ + 'type': 'auth', + if (includeAuthContextId) 'authContextId': authContextId ?? context.id, + 'clientNonce': base64Encode(clientNonce), + 'userUUID': selectedUserUuid, + 'clientIdentifier': 'raw-remote-client', + 'deviceName': deviceName, + 'platform': 'test', + 'authTag': authTag ?? computedAuthTag, + }); +} + +Future _flushEventQueue() async { + await Future.delayed(Duration.zero); + await Future.delayed(Duration.zero); +} + +class _RawWebSocketClient { + _RawWebSocketClient._(this.socket) { + _subscription = socket.listen( + (data) { + if (data is! String) return; + try { + final decoded = jsonDecode(data); + if (decoded is! Map) return; + if (decoded['type'] == 'challenge' && !_challenge.isCompleted) { + _challengeValue = decoded; + _challenge.complete(decoded); + } else if (decoded['type'] == 'authFailed' && !_authFailed.isCompleted) { + _authFailed.complete(); + } + } catch (_) { + // Encrypted and malformed server messages are not part of this helper. + } + }, + onDone: () { + if (!_closed.isCompleted) _closed.complete(socket.closeCode); + }, + onError: (Object error, StackTrace stackTrace) { + if (!_closed.isCompleted) _closed.completeError(error, stackTrace); + }, + cancelOnError: true, + ); + } + + static Future<_RawWebSocketClient> connect(int port) async { + final socket = await WebSocket.connect('ws://127.0.0.1:$port/ws').timeout(_ioTimeout); + final client = _RawWebSocketClient._(socket); + try { + await client._challenge.future.timeout(_ioTimeout); + return client; + } catch (_) { + await client.dispose(); + rethrow; + } + } + + final WebSocket socket; + final Completer> _challenge = Completer>(); + final Completer _authFailed = Completer(); + final Completer _closed = Completer(); + late final StreamSubscription _subscription; + late Map _challengeValue; + + Map get challenge => _challengeValue; + Future get authFailed => _authFailed.future; + Future get closed => _closed.future; + bool get isOpen => socket.readyState == WebSocket.open; + + void add(dynamic data) => socket.add(data); + + Future dispose() async { + if (!_closed.isCompleted) { + try { + await socket.close(WebSocketStatus.normalClosure); + } catch (_) { + // The peer may already have closed the transport. + } + } + if (!_closed.isCompleted) { + await _closed.future.timeout(_ioTimeout); + } + await _subscription.cancel(); + } +} + +class _RawUpgradedSocket { + _RawUpgradedSocket._(this.socket, this._subscription, this._done); + + static Future<_RawUpgradedSocket> connect(int port) async { + final socket = await Socket.connect(InternetAddress.loopbackIPv4, port).timeout(_ioTimeout); + final responseHeaders = StringBuffer(); + final upgraded = Completer(); + final done = Completer(); + late final StreamSubscription> subscription; + subscription = socket.listen( + (bytes) { + responseHeaders.write(latin1.decode(bytes)); + if (!upgraded.isCompleted && responseHeaders.toString().contains('\r\n\r\n')) { + upgraded.complete(); + } + }, + onError: (Object error, StackTrace stackTrace) { + if (!upgraded.isCompleted) upgraded.completeError(error, stackTrace); + if (!done.isCompleted) done.complete(); + }, + onDone: () { + if (!upgraded.isCompleted) { + upgraded.completeError(StateError('Socket closed before WebSocket upgrade')); + } + if (!done.isCompleted) done.complete(); + }, + ); + try { + final key = base64Encode(List.generate(16, (index) => index)); + socket.write( + 'GET /ws HTTP/1.1\r\n' + 'Host: 127.0.0.1:$port\r\n' + 'Connection: Upgrade\r\n' + 'Upgrade: websocket\r\n' + 'Sec-WebSocket-Version: 13\r\n' + 'Sec-WebSocket-Key: $key\r\n' + '\r\n', + ); + await socket.flush(); + await upgraded.future.timeout(_ioTimeout); + expect(responseHeaders.toString(), startsWith('HTTP/1.1 101')); + return _RawUpgradedSocket._(socket, subscription, done); + } catch (_) { + socket.destroy(); + await subscription.cancel(); + rethrow; + } + } + + final Socket socket; + final StreamSubscription> _subscription; + final Completer _done; + + Future dispose() async { + if (!_done.isCompleted) { + socket.destroy(); + await _done.future.timeout(_ioTimeout); + } + await _subscription.cancel(); + } +} diff --git a/test/services/settings_export_service_test.dart b/test/services/settings_export_service_test.dart index 1cc28ad3..b24fc8b4 100644 --- a/test/services/settings_export_service_test.dart +++ b/test/services/settings_export_service_test.dart @@ -90,6 +90,7 @@ void main() { await prefs.setString('custom_relay_url', 'https://${canaries[2]}.invalid'); await prefs.setString('update_last_check_time', canaries[3]); await prefs.setString('watch_together_recent_rooms', canaries[1]); + await prefs.setString('user_alice_watch_together_recent_rooms', canaries[1]); await prefs.setString('future_runtime_key', 'unknown'); final out = SettingsExportService.buildExportMap(prefs, currentUserUuid: 'alice'); @@ -105,6 +106,7 @@ void main() { expect(exported.keys, isNot(contains('custom_relay_url'))); expect(exported.keys, isNot(contains('update_last_check_time'))); expect(exported.keys, isNot(contains('watch_together_recent_rooms'))); + expect(exported.keys, isNot(contains('user_alice_watch_together_recent_rooms'))); expect(exported.keys, isNot(contains('future_runtime_key'))); for (final canary in canaries) { expect(encoded, isNot(contains(canary))); @@ -191,6 +193,8 @@ void main() { 'custom_download_path': {'type': 'string', 'value': '/source/device/downloads'}, 'custom_download_path_type': {'type': 'string', 'value': 'saf'}, 'crash_reporting': {'type': 'bool', 'value': true}, + 'custom_relay_url': {'type': 'string', 'value': 'https://relay.example.test'}, + 'watch_together_recent_rooms': {'type': 'string', 'value': '[]'}, 'unknown_future_key': {'type': 'bool', 'value': true}, }, }, @@ -199,7 +203,7 @@ void main() { ); expect(result.keysImported, 4); - expect(result.keysSkipped, 6); + expect(result.keysSkipped, 8); expect(prefs.getBool('enable_hardware_decoding'), isTrue); expect(prefs.getDouble('default_playback_speed'), 1.0); expect(prefs.getStringList('user_alice_library_order'), ['movies']); @@ -210,6 +214,8 @@ void main() { expect(prefs.getString('custom_download_path'), '/target/device/downloads'); expect(prefs.getString('custom_download_path_type'), 'file'); expect(prefs.getBool('crash_reporting'), isFalse); + expect(prefs.getString('custom_relay_url'), isNull); + expect(prefs.getString('user_alice_watch_together_recent_rooms'), isNull); expect(prefs.getBool('unknown_future_key'), isNull); }); diff --git a/test/services/settings_service_test.dart b/test/services/settings_service_test.dart index 07fd67e7..5c0e9385 100644 --- a/test/services/settings_service_test.dart +++ b/test/services/settings_service_test.dart @@ -1,3 +1,5 @@ +import 'dart:async'; + import 'package:flutter/services.dart'; import 'package:flutter_test/flutter_test.dart'; import 'package:plezy/models/audio_quality_preset.dart'; @@ -226,6 +228,40 @@ void main() { }); }); + group('SettingsService Watch Together relay', () { + test('typed writes canonicalize valid bases and reject invalid replacement', () async { + final settings = await SettingsService.getInstance(); + + await settings.write(SettingsService.customRelayUrl, ' HTTPS://Relay.Example.Test/path/// '); + expect(settings.read(SettingsService.customRelayUrl), 'https://relay.example.test/path'); + + await expectLater( + settings.write(SettingsService.customRelayUrl, 'ws://relay.example.test'), + throwsFormatException, + ); + expect(settings.read(SettingsService.customRelayUrl), 'https://relay.example.test/path'); + + await settings.write(SettingsService.customRelayUrl, ' '); + expect(settings.read(SettingsService.customRelayUrl), isNull); + }); + + test('startup canonicalizes valid history and removes invalid history', () async { + resetSharedPreferencesForTest( + initialAsync: {SettingsService.customRelayUrl.key: ' http://Relay.Example.Test:8080/prefix// '}, + ); + SettingsService.resetForTesting(); + var settings = await SettingsService.getInstance(); + expect(settings.read(SettingsService.customRelayUrl), 'http://relay.example.test:8080/prefix'); + + resetSharedPreferencesForTest( + initialAsync: {SettingsService.customRelayUrl.key: 'https://relay.example.test/path?wrong=route'}, + ); + SettingsService.resetForTesting(); + settings = await SettingsService.getInstance(); + expect(settings.read(SettingsService.customRelayUrl), isNull); + }); + }); + group('SettingsService listenables', () { test('refreshListenables updates active prefs outside the resettable surface', () async { final settings = await SettingsService.getInstance(); @@ -276,4 +312,48 @@ void main() { expect(settings.isLibraryAllowedForTracker(TrackerService.trakt, 'server:allowed'), isTrue); }); }); + group('BaseSharedPreferencesService initialization generations', () { + test('a reset-raced initialization resolves to the replacement backend', () async { + resetSharedPreferencesForTest(initialAsync: const {'generation_marker': 'old'}); + final firstOnInit = Completer(); + final firstStarted = Completer(); + var constructions = 0; + + final raced = BaseSharedPreferencesService.initializeInstance<_GenerationTestPreferences>(() { + final construction = constructions++; + return _GenerationTestPreferences( + construction, + onInitStarted: construction == 0 ? firstStarted : null, + onInitGate: construction == 0 ? firstOnInit : null, + ); + }); + await firstStarted.future; + + resetSharedPreferencesForTest(initialAsync: const {'generation_marker': 'new'}); + final replacement = await BaseSharedPreferencesService.initializeInstance<_GenerationTestPreferences>( + () => _GenerationTestPreferences(constructions++), + ); + firstOnInit.complete(); + final recovered = await raced; + + expect(recovered, same(replacement)); + expect(recovered.construction, 1); + expect(recovered.readString('generation_marker'), 'new'); + expect(constructions, 2); + }); + }); +} + +class _GenerationTestPreferences extends BaseSharedPreferencesService { + _GenerationTestPreferences(this.construction, {this.onInitStarted, this.onInitGate}); + + final int construction; + final Completer? onInitStarted; + final Completer? onInitGate; + + @override + Future onInit() async { + onInitStarted?.complete(); + await onInitGate?.future; + } } diff --git a/test/test_helpers/watch_together_fakes.dart b/test/test_helpers/watch_together_fakes.dart index 7d889b11..a8b71f87 100644 --- a/test/test_helpers/watch_together_fakes.dart +++ b/test/test_helpers/watch_together_fakes.dart @@ -3,6 +3,7 @@ import 'dart:async'; import 'package:plezy/mpv/mpv.dart'; import 'package:plezy/watch_together/models/sync_message.dart'; import 'package:plezy/watch_together/services/watch_together_peer_service.dart'; +import 'package:plezy/watch_together/services/watch_together_relay_endpoint.dart'; /// Rich fake [Player] for Watch Together sync tests. /// @@ -196,6 +197,9 @@ class FakeRelayHub { final Map _peers = {}; HubPeerService register(String peerId) { + if (_peers.containsKey(peerId)) { + throw StateError('Peer ID is already registered: $peerId'); + } final service = HubPeerService._(peerId, this); for (final existing in _peers.values) { existing._peerConnected.add(peerId); @@ -235,7 +239,7 @@ class FakeRelayHub { } class HubPeerService extends WatchTogetherPeerService { - HubPeerService._(this.peerId, this._hub) : super(customBaseUrl: 'http://localhost'); + HubPeerService._(this.peerId, this._hub) : super(endpoint: WatchTogetherRelayEndpoint.resolve('http://localhost')); final String peerId; final FakeRelayHub _hub; diff --git a/test/watch_together/host_playback_coordinator_test.dart b/test/watch_together/host_playback_coordinator_test.dart index e592fe50..9638abc5 100644 --- a/test/watch_together/host_playback_coordinator_test.dart +++ b/test/watch_together/host_playback_coordinator_test.dart @@ -16,9 +16,11 @@ class _Harness { FakeAsync async, { ControlMode controlMode = ControlMode.hostOnly, HostCoordinatorCallbacks callbacks = const HostCoordinatorCallbacks(), + Duration duration = const Duration(minutes: 45), + bool seekable = true, }) { int nowMs() => _epochMs + async.elapsed.inMilliseconds; - player = FakeSyncPlayer(position: const Duration(minutes: 2)); + player = FakeSyncPlayer(position: const Duration(minutes: 2), duration: duration, seekable: seekable); coordinator = HostPlaybackCoordinator( myPeerId: 'host', controlMode: controlMode, @@ -446,6 +448,132 @@ void main() { }); }); + test('invalid remote seeks and rates are rejected without side effects', () { + fakeAsync((async) { + final actions = <(String, PlaybackActionHint)>[]; + final h = _Harness( + async, + controlMode: ControlMode.anyone, + callbacks: HostCoordinatorCallbacks(onRemoteAction: (peer, hint) => actions.add((peer, hint))), + ); + h.attachForMedia(async); + h.hostBecomesReady(async); + final durationMs = h.player.state.duration.inMilliseconds; + final broadcastsBefore = h.broadcasts.length; + final seqBefore = h.last.seq; + final anchorBefore = h.last.anchorPositionMs; + final rateBefore = h.last.rate; + h.player.commandLog.clear(); + + for (final targetMs in [-1, durationMs + 1]) { + h.coordinator.onControlRequest('guest', ControlRequest(kind: ControlRequestKind.seek, positionMs: targetMs)); + } + for (final rate in [0.25 - 0.000001, 4.0 + 0.000001, double.nan, double.infinity, double.negativeInfinity]) { + h.coordinator.onControlRequest('guest', ControlRequest(kind: ControlRequestKind.rate, rate: rate)); + } + async.flushMicrotasks(); + + expect(h.player.commandLog, isEmpty); + expect(actions, isEmpty); + expect(h.broadcasts, hasLength(broadcastsBefore)); + expect(h.last.seq, seqBefore); + expect(h.last.anchorPositionMs, anchorBefore); + expect(h.last.rate, rateBefore); + h.dispose(); + }); + }); + + test('inclusive remote seek and rate boundaries are applied exactly', () { + fakeAsync((async) { + final actions = <(String, PlaybackActionHint)>[]; + final h = _Harness( + async, + controlMode: ControlMode.anyone, + callbacks: HostCoordinatorCallbacks(onRemoteAction: (peer, hint) => actions.add((peer, hint))), + ); + h.attachForMedia(async); + h.hostBecomesReady(async); + final durationMs = h.player.state.duration.inMilliseconds; + final seqBefore = h.last.seq; + h.player.commandLog.clear(); + + for (final targetMs in [0, durationMs]) { + h.coordinator.onControlRequest('guest', ControlRequest(kind: ControlRequestKind.seek, positionMs: targetMs)); + async.flushMicrotasks(); + expect(h.last.anchorPositionMs, targetMs); + expect(h.last.actorPeerId, 'guest'); + expect(h.last.actionHint, PlaybackActionHint.seek); + } + for (final rate in [0.25, 4.0]) { + h.coordinator.onControlRequest('guest', ControlRequest(kind: ControlRequestKind.rate, rate: rate)); + async.flushMicrotasks(); + expect(h.last.rate, rate); + expect(h.last.actorPeerId, 'guest'); + expect(h.last.actionHint, PlaybackActionHint.rate); + } + + expect(h.player.commandLog, ['seek:0', 'seek:$durationMs', 'rate:0.25', 'rate:4.0']); + expect(h.last.seq, seqBefore + 4); + expect(actions, [ + ('guest', PlaybackActionHint.seek), + ('guest', PlaybackActionHint.seek), + ('guest', PlaybackActionHint.rate), + ('guest', PlaybackActionHint.rate), + ]); + h.dispose(); + }); + }); + + test('remote seeks require seekability and a positive known duration', () { + fakeAsync((async) { + for (final config in [ + (seekable: false, duration: const Duration(minutes: 45)), + (seekable: true, duration: Duration.zero), + ]) { + final actions = <(String, PlaybackActionHint)>[]; + final h = _Harness( + async, + controlMode: ControlMode.anyone, + duration: config.duration, + seekable: config.seekable, + callbacks: HostCoordinatorCallbacks(onRemoteAction: (peer, hint) => actions.add((peer, hint))), + ); + h.attachForMedia(async); + h.hostBecomesReady(async); + async.elapse(Duration(milliseconds: h.last.anchorHostTimeMs - (_epochMs + async.elapsed.inMilliseconds))); + h.player.commandLog.clear(); + actions.clear(); + final broadcastsBefore = h.broadcasts.length; + final seqBefore = h.last.seq; + + h.coordinator.onControlRequest( + 'guest', + const ControlRequest(kind: ControlRequestKind.seek, positionMs: 1000), + ); + async.flushMicrotasks(); + + expect(h.player.commandLog, isEmpty); + expect(actions, isEmpty); + expect(h.broadcasts, hasLength(broadcastsBefore)); + expect(h.last.seq, seqBefore); + + h.coordinator.onControlRequest('guest', const ControlRequest(kind: ControlRequestKind.rate, rate: 0.25)); + h.coordinator.onControlRequest('guest', const ControlRequest(kind: ControlRequestKind.pause)); + h.coordinator.onControlRequest('guest', const ControlRequest(kind: ControlRequestKind.play)); + async.flushMicrotasks(); + async.elapse(Duration(milliseconds: h.last.anchorHostTimeMs - (_epochMs + async.elapsed.inMilliseconds))); + + expect(h.player.commandLog, ['rate:0.25', 'pause', 'play']); + expect(actions, [ + ('guest', PlaybackActionHint.rate), + ('guest', PlaybackActionHint.pause), + ('guest', PlaybackActionHint.play), + ]); + h.dispose(); + } + }); + }); + test('local seeks debounce into a single re-anchor broadcast', () { fakeAsync((async) { final h = _Harness(async); diff --git a/test/watch_together/playback_state_test.dart b/test/watch_together/playback_state_test.dart index 1bd5e9d7..2f89018e 100644 --- a/test/watch_together/playback_state_test.dart +++ b/test/watch_together/playback_state_test.dart @@ -112,10 +112,11 @@ void main() { }); }); - group('SyncMessage v2 envelope', () { - test('join carries the protocol version', () { + group('SyncMessage v3 envelope', () { + test('join carries sync protocol 3', () { final join = SyncMessage.join(peerId: 'p', displayName: 'Name', isHost: false); final decoded = SyncMessage.fromJson(join.toJson()); + expect(decoded.version, 3); expect(decoded.version, SyncMessage.protocolVersion); }); @@ -125,7 +126,7 @@ void main() { expect(decoded.peerId, 'p'); }); - test('copyWith preserves v2 payloads', () { + test('copyWith preserves v3 payloads', () { final relabeled = SyncMessage.state(fullState).copyWith(peerId: 'relay-id'); expect(relabeled.state, fullState); expect(relabeled.peerId, 'relay-id'); diff --git a/test/watch_together/primitives_test.dart b/test/watch_together/primitives_test.dart index 100012b2..d86c2d3f 100644 --- a/test/watch_together/primitives_test.dart +++ b/test/watch_together/primitives_test.dart @@ -2,12 +2,6 @@ import 'package:flutter_test/flutter_test.dart'; import 'package:plezy/watch_together/primitives.dart'; void main() { - test('stored room codes preserve the established host peer wire format', () { - const persistedSessionId = 'Ab12z'; - - expect(watchTogetherHostPeerId(persistedSessionId), 'wt-AB12Z'); - }); - test('orderedStringListsEqual preserves order and multiplicity', () { expect(orderedStringListsEqual(const ['a', 'b'], const ['a', 'b']), isTrue); expect(orderedStringListsEqual(const ['a'], const ['a', 'b']), isFalse); diff --git a/test/watch_together/recent_rooms_service_test.dart b/test/watch_together/recent_rooms_service_test.dart new file mode 100644 index 00000000..9c1ae298 --- /dev/null +++ b/test/watch_together/recent_rooms_service_test.dart @@ -0,0 +1,116 @@ +import 'dart:convert'; + +import 'package:flutter_test/flutter_test.dart'; +import 'package:plezy/services/settings_service.dart'; +import 'package:plezy/watch_together/models/watch_session.dart'; +import 'package:plezy/watch_together/services/recent_rooms_service.dart'; +import 'package:plezy/watch_together/services/watch_together_relay_endpoint.dart'; + +import '../test_helpers/prefs.dart'; + +void main() { + late SettingsService settings; + final endpointA = WatchTogetherRelayEndpoint.resolve('https://relay-a.example.test/base/'); + final endpointB = WatchTogetherRelayEndpoint.resolve('http://relay-b.example.test:8080'); + + setUp(() async { + resetSharedPreferencesForTest(); + SettingsService.resetForTesting(); + settings = await SettingsService.getInstance(); + }); + + test('rooms are isolated by profile', () async { + await RecentRoomsService.addOrUpdateRoom( + 'ROOM1', + profileId: 'profile-a', + endpoint: endpointA, + name: 'Profile A room', + controlMode: ControlMode.hostOnly, + ); + + expect(RecentRoomsService.getRecentRooms(profileId: 'profile-b', endpoint: endpointA), isEmpty); + + await RecentRoomsService.addOrUpdateRoom( + 'ROOM1', + profileId: 'profile-b', + endpoint: endpointA, + name: 'Profile B room', + controlMode: ControlMode.anyone, + ); + + final profileA = RecentRoomsService.getRecentRooms(profileId: 'profile-a', endpoint: endpointA); + final profileB = RecentRoomsService.getRecentRooms(profileId: 'profile-b', endpoint: endpointA); + expect(profileA.single.name, 'Profile A room'); + expect(profileA.single.controlMode, ControlMode.hostOnly); + expect(profileB.single.name, 'Profile B room'); + expect(profileB.single.controlMode, ControlMode.anyone); + }); + + test('same code is independently mutable on different relay bases', () async { + await RecentRoomsService.addOrUpdateRoom( + 'SAME1', + profileId: 'profile-a', + endpoint: endpointA, + name: 'Relay A', + controlMode: ControlMode.hostOnly, + ); + await RecentRoomsService.addOrUpdateRoom( + 'SAME1', + profileId: 'profile-a', + endpoint: endpointB, + name: 'Relay B', + controlMode: ControlMode.anyone, + ); + + await RecentRoomsService.renameRoom('SAME1', 'Renamed A', profileId: 'profile-a', endpoint: endpointA); + expect(RecentRoomsService.getRecentRooms(profileId: 'profile-a', endpoint: endpointA).single.name, 'Renamed A'); + expect(RecentRoomsService.getRecentRooms(profileId: 'profile-a', endpoint: endpointB).single.name, 'Relay B'); + + await RecentRoomsService.removeRoom('SAME1', profileId: 'profile-a', endpoint: endpointA); + expect(RecentRoomsService.getRecentRooms(profileId: 'profile-a', endpoint: endpointA), isEmpty); + expect(RecentRoomsService.getRecentRooms(profileId: 'profile-a', endpoint: endpointB).single.code, 'SAME1'); + }); + + test('profile history remains bounded across relay scopes and deduplicates tuples', () async { + for (var index = 0; index < 21; index++) { + await RecentRoomsService.addOrUpdateRoom( + 'R${index.toString().padLeft(4, '0')}', + profileId: 'profile-a', + endpoint: index.isEven ? endpointA : endpointB, + ); + await Future.delayed(const Duration(milliseconds: 2)); + } + + final raw = settings.read(SettingsService.recentRoomsForProfile('profile-a')); + final rows = jsonDecode(raw!) as List; + expect(rows, hasLength(20)); + expect(rows.map((row) => (row as Map)['code']), isNot(contains('R0000'))); + + await RecentRoomsService.addOrUpdateRoom('R0020', profileId: 'profile-a', endpoint: endpointA, name: 'Updated'); + final updatedRows = jsonDecode(settings.read(SettingsService.recentRoomsForProfile('profile-a'))!) as List; + expect(updatedRows, hasLength(20)); + expect( + updatedRows.where( + (row) => + (row as Map)['code'] == 'R0020' && + row['relayScope'] == RecentRoomsService.relayScopeFor(endpointA), + ), + hasLength(1), + ); + }); + + test('startup drops unattributable legacy history', () async { + SettingsService.resetForTesting(); + resetSharedPreferencesForTest( + initialAsync: { + 'watch_together_recent_rooms': jsonEncode([ + {'code': 'OLD01', 'lastUsed': 1}, + ]), + }, + ); + + final initialized = await SettingsService.getInstance(); + expect(initialized.prefs.containsKey('watch_together_recent_rooms'), isFalse); + expect(RecentRoomsService.getRecentRooms(profileId: 'profile-a', endpoint: endpointA), isEmpty); + }); +} diff --git a/test/watch_together/watch_together_controller_test.dart b/test/watch_together/watch_together_controller_test.dart index 70de40c6..c12b6680 100644 --- a/test/watch_together/watch_together_controller_test.dart +++ b/test/watch_together/watch_together_controller_test.dart @@ -209,6 +209,74 @@ void main() { }); }); + test('anyone-mode: invalid controls are rejected and the queue continues', () { + fakeAsync((async) { + final room = _Room(async, controlMode: ControlMode.anyone); + final actions = <(String, PlaybackActionHint)>[]; + room.host.onRemoteAction = (peer, hint) => actions.add((peer, hint)); + room.hostStartsMedia(); + room.guestJoinsMedia(); + room.bothBecomeReady(); + final delay = room.lastHostState().anchorHostTimeMs - room.nowMs(); + async.elapse(Duration(milliseconds: delay + 100)); + room.hostPlayer.commandLog.clear(); + final statesBefore = room.hostService.outgoingLog + .where((message) => message.type == SyncMessageType.state) + .length; + final stateBefore = room.lastHostState(); + + SyncMessage wireControl(ControlRequest request) { + return SyncMessage.fromJson(SyncMessage.control(request, peerId: 'forged-peer').toJson()); + } + + room.guestService.sendTo( + 'host', + wireControl( + ControlRequest(kind: ControlRequestKind.seek, positionMs: room.hostPlayer.state.duration.inMilliseconds + 1), + ), + ); + room.guestService.sendTo( + 'host', + wireControl(const ControlRequest(kind: ControlRequestKind.rate, rate: 4.000001)), + ); + async.flushMicrotasks(); + + expect(room.hostPlayer.commandLog, isEmpty); + expect(actions, isEmpty); + expect( + room.hostService.outgoingLog.where((message) => message.type == SyncMessageType.state), + hasLength(statesBefore), + ); + expect(room.lastHostState().seq, stateBefore.seq); + expect(room.lastHostState().anchorPositionMs, stateBefore.anchorPositionMs); + expect(room.lastHostState().rate, stateBefore.rate); + + room.guestService.sendTo( + 'host', + wireControl(const ControlRequest(kind: ControlRequestKind.seek, positionMs: 600000)), + ); + async.flushMicrotasks(); + room.guestService.sendTo('host', wireControl(const ControlRequest(kind: ControlRequestKind.rate, rate: 0.25))); + async.flushMicrotasks(); + + expect(room.hostPlayer.commandLog, ['seek:600000', 'rate:0.25']); + expect(actions, [('guest', PlaybackActionHint.seek), ('guest', PlaybackActionHint.rate)]); + final acceptedStates = room.hostService.outgoingLog + .where((message) => message.type == SyncMessageType.state) + .skip(statesBefore) + .map((message) => message.state!) + .toList(); + expect(acceptedStates, hasLength(2)); + expect(acceptedStates[0].anchorPositionMs, 600000); + expect(acceptedStates[0].actionHint, PlaybackActionHint.seek); + expect(acceptedStates[0].actorPeerId, 'guest'); + expect(acceptedStates[1].rate, 0.25); + expect(acceptedStates[1].actionHint, PlaybackActionHint.rate); + expect(acceptedStates[1].actorPeerId, 'guest'); + room.dispose(); + }); + }); + test('guest controller starts clock-sync pings automatically', () { fakeAsync((async) { final room = _Room(async); @@ -221,15 +289,14 @@ void main() { }); }); - test('v1 peers are flagged and never gate the start', () { + test('v2 peers are flagged and never gate the start', () { fakeAsync((async) { final needsUpdate = []; final room = _Room(async); room.host.onPeerNeedsUpdate = needsUpdate.add; - // A legacy client joins on its own connection: its join message has no - // version field (the relay stamps the sender id, so it must really - // connect as itself — peerId spoofing is rewritten). + // A sync-protocol-2 client joins on its own connection. The relay + // stamps the sender ID, so it must really connect as itself. final legacyService = room.hub.register('legacy'); legacyService.sendTo( 'host', @@ -239,6 +306,7 @@ void main() { peerId: 'legacy', displayName: 'Old App', isHost: false, + version: 2, ), ); async.flushMicrotasks(); @@ -253,6 +321,34 @@ void main() { }); }); + test('versionless legacy peers are flagged and never gate the start', () { + fakeAsync((async) { + final needsUpdate = []; + final room = _Room(async); + room.host.onPeerNeedsUpdate = needsUpdate.add; + + final versionlessService = room.hub.register('versionless'); + versionlessService.sendTo( + 'host', + SyncMessage( + type: SyncMessageType.join, + timestamp: room.nowMs(), + peerId: 'versionless', + displayName: 'Old App', + isHost: false, + ), + ); + async.flushMicrotasks(); + expect(needsUpdate, ['versionless']); + + room.hostStartsMedia(); + room.guestJoinsMedia(); + room.bothBecomeReady(); + expect(room.lastHostState().phase, PlaybackPhase.playing); + room.dispose(); + }); + }); + test('guest reconnect re-requests state and the host answers directly', () { fakeAsync((async) { final room = _Room(async); @@ -274,6 +370,46 @@ void main() { }); }); + test('only relay-stamped state from the declared host reaches guest reconciliation', () { + fakeAsync((async) { + final room = _Room(async); + final mediaDispatches = []; + room.guest.onMediaStateReceived = (ratingKey, serverId, title) => mediaDispatches.add(ratingKey); + const state = PlaybackState( + seq: 10, + ratingKey: 'relay-authority', + serverId: 'srv', + mediaTitle: 'Authorized', + phase: PlaybackPhase.loading, + anchorPositionMs: 0, + anchorHostTimeMs: _epochMs, + rate: 1, + controlMode: ControlMode.hostOnly, + ); + final unprivileged = room.hub.register('unprivileged'); + + // The fake relay overwrites the payload claim with the connection's + // routing ID, just like the production relay envelope parser. + unprivileged.broadcast(SyncMessage.state(state, peerId: 'host')); + async.flushMicrotasks(); + expect(mediaDispatches, isEmpty); + + room.hostService.broadcast(SyncMessage.state(state, peerId: 'host')); + async.flushMicrotasks(); + expect(mediaDispatches, ['relay-authority']); + room.dispose(); + }); + }); + + test('fake relay rejects duplicate routing IDs instead of replacing authority', () async { + final hub = FakeRelayHub(); + hub.register('reserved'); + + expect(() => hub.register('reserved'), throwsStateError); + + await hub.dispose(); + }); + group('hostExitedPlayer routing', () { test('rides the ordered queue: never overtakes states sent before it', () { fakeAsync((async) { diff --git a/test/watch_together/watch_together_overlay_test.dart b/test/watch_together/watch_together_overlay_test.dart index 266264ba..ea79c292 100644 --- a/test/watch_together/watch_together_overlay_test.dart +++ b/test/watch_together/watch_together_overlay_test.dart @@ -55,6 +55,20 @@ void main() { expect(harness.onLeaveSessionCalls, 0); }); } + + testWidgets('best-effort leave failure is handled by the overlay', (tester) async { + final harness = _OverlayHarness(isHost: false, leaveError: StateError('release failed')); + addTearDown(harness.dispose); + await tester.pumpWidget(harness.build()); + + await _openLeaveConfirmation(tester, harness); + await tester.tap(find.text(t.watchTogether.leave)); + await tester.pumpAndSettle(); + + expect(harness.provider.leaveCalls, 1); + expect(harness.onLeaveSessionCalls, 1); + expect(tester.takeException(), isNull); + }); } Future _openLeaveConfirmation(WidgetTester tester, _OverlayHarness harness) async { @@ -69,7 +83,8 @@ Future _openLeaveConfirmation(WidgetTester tester, _OverlayHarness harness } class _OverlayHarness { - _OverlayHarness({required bool isHost}) : provider = _FakeWatchTogetherProvider(isHostValue: isHost); + _OverlayHarness({required bool isHost, Object? leaveError}) + : provider = _FakeWatchTogetherProvider(isHostValue: isHost, leaveError: leaveError); static const indicatorKey = Key('watch-together-session-indicator'); @@ -101,9 +116,10 @@ class _OverlayHarness { } class _FakeWatchTogetherProvider extends WatchTogetherProvider { - _FakeWatchTogetherProvider({required this.isHostValue}); + _FakeWatchTogetherProvider({required this.isHostValue, this.leaveError}); final bool isHostValue; + final Object? leaveError; var leaveCalls = 0; var _isDisposing = false; @@ -127,6 +143,8 @@ class _FakeWatchTogetherProvider extends WatchTogetherProvider { @override Future leaveSession() async { if (!_isDisposing) leaveCalls++; + final error = leaveError; + if (error != null) throw error; } @override diff --git a/test/watch_together/watch_together_peer_service_test.dart b/test/watch_together/watch_together_peer_service_test.dart index 9812fe45..b88a5b5d 100644 --- a/test/watch_together/watch_together_peer_service_test.dart +++ b/test/watch_together/watch_together_peer_service_test.dart @@ -3,10 +3,14 @@ import 'dart:convert'; import 'dart:io'; import 'package:flutter_test/flutter_test.dart'; +import 'package:stream_channel/stream_channel.dart'; +import 'package:web_socket_channel/web_socket_channel.dart'; import 'package:plezy/watch_together/services/watch_together_peer_service.dart'; +import 'package:plezy/watch_together/services/watch_together_relay_endpoint.dart'; import 'package:plezy/watch_together/models/sync_message.dart'; typedef _MessageHandler = FutureOr Function(int connection, WebSocket socket, Map message); +const _relayHostId = 'relay-host-7'; Future _withShortenedTimer({ required Duration original, @@ -23,6 +27,70 @@ Future _withShortenedTimer({ ); } +Future _withSetupTimersShortened(Future Function() body) { + return _withShortenedTimer( + original: const Duration(seconds: 10), + replacement: const Duration(milliseconds: 500), + body: () => _withShortenedTimer( + original: const Duration(milliseconds: 250), + replacement: const Duration(milliseconds: 1), + body: () => _withShortenedTimer( + original: const Duration(milliseconds: 500), + replacement: const Duration(milliseconds: 1), + body: body, + ), + ), + ); +} + +class _TrackingWebSocketSink implements WebSocketSink { + final Completer _done = Completer(); + bool closed = false; + + @override + void add(dynamic data) {} + + @override + void addError(Object error, [StackTrace? stackTrace]) {} + + @override + Future addStream(Stream stream) => stream.drain(); + + @override + Future close([int? closeCode, String? closeReason]) { + closed = true; + if (!_done.isCompleted) _done.complete(); + return _done.future; + } + + @override + Future get done => _done.future; +} + +class _PendingWebSocketChannel extends StreamChannelMixin implements WebSocketChannel { + final Completer _readyCompleter = Completer(); + + @override + final _TrackingWebSocketSink sink = _TrackingWebSocketSink(); + + @override + Stream get stream => const Stream.empty(); + + @override + Future get ready => _readyCompleter.future; + + @override + String? get protocol => null; + + @override + int? get closeCode => null; + + @override + String? get closeReason => null; + + Future closeForTesting() => sink.close(); +} + class _RelayServer { _RelayServer._(this._server, this._handler); @@ -64,8 +132,16 @@ void main() { final services = []; final relays = <_RelayServer>[]; - WatchTogetherPeerService serviceFor(_RelayServer relay) { - final service = WatchTogetherPeerService(customBaseUrl: relay.baseUrl); + WatchTogetherPeerService serviceFor( + _RelayServer relay, { + Future Function()? debugReconnectSetupSucceededBarrier, + WebSocketChannel Function(Uri uri)? debugChannelFactory, + }) { + final service = WatchTogetherPeerService( + endpoint: WatchTogetherRelayEndpoint.resolve(relay.baseUrl), + debugReconnectSetupSucceededBarrier: debugReconnectSetupSucceededBarrier, + debugChannelFactory: debugChannelFactory, + ); services.add(service); return service; } @@ -88,6 +164,59 @@ void main() { relays.clear(); }); + group('WatchTogetherRelayEndpoint', () { + test('resolves defaults and canonical endpoint paths', () { + for (final value in [null, '', ' ']) { + final endpoint = WatchTogetherRelayEndpoint.resolve(value); + expect(endpoint.canonicalBaseUrl, 'https://ice.plezy.app'); + expect(endpoint.healthUri.toString(), 'https://ice.plezy.app/health'); + expect(endpoint.webSocketUri.toString(), 'wss://ice.plezy.app/relay'); + } + + final endpoint = WatchTogetherRelayEndpoint.resolve(' HTTP://Example.COM:8080/old/../prefix/// '); + expect(endpoint.canonicalBaseUrl, 'http://example.com:8080/prefix'); + expect(endpoint.healthUri.toString(), 'http://example.com:8080/prefix/health'); + expect(endpoint.webSocketUri.toString(), 'ws://example.com:8080/prefix/relay'); + }); + + test('default ports share canonical identity while non-default ports remain distinct', () { + final http = WatchTogetherRelayEndpoint.resolve('http://relay.example.test:80/base/'); + final implicitHttp = WatchTogetherRelayEndpoint.resolve('http://relay.example.test/base'); + final https = WatchTogetherRelayEndpoint.resolve('https://relay.example.test:443/base/'); + final implicitHttps = WatchTogetherRelayEndpoint.resolve('https://relay.example.test/base'); + final nonDefault = WatchTogetherRelayEndpoint.resolve('https://relay.example.test:8443/base'); + + expect(http.canonicalBaseUrl, 'http://relay.example.test/base'); + expect(https.canonicalBaseUrl, 'https://relay.example.test/base'); + expect(http, implicitHttp); + expect(http.hashCode, implicitHttp.hashCode); + expect(https, implicitHttps); + expect(https.hashCode, implicitHttps.hashCode); + expect(nonDefault.canonicalBaseUrl, 'https://relay.example.test:8443/base'); + expect(nonDefault, isNot(implicitHttps)); + expect(http, isNot(https)); + }); + + test('accepts supported host forms and rejects unusable bases', () { + final ipv6 = WatchTogetherRelayEndpoint.tryParseCustom('https://[2001:db8::1]:8443/base/'); + expect(ipv6?.canonicalBaseUrl, 'https://[2001:db8::1]:8443/base'); + + for (final value in [ + 'relay.example.test', + '/relative', + 'ftp://relay.example.test', + 'ws://relay.example.test', + 'https://', + 'https://opaque@relay.example.test', + 'https://relay.example.test/path?mode=test', + 'https://relay.example.test/path#fragment', + 'https://relay.example.test:99999', + ]) { + expect(WatchTogetherRelayEndpoint.tryParseCustom(value), isNull, reason: value); + } + }); + }); + test('invalid relay identifiers fail before network access', () async { final service = WatchTogetherPeerService(); services.add(service); @@ -100,31 +229,47 @@ void main() { ); }); - test('host connects, listens, and announces create with the existing wire format', () async { + test('host stores relay authority and uses a random routing ID', () async { late final _RelayServer relay; relay = await relayWith((_, socket, message) { if (message['type'] == 'create') { - relay.send(socket, {'type': 'created', 'sessionId': message['sessionId']}); + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); } }); final service = serviceFor(relay); expect(await service.createSession(sessionId: 'abc12'), 'ABC12'); expect(service.isHost, isTrue); - expect(service.myPeerId, 'wt-ABC12'); - expect(relay.messages.single, [ - {'type': 'create', 'sessionId': 'ABC12', 'peerId': 'wt-ABC12'}, - ]); + expect(service.myPeerId, isNot('wt-ABC12')); + expect(service.myPeerId, matches(RegExp(r'^[0-9a-f-]{36}$'))); + expect(service.hostPeerId, service.myPeerId); + final create = relay.messages.single.single; + expect(create, { + 'type': 'create', + 'sessionId': 'ABC12', + 'peerId': service.myPeerId, + 'reconnectToken': matches(RegExp(r'^[A-Za-z0-9_-]{43}$')), + 'protocolVersion': 2, + }); }); - test('guest connects, listens, and announces join with the existing wire format', () async { + test('guest stores the relay-declared host identity', () async { late final _RelayServer relay; relay = await relayWith((_, socket, message) { if (message['type'] == 'join') { relay.send(socket, { 'type': 'joined', 'sessionId': message['sessionId'], - 'peers': ['wt-ROOM1'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], }); } }); @@ -136,22 +281,169 @@ void main() { await service.joinSession('room1'); expect(service.isHost, isFalse); - expect(service.connectedPeers, ['wt-ROOM1']); - expect(connectedPeers, ['wt-ROOM1']); - expect(relay.messages.single, [ - {'type': 'join', 'sessionId': 'ROOM1', 'peerId': service.myPeerId}, + expect(service.hostPeerId, _relayHostId); + expect(service.connectedPeers, [_relayHostId]); + expect(connectedPeers, [_relayHostId]); + expect(relay.messages.single.single, { + 'type': 'join', + 'sessionId': 'ROOM1', + 'peerId': service.myPeerId, + 'reconnectToken': matches(RegExp(r'^[A-Za-z0-9_-]{43}$')), + 'protocolVersion': 2, + }); + }); + + test('guest reconnect sends its retained capability', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } + }); + final service = serviceFor(relay); + final reconnected = Completer(); + service.onReconnected = reconnected.complete; + + await _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () => service.joinSession('guest1'), + ); + final guestPeerId = service.myPeerId; + await relay.sockets.single.close(); + await reconnected.future.timeout(const Duration(seconds: 6)); + + final initialToken = relay.messages[0].single['reconnectToken']; + expect(relay.messages[1], [ + { + 'type': 'join', + 'sessionId': 'GUEST1', + 'peerId': guestPeerId, + 'reconnectToken': initialToken, + 'protocolVersion': 2, + }, + ]); + expect(service.hostPeerId, _relayHostId); + }); + + test('guest reconnect releases an admitted identity when the relay host changed', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': connection == 0 ? _relayHostId : 'replacement-host', + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [connection == 0 ? _relayHostId : 'replacement-host'], + }); + } else if (message['type'] == 'leave') { + relay.send(socket, { + 'type': 'left', + 'sessionId': message['sessionId'], + 'peerId': message['peerId'], + 'protocolVersion': 2, + }); + } + }); + final service = serviceFor(relay); + var reconnectCallbacks = 0; + service.onReconnected = () => reconnectCallbacks++; + final identityError = service.onError.firstWhere( + (error) => error.type == PeerErrorType.serverError && error.message.contains('invalid joined response'), + ); + + await _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () => service.joinSession('guest2'), + ); + await relay.sockets.single.close(); + await identityError.timeout(const Duration(seconds: 1)); + await Future.delayed(const Duration(milliseconds: 30)); + + expect(service.hostPeerId, _relayHostId); + expect(reconnectCallbacks, 0); + expect(relay.sockets, hasLength(2)); + final reconnect = relay.messages[1].first; + expect(relay.messages[1], [ + reconnect, + { + 'type': 'leave', + 'sessionId': 'GUEST2', + 'peerId': reconnect['peerId'], + 'reconnectToken': reconnect['reconnectToken'], + 'protocolVersion': 2, + }, ]); }); - test('host reconnect joins first and re-creates a missing room on the same socket', () async { + test('guest reconnect closes after a rejected admission leave ACK is lost', () async { + final leaveSeen = Completer(); + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': connection == 0 ? _relayHostId : 'replacement-host', + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [connection == 0 ? _relayHostId : 'replacement-host'], + }); + } else if (message['type'] == 'leave' && !leaveSeen.isCompleted) { + leaveSeen.complete(); + } + }); + final service = serviceFor(relay); + + await _withShortenedTimer( + original: const Duration(seconds: 10), + replacement: const Duration(milliseconds: 500), + body: () => _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () async { + await service.joinSession('guest3'); + await relay.sockets.single.close(); + await leaveSeen.future.timeout(const Duration(seconds: 1)); + await relay.sockets[1].done.timeout(const Duration(seconds: 1)); + }, + ), + ); + + expect(relay.messages[1].map((message) => message['type']), ['join', 'leave']); + }); + + test('host reconnect proves ownership and re-creates with the retained authority', () async { late final _RelayServer relay; relay = await relayWith((connection, socket, message) { if (connection == 0 && message['type'] == 'create') { - relay.send(socket, {'type': 'created', 'sessionId': message['sessionId']}); + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); } else if (connection == 1 && message['type'] == 'join') { relay.send(socket, {'type': 'error', 'code': 'room_not_found', 'message': 'Room not found'}); } else if (connection == 1 && message['type'] == 'create') { - relay.send(socket, {'type': 'created', 'sessionId': message['sessionId']}); + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); } }); final service = serviceFor(relay); @@ -167,18 +459,71 @@ void main() { replacement: const Duration(milliseconds: 10), body: () => service.createSession(sessionId: 'room2'), ); + final hostPeerId = service.myPeerId; await relay.sockets.single.close(); await reconnected.future.timeout(const Duration(seconds: 6)); expect(reconnectCallbacks, 1); expect(relay.sockets, hasLength(2)); - expect(relay.messages[0], [ - {'type': 'create', 'sessionId': 'ROOM2', 'peerId': 'wt-ROOM2'}, - ]); + final initialCreate = relay.messages[0].single; + final reconnectToken = initialCreate['reconnectToken']; + expect(initialCreate, { + 'type': 'create', + 'sessionId': 'ROOM2', + 'peerId': hostPeerId, + 'reconnectToken': matches(RegExp(r'^[A-Za-z0-9_-]{43}$')), + 'protocolVersion': 2, + }); expect(relay.messages[1], [ - {'type': 'join', 'sessionId': 'ROOM2', 'peerId': 'wt-ROOM2'}, - {'type': 'create', 'sessionId': 'ROOM2', 'peerId': 'wt-ROOM2'}, + { + 'type': 'join', + 'sessionId': 'ROOM2', + 'peerId': hostPeerId, + 'reconnectToken': reconnectToken, + 'protocolVersion': 2, + }, + { + 'type': 'create', + 'sessionId': 'ROOM2', + 'peerId': hostPeerId, + 'reconnectToken': reconnectToken, + 'protocolVersion': 2, + }, ]); + expect(service.hostPeerId, hostPeerId); + }); + + test('initial create retry reuses its pre-minted identity and capability', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) async { + if (message['type'] != 'create') return; + if (connection == 0) { + await socket.close(); + return; + } + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); + }); + final service = serviceFor(relay); + + await _withShortenedTimer( + original: const Duration(milliseconds: 250), + replacement: const Duration(milliseconds: 1), + body: () => service.createSession(sessionId: 'retry1'), + ); + + expect(relay.messages, hasLength(2)); + final first = relay.messages[0].single; + final retry = relay.messages[1].single; + expect(first['reconnectToken'], matches(RegExp(r'^[A-Za-z0-9_-]{43}$'))); + expect(retry, first); + expect(service.myPeerId, first['peerId']); + expect(service.hostPeerId, first['peerId']); }); test('setup preserves typed timeout and relay errors', () async { @@ -214,11 +559,548 @@ void main() { ); }); + test('relay protocol mismatch remains a typed setup error', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + relay.send(socket, { + 'type': 'error', + 'code': 'protocol_mismatch', + 'message': 'Relay protocol version 2 is required', + }); + }); + final service = serviceFor(relay); + + await expectLater( + service.joinSession('proto1'), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.serverError) + .having((error) => error.serverCode, 'serverCode', 'protocol_mismatch') + .having((error) => error.message, 'message', contains('Relay protocol version 2 is required')), + ), + ); + }); + + test('setup rejects a relay that substitutes the pre-minted capability', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + const tokenA = 'AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA'; + const tokenB = 'BBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBB'; + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'] == tokenA ? tokenB : tokenA, + 'protocolVersion': 2, + }); + }); + final service = serviceFor(relay); + + await expectLater( + service.createSession(sessionId: 'token1'), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.serverError) + .having((error) => error.message, 'message', 'Relay returned an invalid created response'), + ), + ); + expect(service.hostPeerId, isNull); + }); + + test('missing or malformed authority fields fail setup immediately', () async { + late final _RelayServer oldRelay; + oldRelay = await relayWith((_, socket, message) { + oldRelay.send(socket, {'type': 'created', 'sessionId': message['sessionId'], 'protocolVersion': 2}); + }); + final oldRelayService = serviceFor(oldRelay); + + await expectLater( + oldRelayService.createSession(sessionId: 'old01'), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.serverError) + .having((error) => error.message, 'message', 'Relay returned an invalid created response'), + ), + ); + expect(oldRelayService.hostPeerId, isNull); + + late final _RelayServer malformedRelay; + malformedRelay = await relayWith((_, socket, message) { + malformedRelay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': 'not-a-capability', + 'protocolVersion': 2, + }); + }); + final malformedService = serviceFor(malformedRelay); + + await expectLater( + malformedService.joinSession('bad01'), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.serverError) + .having((error) => error.message, 'message', 'Relay returned an invalid joined response'), + ), + ); + expect(malformedService.hostPeerId, isNull); + }); + + test('exhausted create retries end a possibly committed room before clearing credentials', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (connection >= 3 && message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': const [], + }); + } else if (message['type'] == 'endSession') { + relay.send(socket, {'type': 'ended', 'sessionId': message['sessionId'], 'protocolVersion': 2}); + } + }); + final service = serviceFor(relay); + + await expectLater( + _withSetupTimersShortened(() => service.createSession(sessionId: 'lostc')), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.timeout) + .having((error) => error.message, 'message', 'Timed out creating session'), + ), + ); + + expect(relay.messages.map((messages) => messages.map((message) => message['type']).toList()).toList(), [ + ['create'], + ['create'], + ['create'], + ['join', 'endSession'], + ]); + final announcements = relay.messages.map((messages) => messages.first).toList(); + expect(announcements.map((message) => message['peerId']).toSet(), hasLength(1)); + expect(announcements.map((message) => message['reconnectToken']).toSet(), hasLength(1)); + expect(service.sessionId, isNull); + expect(service.myPeerId, isNull); + }); + + test('exhausted join retries leave a possibly committed guest reservation before clearing credentials', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (connection >= 3 && message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } else if (message['type'] == 'leave') { + relay.send(socket, { + 'type': 'left', + 'sessionId': message['sessionId'], + 'peerId': message['peerId'], + 'protocolVersion': 2, + }); + } + }); + final service = serviceFor(relay); + + await expectLater( + _withSetupTimersShortened(() => service.joinSession('lostj')), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.timeout) + .having((error) => error.message, 'message', isNotEmpty), + ), + ); + + expect(relay.messages.map((messages) => messages.map((message) => message['type']).toList()).toList(), [ + ['join'], + ['join'], + ['join'], + ['join', 'leave'], + ]); + final announcements = relay.messages.map((messages) => messages.first).toList(); + expect(announcements.map((message) => message['peerId']).toSet(), hasLength(1)); + expect(announcements.map((message) => message['reconnectToken']).toSet(), hasLength(1)); + expect(service.sessionId, isNull); + expect(service.myPeerId, isNull); + }); + + test('guest initial setup treats ended as terminal rather than joined', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, {'type': 'ended', 'sessionId': message['sessionId'], 'protocolVersion': 2}); + } + }); + final service = serviceFor(relay); + final ended = service.onSessionEnded.first; + + await expectLater( + service.joinSession('ended2'), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.invalidSession) + .having((error) => error.message, 'message', 'Watch Together session ended'), + ), + ); + await ended.timeout(const Duration(seconds: 1)); + + expect(service.sessionId, isNull); + expect(service.isConnected, isFalse); + expect(relay.sockets, hasLength(1)); + }); + + test('guest reconnect treats ended as terminal and cancels further reconnects', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (message['type'] != 'join') return; + if (connection == 0) { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } else { + relay.send(socket, {'type': 'ended', 'sessionId': message['sessionId'], 'protocolVersion': 2}); + } + }); + final service = serviceFor(relay); + var reconnectCallbacks = 0; + service.onReconnected = () => reconnectCallbacks++; + final ended = service.onSessionEnded.first; + + await _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () async { + await service.joinSession('ended3'); + await relay.sockets.single.close(); + await ended.timeout(const Duration(seconds: 1)); + await Future.delayed(const Duration(milliseconds: 50)); + }, + ); + + expect(reconnectCallbacks, 0); + expect(service.isConnected, isFalse); + expect(relay.sockets, hasLength(2)); + }); + + test('guest reconnect converges to session ended after room-not-found retries are exhausted', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (message['type'] != 'join') return; + if (connection == 0) { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } else { + relay.send(socket, {'type': 'error', 'code': 'room_not_found', 'message': 'Room not found'}); + } + }); + final service = serviceFor(relay); + var reconnectCallbacks = 0; + var sessionEndedEvents = 0; + final ended = Completer(); + service.onReconnected = () => reconnectCallbacks++; + service.onSessionEnded.listen((_) { + sessionEndedEvents++; + if (!ended.isCompleted) ended.complete(); + }); + + await _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () => _withShortenedTimer( + original: const Duration(seconds: 4), + replacement: const Duration(milliseconds: 10), + body: () => _withShortenedTimer( + original: const Duration(seconds: 6), + replacement: const Duration(milliseconds: 10), + body: () async { + await service.joinSession('ended4'); + await relay.sockets.single.close(); + await ended.future.timeout(const Duration(seconds: 1)); + await Future.delayed(const Duration(milliseconds: 50)); + }, + ), + ), + ); + + expect(reconnectCallbacks, 0); + expect(sessionEndedEvents, 1); + expect(service.isConnected, isFalse); + expect(relay.sockets, hasLength(4)); + expect( + relay.messages.skip(1).map((messages) => messages.map((message) => message['type']).toList()), + everyElement(['join']), + ); + }); + + test('initial invalid-room join remains a room-not-found error', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, {'type': 'error', 'code': 'room_not_found', 'message': 'Room not found'}); + } + }); + final service = serviceFor(relay); + var sessionEndedEvents = 0; + service.onSessionEnded.listen((_) => sessionEndedEvents++); + + await expectLater( + service.joinSession('miss01'), + throwsA( + isA() + .having((error) => error.type, 'type', PeerErrorType.serverError) + .having((error) => error.serverCode, 'serverCode', 'room_not_found'), + ), + ); + await Future.delayed(Duration.zero); + + expect(sessionEndedEvents, 0); + expect(relay.messages, hasLength(1)); + }); + + test('release timeout covers a pending WebSocket handshake and closes every channel', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } + }); + var channelCalls = 0; + final pendingChannels = <_PendingWebSocketChannel>[]; + addTearDown(() async { + for (final channel in pendingChannels) { + await channel.closeForTesting(); + } + }); + final service = serviceFor( + relay, + debugChannelFactory: (uri) { + if (channelCalls++ == 0) return WebSocketChannel.connect(uri); + final channel = _PendingWebSocketChannel(); + pendingChannels.add(channel); + return channel; + }, + ); + + await service.joinSession('hang01'); + final disconnected = service.onConnectionStateChanged.firstWhere((connected) => !connected); + await relay.sockets.single.close(); + await disconnected.timeout(const Duration(seconds: 1)); + + await expectLater(_withSetupTimersShortened(service.releaseSession), throwsA(isA())); + + expect(pendingChannels, hasLength(3)); + expect(pendingChannels.every((channel) => channel.sink.closed), isTrue); + expect(relay.sockets, hasLength(1)); + }); + + test('guest release accepts peer_id_unavailable after a processed leave loses its ACK', () async { + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (connection == 0 && message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } else if (connection == 1 && message['type'] == 'join') { + relay.send(socket, { + 'type': 'error', + 'code': 'peer_id_unavailable', + 'message': 'Peer identity is no longer available', + }); + } + }); + final service = serviceFor(relay); + + await service.joinSession('lostl'); + await _withSetupTimersShortened(service.releaseSession); + + expect(relay.messages, hasLength(2)); + expect(relay.messages[0].map((message) => message['type']), ['join', 'leave']); + expect(relay.messages[1].map((message) => message['type']), ['join']); + }); + + test('guest leave waits for a protocol-2 left acknowledgement', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } else if (message['type'] == 'leave') { + relay.send(socket, { + 'type': 'left', + 'sessionId': message['sessionId'], + 'peerId': message['peerId'], + 'protocolVersion': 2, + }); + } + }); + final service = serviceFor(relay); + + await service.joinSession('leave1'); + final join = relay.messages.single.single; + await service.releaseSession(); + + expect(relay.messages.single, [ + join, + { + 'type': 'leave', + 'sessionId': 'LEAVE1', + 'peerId': service.myPeerId, + 'reconnectToken': join['reconnectToken'], + 'protocolVersion': 2, + }, + ]); + }); + + test('sequential guest release accepts not-in-room as idempotent success', () async { + var leaveRequests = 0; + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } else if (message['type'] == 'leave') { + leaveRequests++; + if (leaveRequests == 1) { + relay.send(socket, { + 'type': 'left', + 'sessionId': message['sessionId'], + 'peerId': message['peerId'], + 'protocolVersion': 2, + }); + } else { + relay.send(socket, {'type': 'error', 'code': 'not_in_room', 'message': 'Peer is not in the room'}); + } + } + }); + final service = serviceFor(relay); + + await service.joinSession('leave2'); + await service.releaseSession(); + await service.releaseSession(); + + expect(leaveRequests, 2); + expect(relay.messages.single.map((message) => message['type']), ['join', 'leave', 'leave']); + }); + + test('host end waits for a protocol-2 ended acknowledgement', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'create') { + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); + } else if (message['type'] == 'endSession') { + relay.send(socket, {'type': 'ended', 'sessionId': message['sessionId'], 'protocolVersion': 2}); + } + }); + final service = serviceFor(relay); + + await service.createSession(sessionId: 'end01'); + final create = relay.messages.single.single; + await service.releaseSession(); + + expect(relay.messages.single, [ + create, + { + 'type': 'endSession', + 'sessionId': 'END01', + 'peerId': service.myPeerId, + 'reconnectToken': create['reconnectToken'], + 'protocolVersion': 2, + }, + ]); + }); + + test('guest receives ended before close and does not enter reconnect', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _relayHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': [_relayHostId], + }); + } + }); + final service = serviceFor(relay); + + await _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () async { + await service.joinSession('ended1'); + final ended = service.onSessionEnded.first; + relay.send(relay.sockets.single, {'type': 'ended', 'sessionId': 'ENDED1', 'protocolVersion': 2}); + await ended.timeout(const Duration(seconds: 1)); + await relay.sockets.single.close(); + await Future.delayed(const Duration(milliseconds: 100)); + }, + ); + + expect(relay.sockets, hasLength(1)); + }); + test('one setup installs one listener and sends one announcement', () async { late final _RelayServer relay; relay = await relayWith((_, socket, message) { if (message['type'] == 'create') { - relay.send(socket, {'type': 'created', 'sessionId': message['sessionId']}); + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); } }); final service = serviceFor(relay); @@ -239,6 +1121,36 @@ void main() { expect(relay.messages.single.where((message) => message['type'] == 'create'), hasLength(1)); }); + test('explicit disconnect clears relay authority before a new setup', () async { + late final _RelayServer relay; + relay = await relayWith((_, socket, message) { + if (message['type'] == 'create') { + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); + } + }); + final service = serviceFor(relay); + + await service.createSession(sessionId: 'clear1'); + expect(service.hostPeerId, isNotNull); + await service.disconnect(); + expect(service.hostPeerId, isNull); + + await service.createSession(sessionId: 'clear2'); + expect(relay.messages[1].single, { + 'type': 'create', + 'sessionId': 'CLEAR2', + 'peerId': service.myPeerId, + 'reconnectToken': matches(RegExp(r'^[A-Za-z0-9_-]{43}$')), + 'protocolVersion': 2, + }); + }); + test('disconnect cancels an in-flight room announcement without timeout delay', () async { final announcementSeen = Completer(); final relay = await relayWith((_, _, message) { @@ -256,4 +1168,53 @@ void main() { expect(service.sessionId, isNull); expect(service.connectedPeers, isEmpty); }); + + test('dispose invalidates a reconnect after relay setup succeeds', () async { + final reconnectSetupSucceeded = Completer(); + final releaseReconnectPublication = Completer(); + late final _RelayServer relay; + relay = await relayWith((connection, socket, message) { + if (connection == 0 && message['type'] == 'create') { + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); + } else if (connection == 1 && message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + 'peers': const [], + }); + } + }); + final service = serviceFor( + relay, + debugReconnectSetupSucceededBarrier: () { + reconnectSetupSucceeded.complete(); + return releaseReconnectPublication.future; + }, + ); + var reconnectCallbacks = 0; + service.onReconnected = () => reconnectCallbacks++; + + await _withShortenedTimer( + original: const Duration(seconds: 2), + replacement: const Duration(milliseconds: 10), + body: () => service.createSession(sessionId: 'epoch1'), + ); + await relay.sockets.single.close(); + await reconnectSetupSucceeded.future.timeout(const Duration(seconds: 1)); + + service.dispose(); + releaseReconnectPublication.complete(); + await Future.delayed(const Duration(milliseconds: 50)); + + expect(reconnectCallbacks, 0); + }); } diff --git a/test/watch_together/watch_together_provider_test.dart b/test/watch_together/watch_together_provider_test.dart index 8c285b95..d28d368e 100644 --- a/test/watch_together/watch_together_provider_test.dart +++ b/test/watch_together/watch_together_provider_test.dart @@ -1,13 +1,256 @@ import 'dart:async'; +import 'dart:convert'; +import 'dart:io'; +import 'package:fake_async/fake_async.dart'; import 'package:flutter_test/flutter_test.dart'; import 'package:plezy/media/ids.dart'; +import 'package:plezy/watch_together/services/watch_together_relay_endpoint.dart'; +import 'package:plezy/watch_together/models/sync_message.dart'; import 'package:plezy/watch_together/models/watch_session.dart'; import 'package:plezy/watch_together/providers/watch_together_provider.dart'; +import 'package:plezy/watch_together/services/watch_together_peer_service.dart'; +import 'package:plezy/watch_together/services/relay_protocol.g.dart'; import '../test_helpers/prefs.dart'; +class _FakeWatchTogetherPeerService extends WatchTogetherPeerService { + _FakeWatchTogetherPeerService( + this.sequence, { + required bool hostInitiallyConnected, + required this.rejectDisconnectedTargets, + this.releaseError, + this.releaseBarrier, + }) : _hostConnected = hostInitiallyConnected; + + final int sequence; + final bool rejectDisconnectedTargets; + final Object? releaseError; + final Future? releaseBarrier; + final _peerConnectedController = StreamController.broadcast(); + final _peerDisconnectedController = StreamController.broadcast(); + final _messageController = StreamController.broadcast(); + final _errorController = StreamController.broadcast(); + final _sessionEndedController = StreamController.broadcast(); + final List broadcasts = []; + final List<(String, SyncMessage)> directMessages = []; + + String? _sessionId; + String? _myPeerId; + String? _hostPeerId; + bool _isHost = false; + bool _disposed = false; + bool _didDisconnect = false; + int releaseCalls = 0; + bool _hostConnected; + + @override + Stream get onPeerConnected => _peerConnectedController.stream; + + @override + Stream get onPeerDisconnected => _peerDisconnectedController.stream; + + @override + Stream get onMessageReceived => _messageController.stream; + + @override + Stream get onError => _errorController.stream; + + @override + Stream get onSessionEnded => _sessionEndedController.stream; + + @override + String? get sessionId => _sessionId; + + @override + String? get myPeerId => _myPeerId; + + @override + String? get hostPeerId => _hostPeerId; + + @override + bool get isHost => _isHost; + + @override + List get connectedPeers => !_isHost && _sessionId != null && _hostConnected ? ['wt-$_sessionId'] : const []; + + @override + Future createSession({String? sessionId}) { + _sessionId = (sessionId ?? 'ROOM$sequence').toUpperCase(); + _myPeerId = 'wt-$_sessionId'; + _hostPeerId = _myPeerId; + _isHost = true; + return Future.value(_sessionId); + } + + @override + Future joinSession(String sessionId) { + _sessionId = sessionId.toUpperCase(); + _myPeerId = 'guest-$sequence'; + _hostPeerId = 'wt-$_sessionId'; + _isHost = false; + return Future.value(); + } + + @override + void broadcast(SyncMessage message) { + broadcasts.add(message); + } + + @override + void sendTo(String peerId, SyncMessage message) { + if (rejectDisconnectedTargets && !connectedPeers.contains(peerId)) { + _errorController.add( + const PeerError(type: PeerErrorType.serverError, message: 'Peer is not in the room', serverCode: 'not_in_room'), + ); + return; + } + directMessages.add((peerId, message)); + } + + void emitError(PeerError error) => _errorController.add(error); + void emitPeerConnected(String peerId) { + if (peerId == _hostPeerId) _hostConnected = true; + _peerConnectedController.add(peerId); + } + + void emitPeerDisconnected(String peerId) { + if (peerId == _hostPeerId) _hostConnected = false; + _peerDisconnectedController.add(peerId); + } + + void emitSessionEnded() => _sessionEndedController.add(null); + void emitMessage(SyncMessage message) => _messageController.add(message); + + bool get hasRelayListeners => + _peerConnectedController.hasListener || + _peerDisconnectedController.hasListener || + _messageController.hasListener || + _errorController.hasListener || + _sessionEndedController.hasListener; + + bool get isDisposed => _disposed; + bool get didDisconnect => _didDisconnect; + + @override + Future releaseSession() async { + releaseCalls++; + final barrier = releaseBarrier; + if (barrier != null) await barrier; + final error = releaseError; + if (error != null) throw error; + } + + void reconnect() => onReconnected?.call(); + @override + Future disconnect() { + _didDisconnect = true; + _sessionId = null; + _myPeerId = null; + _hostPeerId = null; + _isHost = false; + _hostConnected = false; + return Future.value(); + } + + @override + void dispose() { + if (_disposed) return; + _disposed = true; + _peerConnectedController.close(); + _peerDisconnectedController.close(); + _messageController.close(); + _errorController.close(); + _sessionEndedController.close(); + super.dispose(); + } +} + +class _FakePeerServiceFactory { + _FakePeerServiceFactory({ + this.hostInitiallyConnected = true, + this.rejectDisconnectedTargets = false, + this.releaseError, + this.releaseBarrier, + }); + + final bool hostInitiallyConnected; + final bool rejectDisconnectedTargets; + final Object? releaseError; + final Future? releaseBarrier; + final List<_FakeWatchTogetherPeerService> services = []; + WatchTogetherPeerService call({WatchTogetherRelayEndpoint? endpoint}) { + final service = _FakeWatchTogetherPeerService( + services.length + 1, + hostInitiallyConnected: hostInitiallyConnected, + rejectDisconnectedTargets: rejectDisconnectedTargets, + releaseError: releaseError, + releaseBarrier: releaseBarrier, + ); + services.add(service); + return service; + } +} + +PeerError _transportError([String message = 'WebSocket error: connection reset']) { + return PeerError(type: PeerErrorType.serverError, message: message, originalError: StateError('connection reset')); +} + +Future _flushProviderEvents() async { + await Future.delayed(Duration.zero); + await Future.delayed(Duration.zero); +} + +const _providerHostId = 'relay-authoritative-host'; + +typedef _ProviderRelayHandler = void Function(WebSocket socket, Map message); + +class _ProviderRelay { + _ProviderRelay._(this._server); + + final HttpServer _server; + final List _sockets = []; + final List> messages = []; + + String get baseUrl => 'http://${_server.address.address}:${_server.port}'; + + static Future<_ProviderRelay> start(_ProviderRelayHandler handler) async { + final server = await HttpServer.bind(InternetAddress.loopbackIPv4, 0); + final relay = _ProviderRelay._(server); + server.listen((request) async { + if (request.uri.path != '/relay') { + request.response.statusCode = HttpStatus.notFound; + await request.response.close(); + return; + } + final socket = await WebSocketTransformer.upgrade(request); + relay._sockets.add(socket); + socket.listen((data) { + final message = jsonDecode(data as String) as Map; + relay.messages.add(message); + handler(socket, message); + }); + }); + return relay; + } + + void send(WebSocket socket, Map message) => socket.add(jsonEncode(message)); + + Future close() async { + for (final socket in _sockets) { + await socket.close(); + } + await _server.close(force: true); + } +} + void main() { + test('generated relay versions match the protocol specification', () { + final spec = (jsonDecode(File('relay_protocol.json').readAsStringSync()) as Map).cast(); + expect(RelayProtocol.protocolVersion, spec['protocolVersion']); + expect(RelayProtocol.legacyProtocolVersion, spec['legacyProtocolVersion']); + }); + TestWidgetsFlutterBinding.ensureInitialized(); setUp(() { @@ -210,6 +453,389 @@ void main() { }); }); + group('WatchTogetherProvider — reconnect recovery', () { + test('restores only the matching host transport error and preserves session fields', () async { + final factory = _FakePeerServiceFactory(); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + final observedStates = []; + provider.addListener(() => observedStates.add(provider.session?.state)); + + await provider.createSession( + controlMode: ControlMode.anyone, + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + displayName: 'Host', + sessionId: 'room1', + mediaRatingKey: 'rating-1', + mediaServerId: 'server-1', + mediaTitle: 'Episode 1', + ); + await _flushProviderEvents(); + final established = provider.session!; + final service = factory.services.single; + observedStates.clear(); + service.broadcasts.clear(); + + service.emitError(_transportError()); + await _flushProviderEvents(); + expect(provider.session?.state, SessionState.error); + expect(provider.isConnected, isFalse); + + service.reconnect(); + await _flushProviderEvents(); + + expect(observedStates, [SessionState.error, SessionState.connected]); + expect(provider.isConnected, isTrue); + expect(provider.session?.errorMessage, isNull); + expect(provider.session?.sessionId, established.sessionId); + expect(provider.session?.role, established.role); + expect(provider.session?.controlMode, established.controlMode); + expect(provider.session?.mediaRatingKey, established.mediaRatingKey); + expect(provider.session?.mediaServerId, established.mediaServerId); + expect(provider.session?.mediaTitle, established.mediaTitle); + expect(service.broadcasts.where((message) => message.type == SyncMessageType.join), hasLength(1)); + + await provider.leaveSession(); + provider.dispose(); + }); + + test('guest waits for a reconnecting declared host before requesting state', () async { + final factory = _FakePeerServiceFactory(hostInitiallyConnected: false, rejectDisconnectedTargets: true); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + + await provider.joinSession( + 'late1', + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + displayName: 'Guest', + ); + await _flushProviderEvents(); + final service = factory.services.single; + final hostPeerId = provider.session!.hostPeerId!; + + expect(provider.session?.state, SessionState.connected); + expect(service.connectedPeers, isEmpty); + expect(service.directMessages.where((entry) => entry.$2.type == SyncMessageType.requestState), isEmpty); + + service.emitPeerConnected(hostPeerId); + await _flushProviderEvents(); + + expect(service.connectedPeers, [hostPeerId]); + expect( + service.directMessages.where( + (entry) => entry.$1 == hostPeerId && entry.$2.type == SyncMessageType.requestState, + ), + hasLength(1), + ); + expect(provider.session?.state, SessionState.connected); + + await provider.leaveSession(); + provider.dispose(); + }); + + test('guest recovery re-announces and re-requests authoritative state once', () async { + final factory = _FakePeerServiceFactory(); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + await provider.joinSession( + 'room2', + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + displayName: 'Guest', + ); + await _flushProviderEvents(); + final service = factory.services.single; + service.broadcasts.clear(); + service.directMessages.clear(); + + service.emitError(_transportError()); + await _flushProviderEvents(); + service.reconnect(); + await _flushProviderEvents(); + + expect(provider.session?.state, SessionState.connected); + expect(service.broadcasts.where((message) => message.type == SyncMessageType.join), hasLength(1)); + expect(service.directMessages.where((entry) => entry.$2.type == SyncMessageType.requestState), hasLength(1)); + + await provider.leaveSession(); + provider.dispose(); + }); + + test('relay errors remain terminal when the current service reconnects', () async { + final factory = _FakePeerServiceFactory(); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + final observedStates = []; + provider.addListener(() => observedStates.add(provider.session?.state)); + await provider.createSession( + controlMode: ControlMode.hostOnly, + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + sessionId: 'room3', + ); + await _flushProviderEvents(); + final service = factory.services.single; + observedStates.clear(); + + service.emitError( + const PeerError(type: PeerErrorType.serverError, message: 'Room was rejected', serverCode: 'room_rejected'), + ); + await _flushProviderEvents(); + service.reconnect(); + await _flushProviderEvents(); + + expect(observedStates, [SessionState.error]); + expect(provider.session?.state, SessionState.error); + expect(provider.session?.errorMessage, 'Room was rejected'); + expect(provider.isConnected, isFalse); + + await provider.leaveSession(); + provider.dispose(); + }); + + test('host-loss expiry supersedes a recoverable guest transport error', () { + fakeAsync((async) { + final factory = _FakePeerServiceFactory(); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + var joined = false; + unawaited( + provider + .joinSession('room4', relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, displayName: 'Guest') + .then((_) => joined = true), + ); + async.flushMicrotasks(); + expect(joined, isTrue); + final service = factory.services.single; + final hostPeerId = provider.session!.hostPeerId!; + + service.emitError(_transportError()); + async.flushMicrotasks(); + expect(provider.session?.state, SessionState.error); + + service.emitPeerDisconnected(hostPeerId); + async.flushMicrotasks(); + expect(provider.isWaitingForHostReconnect, isTrue); + async.elapse(const Duration(seconds: 15)); + async.flushMicrotasks(); + expect(provider.session?.errorMessage, 'Host left the session'); + + service.reconnect(); + async.flushMicrotasks(); + expect(provider.session?.state, SessionState.error); + expect(provider.session?.errorMessage, 'Host left the session'); + + unawaited(provider.leaveSession()); + async.flushMicrotasks(); + provider.dispose(); + async.flushMicrotasks(); + }); + }); + + test('stale service callback cannot mutate or resync a replacement session', () async { + final factory = _FakePeerServiceFactory(); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + var notifications = 0; + provider.addListener(() => notifications++); + + await provider.createSession( + controlMode: ControlMode.hostOnly, + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + sessionId: 'old1', + ); + await _flushProviderEvents(); + final oldService = factory.services.single; + final staleReconnect = oldService.onReconnected!; + oldService.emitError(_transportError('old transport error')); + await _flushProviderEvents(); + + await provider.createSession( + controlMode: ControlMode.anyone, + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + sessionId: 'new1', + ); + await _flushProviderEvents(); + final currentService = factory.services.last; + currentService.broadcasts.clear(); + notifications = 0; + final replacement = provider.session; + + staleReconnect(); + await _flushProviderEvents(); + expect(provider.session, replacement); + expect(notifications, 0); + expect(currentService.broadcasts, isEmpty); + + currentService.emitError(_transportError('current transport error')); + await _flushProviderEvents(); + currentService.reconnect(); + await _flushProviderEvents(); + expect(provider.session?.state, SessionState.connected); + expect(provider.session?.errorMessage, isNull); + expect(currentService.broadcasts.where((message) => message.type == SyncMessageType.join), hasLength(1)); + + await provider.leaveSession(); + provider.dispose(); + }); + }); + + group('WatchTogetherProvider — release cleanup', () { + test('release failure is surfaced after local session teardown completes', () async { + final releaseFailure = StateError('relay release failed'); + final factory = _FakePeerServiceFactory(releaseError: releaseFailure); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + addTearDown(provider.dispose); + + await provider.createSession( + controlMode: ControlMode.hostOnly, + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + sessionId: 'fail1', + ); + final service = factory.services.single; + + await expectLater(provider.leaveSession(), throwsA(same(releaseFailure))); + + expect(service.releaseCalls, 1); + expect(service.didDisconnect, isTrue); + expect(service.isDisposed, isTrue); + expect(provider.session, isNull); + expect(provider.isInSession, isFalse); + expect(provider.participants, isEmpty); + }); + }); + + group('WatchTogetherProvider — relay authority', () { + test('an empty successful probe joins the reserved room without creating', () async { + late final _ProviderRelay relay; + relay = await _ProviderRelay.start((socket, message) { + if (message['type'] == 'join') { + relay.send(socket, { + 'type': 'joined', + 'sessionId': message['sessionId'], + 'hostPeerId': _providerHostId, + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); + } else if (message['type'] == 'leave') { + relay.send(socket, { + 'type': 'left', + 'sessionId': message['sessionId'], + 'peerId': message['peerId'], + 'protocolVersion': 2, + }); + } + }); + addTearDown(relay.close); + final endpoint = WatchTogetherRelayEndpoint.resolve(relay.baseUrl); + final provider = WatchTogetherProvider(); + addTearDown(() async { + await provider.leaveSession(); + provider.dispose(); + }); + + final becameHost = await provider.enterRoom('empty1', relayEndpoint: endpoint, displayName: 'Guest'); + + expect(becameHost, isFalse); + expect(provider.isHost, isFalse); + expect(provider.session?.hostPeerId, _providerHostId); + final joins = relay.messages.where((message) => message['type'] == 'join').toList(); + expect(joins, hasLength(2)); + for (final join in joins) { + expect(join['protocolVersion'], 2); + expect(join['reconnectToken'], matches(RegExp(r'^[A-Za-z0-9_-]{43}$'))); + } + final leave = relay.messages.singleWhere((message) => message['type'] == 'leave'); + expect(leave['peerId'], joins.first['peerId']); + expect(leave['reconnectToken'], joins.first['reconnectToken']); + expect(leave['protocolVersion'], RelayProtocol.protocolVersion); + expect(joins.last['peerId'], isNot(joins.first['peerId'])); + expect(relay.messages.map((message) => message['type']).take(3), ['join', 'leave', 'join']); + expect(relay.messages.where((message) => message['type'] == 'create'), isEmpty); + }); + + test('a room-not-found probe creates with relay-declared host authority', () async { + late final _ProviderRelay relay; + relay = await _ProviderRelay.start((socket, message) { + if (message['type'] == 'join') { + relay.send(socket, {'type': 'error', 'code': 'room_not_found', 'message': 'Room not found'}); + } else if (message['type'] == 'create') { + relay.send(socket, { + 'type': 'created', + 'sessionId': message['sessionId'], + 'hostPeerId': message['peerId'], + 'reconnectToken': message['reconnectToken'], + 'protocolVersion': 2, + }); + } else if (message['type'] == 'endSession') { + relay.send(socket, {'type': 'ended', 'sessionId': message['sessionId'], 'protocolVersion': 2}); + } + }); + addTearDown(relay.close); + final endpoint = WatchTogetherRelayEndpoint.resolve(relay.baseUrl); + final provider = WatchTogetherProvider(); + addTearDown(() async { + await provider.leaveSession(); + provider.dispose(); + }); + + final becameHost = await provider.enterRoom('new01', relayEndpoint: endpoint, displayName: 'Host'); + + expect(becameHost, isTrue); + expect(provider.isHost, isTrue); + final create = relay.messages.singleWhere((message) => message['type'] == 'create'); + expect(provider.session?.hostPeerId, create['peerId']); + expect(create['peerId'], isNot('wt-NEW01')); + expect(create['protocolVersion'], 2); + expect(create['reconnectToken'], matches(RegExp(r'^[A-Za-z0-9_-]{43}$'))); + }); + }); + + group('WatchTogetherProvider — terminal room lifecycle', () { + test('relay ended notification exits immediately without release or reconnect grace', () async { + final factory = _FakePeerServiceFactory(); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + var hostExitCalls = 0; + provider.onHostExitedPlayer = () => hostExitCalls++; + + await provider.joinSession( + 'ended1', + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + displayName: 'Guest', + ); + final service = factory.services.single; + service.emitPeerDisconnected(provider.session!.hostPeerId!); + await _flushProviderEvents(); + expect(provider.isWaitingForHostReconnect, isTrue); + + service.emitSessionEnded(); + await _flushProviderEvents(); + + expect(hostExitCalls, 1); + expect(provider.session, isNull); + expect(provider.isWaitingForHostReconnect, isFalse); + expect(service.releaseCalls, 0); + expect(service.didDisconnect, isTrue); + expect(service.isDisposed, isTrue); + provider.dispose(); + }); + + test('best-effort host-leave cleanup observes release failures', () async { + final releaseFailure = StateError('relay release failed'); + final factory = _FakePeerServiceFactory(releaseError: releaseFailure); + final uncaught = []; + + await runZonedGuarded(() async { + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + await provider.joinSession( + 'leave2', + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + displayName: 'Guest', + ); + final service = factory.services.single; + service.emitMessage(SyncMessage.leave(peerId: provider.session!.hostPeerId!)); + await _flushProviderEvents(); + expect(provider.session, isNull); + expect(service.releaseCalls, 1); + provider.dispose(); + }, (error, _) => uncaught.add(error)); + + expect(uncaught, isEmpty); + }); + }); + group('WatchTogetherProvider — dispose hygiene', () { test('participantEvents stream is closed after dispose', () async { final p = WatchTogetherProvider(); @@ -223,5 +849,37 @@ void main() { await sub.cancel(); expect(streamDone, isTrue); }); + + test('dispose detaches local listeners before relay release completes', () async { + final releaseCompleter = Completer(); + final factory = _FakePeerServiceFactory(releaseBarrier: releaseCompleter.future); + final provider = WatchTogetherProvider(peerServiceFactory: factory.call); + await provider.joinSession( + 'dispose1', + relayEndpoint: WatchTogetherRelayEndpoint.defaultEndpoint, + displayName: 'Guest', + ); + final service = factory.services.single; + expect(service.hasRelayListeners, isTrue); + + provider.dispose(); + + expect(provider.session, isNull); + expect(provider.participants, isEmpty); + expect(provider.isWaitingForHostReconnect, isFalse); + expect(service.hasRelayListeners, isFalse); + expect(service.releaseCalls, 1); + expect(service.didDisconnect, isFalse); + + service.emitPeerConnected('late-peer'); + service.emitMessage(SyncMessage.join(peerId: 'late-peer', displayName: 'Late', isHost: false)); + await _flushProviderEvents(); + expect(provider.participants, isEmpty); + + releaseCompleter.complete(); + await _flushProviderEvents(); + expect(service.didDisconnect, isTrue); + expect(service.isDisposed, isTrue); + }); }); }