diff --git a/mobile/lib/providers/app_life_cycle.provider.dart b/mobile/lib/providers/app_life_cycle.provider.dart index ad0940b776..45e6fb089d 100644 --- a/mobile/lib/providers/app_life_cycle.provider.dart +++ b/mobile/lib/providers/app_life_cycle.provider.dart @@ -81,6 +81,10 @@ class AppLifeCycleNotifier extends StateNotifier { await _ref.read(serverInfoProvider.notifier).getServerVersion(); } + if (!_shouldContinueOperation()) { + _wasPaused = true; + return; + } _ref.read(websocketProvider.notifier).connect(); await _handleBetaTimelineResume(); diff --git a/mobile/lib/providers/websocket.provider.dart b/mobile/lib/providers/websocket.provider.dart index 3a0df83612..6cf95ad1d1 100644 --- a/mobile/lib/providers/websocket.provider.dart +++ b/mobile/lib/providers/websocket.provider.dart @@ -42,13 +42,10 @@ class WebsocketState { } class WebsocketNotifier extends StateNotifier { - WebsocketNotifier(this._ref, {Socket Function(dynamic, dynamic)? createSocket}) - : _createSocket = createSocket ?? io, - super(const WebsocketState(socket: null, isConnected: false)); + WebsocketNotifier(this._ref) : super(const WebsocketState(socket: null, isConnected: false)); final _log = Logger('WebsocketNotifier'); final Ref _ref; - final Socket Function(dynamic, dynamic) _createSocket; final Debouncer _batchDebouncer = Debouncer( interval: const Duration(seconds: 5), @@ -63,11 +60,12 @@ class WebsocketNotifier extends StateNotifier { super.dispose(); } - /// Connects websocket to server unless a socket already exists + /// Connects websocket to server unless an active socket already exists void connect() { - if (state.socket != null) { + if (state.socket?.active == true) { return; } + state.socket?.dispose(); final authenticationState = _ref.read(authProvider); if (authenticationState.isAuthenticated) { @@ -75,7 +73,7 @@ class WebsocketNotifier extends StateNotifier { final endpoint = Uri.parse(Store.get(StoreKey.serverEndpoint)); dPrint(() => "Attempting to connect to websocket"); // Configure socket transports must be specified - Socket socket = _createSocket( + Socket socket = io( endpoint.origin, OptionBuilder() .setPath("${endpoint.path}/socket.io") @@ -88,7 +86,6 @@ class WebsocketNotifier extends StateNotifier { .build(), ); - // Hold the socket now so disconnect() can tear it down even if it never connects state = WebsocketState(isConnected: false, socket: socket); socket.onConnect((_) { diff --git a/mobile/test/providers/app_life_cycle_provider_test.dart b/mobile/test/providers/app_life_cycle_provider_test.dart new file mode 100644 index 0000000000..d4619e8bc0 --- /dev/null +++ b/mobile/test/providers/app_life_cycle_provider_test.dart @@ -0,0 +1,202 @@ +import 'dart:async'; + +import 'package:flutter_test/flutter_test.dart'; +import 'package:hooks_riverpod/hooks_riverpod.dart'; +import 'package:immich_mobile/domain/models/config/app_config.dart'; +import 'package:immich_mobile/domain/models/log.model.dart'; +import 'package:immich_mobile/domain/services/background_worker.service.dart'; +import 'package:immich_mobile/domain/services/log.service.dart'; +import 'package:immich_mobile/models/auth/auth_state.model.dart'; +import 'package:immich_mobile/models/server_info/server_version.model.dart'; +import 'package:immich_mobile/providers/app_life_cycle.provider.dart'; +import 'package:immich_mobile/providers/auth.provider.dart'; +import 'package:immich_mobile/providers/backup/drift_backup.provider.dart'; +import 'package:immich_mobile/providers/infrastructure/platform.provider.dart'; +import 'package:immich_mobile/providers/server_info.provider.dart'; +import 'package:immich_mobile/providers/websocket.provider.dart'; +import 'package:immich_mobile/services/auth.service.dart'; +import 'package:immich_mobile/services/background_upload.service.dart'; +import 'package:immich_mobile/services/foreground_upload.service.dart'; +import 'package:immich_mobile/services/secure_storage.service.dart'; +import 'package:immich_mobile/services/server_info.service.dart'; +import 'package:immich_mobile/services/widget.service.dart'; +import 'package:immich_mobile/utils/upload_speed_calculator.dart'; +import 'package:mocktail/mocktail.dart'; + +import '../domain/service.mock.dart'; +import '../infrastructure/repository.mock.dart'; +import '../service.mocks.dart'; + +class MockAuthService extends Mock implements AuthService {} + +class MockSecureStorageService extends Mock implements SecureStorageService {} + +class MockWidgetService extends Mock implements WidgetService {} + +class MockServerInfoService extends Mock implements ServerInfoService {} + +class MockForegroundUploadService extends Mock implements ForegroundUploadService {} + +class MockBackgroundUploadService extends Mock implements BackgroundUploadService {} + +class MockBackgroundWorkerLockService extends Mock implements BackgroundWorkerLockService {} + +class FakeLogMessage extends Fake implements LogMessage {} + +class TestAuthNotifier extends AuthNotifier { + TestAuthNotifier(Ref ref) + : super( + MockAuthService(), + MockApiService(), + MockUserService(), + MockSecureStorageService(), + MockWidgetService(), + ref, + ) { + state = const AuthState( + deviceId: 'device-1', + userId: 'user-1', + userEmail: 'user@example.com', + name: 'User', + profileImagePath: '', + isAdmin: false, + isAuthenticated: true, + ); + } + + @override + Future setOpenApiServiceEndpoint() async => 'http://test-server.com'; +} + +class TestWebsocketNotifier extends WebsocketNotifier { + TestWebsocketNotifier(super.ref); + + int connectCount = 0; + int disconnectCount = 0; + final connectCalled = Completer(); + + @override + void connect() { + connectCount++; + if (!connectCalled.isCompleted) { + connectCalled.complete(); + } + throw StateError('unexpected websocket connection'); + } + + @override + void disconnect() => disconnectCount++; +} + +class TestDriftBackupNotifier extends DriftBackupNotifier { + TestDriftBackupNotifier() : super(MockForegroundUploadService(), MockBackgroundUploadService(), UploadSpeedManager()); +} + +void main() { + late LogService logService; + late Completer serverVersion; + late MockServerInfoService serverInfoService; + late MockBackgroundWorkerLockService lockService; + late ProviderContainer container; + late TestWebsocketNotifier websocket; + late AppLifeCycleNotifier lifeCycle; + late int serverVersionCount; + + setUpAll(() async { + final logRepository = MockLogRepository(); + final settingsRepository = MockSettingsRepository(); + registerFallbackValue(FakeLogMessage()); + when(() => logRepository.truncate(limit: any(named: 'limit'))).thenAnswer((_) async {}); + when(() => logRepository.insert(any())).thenAnswer((_) async => true); + when(() => settingsRepository.appConfig).thenReturn(const AppConfig(logLevel: LogLevel.info)); + logService = await LogService.init( + logRepository: logRepository, + settingsRepository: settingsRepository, + shouldBuffer: false, + ); + }); + + tearDownAll(() => logService.dispose()); + + setUp(() { + serverVersion = Completer(); + serverInfoService = MockServerInfoService(); + lockService = MockBackgroundWorkerLockService(); + serverVersionCount = 0; + + when(() => serverInfoService.getServerVersion()).thenAnswer((_) { + serverVersionCount++; + return serverVersionCount == 1 ? serverVersion.future : Future.value(); + }); + when(() => lockService.lock()).thenAnswer((_) async {}); + when(() => lockService.unlock()).thenAnswer((_) async {}); + + container = ProviderContainer( + overrides: [ + authProvider.overrideWith(TestAuthNotifier.new), + serverInfoProvider.overrideWith((_) => ServerInfoNotifier(serverInfoService)), + websocketProvider.overrideWith((ref) { + return websocket = TestWebsocketNotifier(ref); + }), + driftBackupProvider.overrideWith((_) => TestDriftBackupNotifier()), + backgroundWorkerLockServiceProvider.overrideWithValue(lockService), + ], + ); + lifeCycle = container.read(appStateProvider.notifier); + }); + + tearDown(() => container.dispose()); + + Future startResume() async { + await lifeCycle.handleAppPause(); + lifeCycle.handleAppResume(); + await untilCalled(() => serverInfoService.getServerVersion()); + } + + Future releaseResume() async { + serverVersion.complete(); + await Future.delayed(Duration.zero); + } + + test('pause during resume does not reconnect websocket', () async { + await startResume(); + await lifeCycle.handleAppPause(); + await releaseResume(); + + expect(lifeCycle.getAppState(), AppLifeCycleEnum.paused); + expect(serverVersionCount, 1); + expect(websocket.disconnectCount, 2); + expect(websocket.connectCount, 0); + }); + + test('inactive resume retries when the app resumes again', () async { + await startResume(); + lifeCycle.handleAppInactivity(); + await releaseResume(); + + lifeCycle.handleAppResume(); + await websocket.connectCalled.future; + + expect(lifeCycle.getAppState(), AppLifeCycleEnum.resumed); + expect(serverVersionCount, 2); + expect(websocket.disconnectCount, 1); + expect(websocket.connectCount, 1); + }); + + test('pause after an inactive abort resumes once', () async { + await startResume(); + lifeCycle.handleAppInactivity(); + await releaseResume(); + await lifeCycle.handleAppPause(); + + lifeCycle.handleAppResume(); + await websocket.connectCalled.future; + lifeCycle.handleAppResume(); + await Future.delayed(Duration.zero); + + expect(lifeCycle.getAppState(), AppLifeCycleEnum.resumed); + expect(serverVersionCount, 2); + expect(websocket.disconnectCount, 2); + expect(websocket.connectCount, 1); + }); +} diff --git a/mobile/test/providers/websocket_provider_test.dart b/mobile/test/providers/websocket_provider_test.dart deleted file mode 100644 index 9e0f31747e..0000000000 --- a/mobile/test/providers/websocket_provider_test.dart +++ /dev/null @@ -1,173 +0,0 @@ -import 'package:drift/drift.dart' hide isNull, isNotNull; -import 'package:drift/native.dart'; -import 'package:flutter/services.dart'; -import 'package:flutter_test/flutter_test.dart'; -import 'package:hooks_riverpod/hooks_riverpod.dart'; -import 'package:immich_mobile/domain/models/store.model.dart'; -import 'package:immich_mobile/domain/services/store.service.dart'; -import 'package:immich_mobile/domain/services/user.service.dart'; -import 'package:immich_mobile/entities/store.entity.dart'; -import 'package:immich_mobile/infrastructure/repositories/db.repository.dart'; -import 'package:immich_mobile/infrastructure/repositories/store.repository.dart'; -import 'package:immich_mobile/models/auth/auth_state.model.dart'; -import 'package:immich_mobile/providers/auth.provider.dart'; -import 'package:immich_mobile/providers/websocket.provider.dart'; -import 'package:immich_mobile/services/api.service.dart'; -import 'package:immich_mobile/services/auth.service.dart'; -import 'package:immich_mobile/services/secure_storage.service.dart'; -import 'package:immich_mobile/services/widget.service.dart'; -import 'package:mocktail/mocktail.dart'; -import 'package:socket_io_client/socket_io_client.dart'; - -class MockAuthService extends Mock implements AuthService {} - -class MockApiService extends Mock implements ApiService {} - -class MockUserService extends Mock implements UserService {} - -class MockSecureStorageService extends Mock implements SecureStorageService {} - -class MockWidgetService extends Mock implements WidgetService {} - -class TestAuthNotifier extends AuthNotifier { - TestAuthNotifier(Ref ref, AuthState initial) - : super( - MockAuthService(), - MockApiService(), - MockUserService(), - MockSecureStorageService(), - MockWidgetService(), - ref, - ) { - state = initial; - } -} - -class _FakeSocket extends Fake implements Socket { - int disposeCount = 0; - final Map handlers = {}; - - @override - Function() on(String event, dynamic Function(dynamic) handler) { - handlers[event] = handler; - return () {}; - } - - @override - void dispose() => disposeCount++; -} - -AuthState _authState({required bool isAuthenticated}) { - return AuthState( - deviceId: 'device-1', - userId: 'user-1', - userEmail: 'user@example.com', - isAuthenticated: isAuthenticated, - name: 'User', - isAdmin: false, - profileImagePath: '', - ); -} - -void main() { - TestWidgetsFlutterBinding.ensureInitialized(); - - late Drift db; - - setUpAll(() async { - TestDefaultBinaryMessengerBinding.instance.defaultBinaryMessenger.setMockMethodCallHandler( - const MethodChannel('plugins.flutter.io/path_provider'), - (MethodCall methodCall) async => 'test', - ); - db = Drift(DatabaseConnection(NativeDatabase.memory(), closeStreamsSynchronously: true)); - await StoreService.init(storeRepository: DriftStoreRepository(db)); - await Store.put(StoreKey.serverEndpoint, 'http://test-server.com'); - }); - - tearDownAll(() async { - await db.close(); - }); - - ProviderContainer buildContainer(List<_FakeSocket> created, {required bool isAuthenticated}) { - return ProviderContainer( - overrides: [ - authProvider.overrideWith((ref) => TestAuthNotifier(ref, _authState(isAuthenticated: isAuthenticated))), - websocketProvider.overrideWith( - (ref) => WebsocketNotifier( - ref, - createSocket: (_, _) { - final socket = _FakeSocket(); - created.add(socket); - return socket; - }, - ), - ), - ], - ); - } - - test('connect creates a socket and holds it in state before it connects', () { - final created = <_FakeSocket>[]; - final container = buildContainer(created, isAuthenticated: true); - addTearDown(container.dispose); - - container.read(websocketProvider.notifier).connect(); - - expect(created, hasLength(1)); - expect(container.read(websocketProvider).socket, isNotNull); - expect(container.read(websocketProvider).isConnected, isFalse); - }); - - test('connect does not create a second socket while one already exists', () { - final created = <_FakeSocket>[]; - final container = buildContainer(created, isAuthenticated: true); - addTearDown(container.dispose); - - final notifier = container.read(websocketProvider.notifier); - notifier.connect(); - notifier.connect(); - - expect(created, hasLength(1)); - }); - - test('disconnect disposes a socket that never connected (unreachable server)', () { - final created = <_FakeSocket>[]; - final container = buildContainer(created, isAuthenticated: true); - addTearDown(container.dispose); - - final notifier = container.read(websocketProvider.notifier); - notifier.connect(); - notifier.disconnect(); - - expect(created.single.disposeCount, 1); - expect(container.read(websocketProvider).socket, isNull); - }); - - test('socket stays disposable after an error so disconnect can stop it', () { - final created = <_FakeSocket>[]; - final container = buildContainer(created, isAuthenticated: true); - addTearDown(container.dispose); - - final notifier = container.read(websocketProvider.notifier); - notifier.connect(); - - // Simulate the cert-error the reporter hits off-network. - created.single.handlers['error']?.call('CERTIFICATE_VERIFY_FAILED'); - expect(container.read(websocketProvider).socket, isNotNull); - - notifier.disconnect(); - expect(created.single.disposeCount, 1); - expect(container.read(websocketProvider).socket, isNull); - }); - - test('connect is a no-op when not authenticated', () { - final created = <_FakeSocket>[]; - final container = buildContainer(created, isAuthenticated: false); - addTearDown(container.dispose); - - container.read(websocketProvider.notifier).connect(); - - expect(created, isEmpty); - expect(container.read(websocketProvider).socket, isNull); - }); -}