diff --git a/mobile/lib/providers/websocket.provider.dart b/mobile/lib/providers/websocket.provider.dart index 8d9bd5bfe3..3a0df83612 100644 --- a/mobile/lib/providers/websocket.provider.dart +++ b/mobile/lib/providers/websocket.provider.dart @@ -42,10 +42,13 @@ class WebsocketState { } class WebsocketNotifier extends StateNotifier { - WebsocketNotifier(this._ref) : super(const WebsocketState(socket: null, isConnected: false)); + WebsocketNotifier(this._ref, {Socket Function(dynamic, dynamic)? createSocket}) + : _createSocket = createSocket ?? io, + 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), @@ -56,12 +59,13 @@ class WebsocketNotifier extends StateNotifier { @override void dispose() { _batchDebouncer.dispose(); + state.socket?.dispose(); super.dispose(); } - /// Connects websocket to server unless already connected + /// Connects websocket to server unless a socket already exists void connect() { - if (state.isConnected) { + if (state.socket != null) { return; } final authenticationState = _ref.read(authProvider); @@ -71,7 +75,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 = io( + Socket socket = _createSocket( endpoint.origin, OptionBuilder() .setPath("${endpoint.path}/socket.io") @@ -84,6 +88,9 @@ 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((_) { dPrint(() => "Established Websocket Connection"); state = WebsocketState(isConnected: true, socket: socket); @@ -91,12 +98,12 @@ class WebsocketNotifier extends StateNotifier { socket.onDisconnect((_) { dPrint(() => "Disconnect to Websocket Connection"); - state = const WebsocketState(isConnected: false, socket: null); + state = WebsocketState(isConnected: false, socket: socket); }); socket.on('error', (errorMessage) { _log.severe("Websocket Error - $errorMessage"); - state = const WebsocketState(isConnected: false, socket: null); + state = WebsocketState(isConnected: false, socket: socket); }); socket.on('AssetUploadReadyV1', _handleSyncAssetUploadReadyV1); diff --git a/mobile/test/providers/websocket_provider_test.dart b/mobile/test/providers/websocket_provider_test.dart new file mode 100644 index 0000000000..9e0f31747e --- /dev/null +++ b/mobile/test/providers/websocket_provider_test.dart @@ -0,0 +1,173 @@ +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); + }); +}