From 9df7449e5ca2d30fb46deade4b87d3cff8b25eb4 Mon Sep 17 00:00:00 2001 From: PhucNguyen Date: Mon, 17 Aug 2026 12:31:48 +0700 Subject: [PATCH] Harden mobile screen sharing lifecycle --- apps/mobile/README.md | 6 + apps/mobile/app/(app)/session/[id].tsx | 137 +++++----- apps/mobile/src/hooks/useScreenShare.test.ts | 268 +++++++++++++++++++ apps/mobile/src/hooks/useScreenShare.ts | 226 ++++++++++++++++ apps/mobile/src/hooks/useWebRTCHost.test.ts | 143 +++++++++- apps/mobile/src/hooks/useWebRTCHost.ts | 212 ++++++++++----- apps/mobile/src/lib/host-session.test.ts | 41 +++ apps/mobile/src/lib/host-session.ts | 26 ++ apps/mobile/src/test/setup.ts | 23 +- 9 files changed, 938 insertions(+), 144 deletions(-) create mode 100644 apps/mobile/src/hooks/useScreenShare.test.ts create mode 100644 apps/mobile/src/hooks/useScreenShare.ts create mode 100644 apps/mobile/src/lib/host-session.test.ts create mode 100644 apps/mobile/src/lib/host-session.ts diff --git a/apps/mobile/README.md b/apps/mobile/README.md index eca5a69f..eae78ec2 100644 --- a/apps/mobile/README.md +++ b/apps/mobile/README.md @@ -87,6 +87,12 @@ notice only in Task Manager when notification permission is denied. The generate camera and system-overlay permissions inherited from the WebRTC dependency because PairUX currently uses screen capture and voice, not camera capture or overlay windows. +The host UI reports sharing as active only after the captured stream has been published to the +current viewers. Capture permission, publication, active sharing, and shutdown are serialized so +repeated taps cannot start duplicate MediaProjection sessions. PairUX also stops and unpublishes the +capture when Android ends it from the system controls, when the host leaves the session tab, or when +the session screen unmounts. + The generated iOS app includes native WebRTC and supports joining sessions and voice chat. Full device screen broadcasting on iOS additionally requires a ReplayKit Broadcast Upload Extension, an App Group, and matching Apple signing entitlements. Those are a separate native milestone; do diff --git a/apps/mobile/app/(app)/session/[id].tsx b/apps/mobile/app/(app)/session/[id].tsx index 337727bf..21f8bc8e 100644 --- a/apps/mobile/app/(app)/session/[id].tsx +++ b/apps/mobile/app/(app)/session/[id].tsx @@ -6,16 +6,16 @@ * - role: 'host' | 'viewer' * - participantId: viewer's participant ID (viewer mode only) */ -import { useState, useCallback, useEffect } from 'react'; -import { View, Text, TouchableOpacity, Alert, ActivityIndicator } from 'react-native'; -import { useLocalSearchParams, useRouter } from 'expo-router'; -import type { MediaStream } from 'react-native-webrtc'; -import { mediaDevices } from 'react-native-webrtc'; +import { useState, useCallback, useEffect, useRef } from 'react'; +import { View, Text, TouchableOpacity, Alert, ActivityIndicator, Platform } from 'react-native'; +import { useFocusEffect, useLocalSearchParams, useRouter } from 'expo-router'; import { useAuth } from '@/contexts/AuthContext'; import { useWebRTCHost } from '@/hooks/useWebRTCHost'; import { useWebRTCViewer } from '@/hooks/useWebRTCViewer'; +import { useScreenShare } from '@/hooks/useScreenShare'; import { useChat } from '@/hooks/useChat'; import { sessionApi } from '@/lib/api/sessions'; +import { endHostSession } from '@/lib/host-session'; import { VideoViewer } from '@/components/VideoViewer'; import { ChatPanel } from '@/components/ChatPanel'; import { SessionInfo } from '@/components/SessionInfo'; @@ -35,8 +35,6 @@ export default function SessionScreen() { const participantId = (params.participantId as string | undefined) ?? user?.id ?? ''; const [joinCode, setJoinCode] = useState(''); - const [localStream, setLocalStream] = useState(null); - const [isSharing, setIsSharing] = useState(false); const [loadingSession, setLoadingSession] = useState(true); // Fetch session details @@ -64,10 +62,6 @@ export default function SessionScreen() { sessionId={sessionId} hostId={user?.id ?? ''} joinCode={joinCode} - localStream={localStream} - isSharing={isSharing} - setLocalStream={setLocalStream} - setIsSharing={setIsSharing} chat={chat} currentUserId={user?.id} router={router} @@ -95,10 +89,6 @@ interface HostSessionProps { sessionId: string; hostId: string; joinCode: string; - localStream: MediaStream | null; - isSharing: boolean; - setLocalStream: (s: MediaStream | null) => void; - setIsSharing: (b: boolean) => void; chat: ReturnType; currentUserId?: string; router: ReturnType; @@ -109,10 +99,6 @@ function HostSession({ sessionId, hostId, joinCode, - localStream, - isSharing, - setLocalStream, - setIsSharing, chat, currentUserId, router, @@ -121,8 +107,13 @@ function HostSession({ const webrtc = useWebRTCHost({ sessionId, hostId, - localStream, }); + const screenShare = useScreenShare({ + publishStream: webrtc.publishStream, + unpublishStream: webrtc.unpublishStream, + }); + const stopScreenShare = screenShare.stop; + const endingSessionRef = useRef(false); // Auto-start hosting when session loads useEffect(() => { @@ -132,28 +123,13 @@ function HostSession({ // eslint-disable-next-line react-hooks/exhaustive-deps }, [loadingSession]); - const startScreenShare = useCallback(async () => { - try { - const stream = await mediaDevices.getDisplayMedia(); - setLocalStream(stream); - setIsSharing(true); - await webrtc.publishStream(stream); - } catch (err) { - console.error('[Session] Failed to start screen share:', err); - Alert.alert('Error', 'Failed to start screen sharing. Please check permissions.'); - } - }, [webrtc, setLocalStream, setIsSharing]); - - const stopScreenShare = useCallback(async () => { - if (localStream) { - localStream.getTracks().forEach((track) => { - track.stop(); - }); - setLocalStream(null); - } - setIsSharing(false); - await webrtc.unpublishStream(); - }, [localStream, webrtc, setLocalStream, setIsSharing]); + useFocusEffect( + useCallback(() => { + return () => { + void stopScreenShare(); + }; + }, [stopScreenShare]) + ); function handleEndSession() { Alert.alert('End Session', 'Are you sure you want to end this session?', [ @@ -162,11 +138,23 @@ function HostSession({ text: 'End', style: 'destructive', onPress: () => { + if (endingSessionRef.current) return; + endingSessionRef.current = true; void (async () => { - await stopScreenShare(); - webrtc.stopHosting(); - await sessionApi.end(sessionId); - router.back(); + try { + await endHostSession({ + stopScreenShare: screenShare.stop, + endSession: () => sessionApi.end(sessionId), + stopHosting: webrtc.stopHosting, + goBack: () => { + router.back(); + }, + }); + } catch (endError) { + console.error('[Session] Failed to end session:', endError); + endingSessionRef.current = false; + Alert.alert('Unable to end session', 'Please check your connection and try again.'); + } })(); }, }, @@ -188,25 +176,50 @@ function HostSession({ {webrtc.error} )} + {screenShare.error && ( + + {screenShare.error} + + )} {/* Main content */} {!webrtc.isHosting ? ( - ) : !isSharing ? ( + ) : screenShare.isBusy ? ( - Ready to share - - Tap the button below to start sharing your screen with viewers. + + + {screenShare.state === 'requesting' + ? 'Waiting for screen capture permission...' + : screenShare.state === 'publishing' + ? 'Connecting the screen share to viewers...' + : 'Stopping screen sharing...'} - { - void startScreenShare(); - }} - className="mt-4 rounded-xl bg-primary-600 px-8 py-4" - > - Share Screen - + + ) : !screenShare.isSharing ? ( + + Ready to share + {Platform.OS === 'ios' ? ( + + Full-device sharing on iOS requires ReplayKit and is not available in this build. + + ) : ( + <> + + Tap the button below to start sharing your screen with viewers. + + { + void screenShare.start(); + }} + disabled={screenShare.isBusy} + className="mt-4 rounded-xl bg-primary-600 px-8 py-4" + > + Share Screen + + + )} ) : ( @@ -219,11 +232,12 @@ function HostSession({ {/* Controls */} - {isSharing && ( + {screenShare.isSharing && ( { - void stopScreenShare(); + void screenShare.stop(); }} + disabled={screenShare.isBusy} className="rounded-lg bg-yellow-600 px-4 py-2" > Stop Sharing @@ -277,6 +291,7 @@ function ViewerSession({ router, loadingSession: _loadingSession, }: ViewerSessionProps) { + const leavingSessionRef = useRef(false); const webrtc = useWebRTCViewer({ sessionId, participantId, @@ -299,6 +314,8 @@ function ViewerSession({ text: 'Leave', style: 'destructive', onPress: () => { + if (leavingSessionRef.current) return; + leavingSessionRef.current = true; webrtc.disconnect(); router.back(); }, diff --git a/apps/mobile/src/hooks/useScreenShare.test.ts b/apps/mobile/src/hooks/useScreenShare.test.ts new file mode 100644 index 00000000..29b8153e --- /dev/null +++ b/apps/mobile/src/hooks/useScreenShare.test.ts @@ -0,0 +1,268 @@ +import { act, renderHook, waitFor } from '@testing-library/react'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { mediaDevices } from 'react-native-webrtc'; +import type { MediaStream, MediaStreamTrack } from 'react-native-webrtc'; +import { useScreenShare } from './useScreenShare'; + +function deferred() { + let resolve!: (value: T | PromiseLike) => void; + const promise = new Promise((resolvePromise) => { + resolve = resolvePromise; + }); + return { promise, resolve }; +} + +function createCapture() { + const endedListeners = new Set<() => void>(); + const track = { + id: 'screen-video', + kind: 'video', + stop: vi.fn(), + addEventListener: vi.fn((event: string, listener: () => void) => { + if (event === 'ended') endedListeners.add(listener); + }), + removeEventListener: vi.fn((event: string, listener: () => void) => { + if (event === 'ended') endedListeners.delete(listener); + }), + } as unknown as MediaStreamTrack; + const stream = { + getTracks: vi.fn(() => [track]), + } as unknown as MediaStream; + + return { + stream, + track, + end: () => { + for (const listener of endedListeners) listener(); + }, + }; +} + +describe('useScreenShare', () => { + beforeEach(() => { + vi.mocked(mediaDevices.getDisplayMedia).mockReset(); + }); + + it('only reports active after capture publishing succeeds', async () => { + const capture = createCapture(); + const publication = deferred(); + const publishStream = vi.fn(() => publication.promise); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(capture.stream); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + + let startPromise!: Promise; + act(() => { + startPromise = result.current.start(); + }); + await waitFor(() => expect(result.current.state).toBe('publishing')); + expect(result.current.isSharing).toBe(false); + + publication.resolve(undefined); + await act(async () => { + await expect(startPromise).resolves.toBe(true); + }); + + expect(result.current.state).toBe('active'); + expect(result.current.isSharing).toBe(true); + }); + + it('prevents duplicate capture requests', async () => { + const captureRequest = deferred(); + const publishStream = vi.fn().mockResolvedValue(undefined); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockReturnValue(captureRequest.promise); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + + let firstStart!: Promise; + let secondStart!: Promise; + act(() => { + firstStart = result.current.start(); + secondStart = result.current.start(); + }); + await expect(secondStart).resolves.toBe(false); + expect(mediaDevices.getDisplayMedia).toHaveBeenCalledTimes(1); + + captureRequest.resolve(createCapture().stream); + await act(async () => { + await expect(firstStart).resolves.toBe(true); + }); + }); + + it('rolls back capture when publishing fails', async () => { + const capture = createCapture(); + const publishStream = vi.fn().mockRejectedValue(new Error('Viewer signaling failed')); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(capture.stream); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + + await act(async () => { + await expect(result.current.start()).resolves.toBe(false); + }); + + expect(capture.track.stop).toHaveBeenCalledTimes(1); + expect(unpublishStream).toHaveBeenCalledTimes(1); + expect(result.current.state).toBe('idle'); + expect(result.current.error).toBe( + 'Could not share your screen with viewers. Check your connection and try again.' + ); + expect(result.current.error).not.toContain('Viewer signaling failed'); + }); + + it('preserves a capture permission error before publishing starts', async () => { + const publishStream = vi.fn().mockResolvedValue(undefined); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockRejectedValue( + new Error('Screen capture permission denied') + ); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + + await act(async () => { + await expect(result.current.start()).resolves.toBe(false); + }); + + expect(publishStream).not.toHaveBeenCalled(); + expect(unpublishStream).not.toHaveBeenCalled(); + expect(result.current.state).toBe('idle'); + expect(result.current.error).toBe('Screen capture permission denied'); + }); + + it('discards a capture that resolves after stop', async () => { + const capture = createCapture(); + const captureRequest = deferred(); + const publishStream = vi.fn().mockResolvedValue(undefined); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockReturnValue(captureRequest.promise); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + + let startPromise!: Promise; + act(() => { + startPromise = result.current.start(); + }); + await act(async () => { + await result.current.stop(); + }); + + captureRequest.resolve(capture.stream); + await act(async () => { + await expect(startPromise).resolves.toBe(false); + }); + + expect(capture.track.stop).toHaveBeenCalledTimes(1); + expect(publishStream).not.toHaveBeenCalled(); + expect(unpublishStream).not.toHaveBeenCalled(); + expect(result.current.state).toBe('idle'); + + const nextCapture = createCapture(); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(nextCapture.stream); + await act(async () => { + await expect(result.current.start()).resolves.toBe(true); + }); + expect(result.current.state).toBe('active'); + }); + + it('waits for an in-flight publication before unpublishing', async () => { + const capture = createCapture(); + const publication = deferred(); + const publishStream = vi.fn(() => publication.promise); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(capture.stream); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + + let startPromise!: Promise; + act(() => { + startPromise = result.current.start(); + }); + await waitFor(() => expect(result.current.state).toBe('publishing')); + + let stopPromise!: Promise; + act(() => { + stopPromise = result.current.stop(); + }); + expect(capture.track.stop).toHaveBeenCalledTimes(1); + expect(unpublishStream).not.toHaveBeenCalled(); + + publication.resolve(undefined); + await act(async () => { + await Promise.all([startPromise, stopPromise]); + }); + + expect(unpublishStream).toHaveBeenCalledTimes(1); + expect(result.current.state).toBe('idle'); + }); + + it('stops sharing when the native capture ends', async () => { + const capture = createCapture(); + const publishStream = vi.fn().mockResolvedValue(undefined); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(capture.stream); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + await act(async () => { + await result.current.start(); + }); + + act(() => { + capture.end(); + }); + await waitFor(() => expect(result.current.state).toBe('idle')); + + expect(capture.track.stop).toHaveBeenCalledTimes(1); + expect(unpublishStream).toHaveBeenCalledTimes(1); + }); + + it('serializes duplicate stop requests', async () => { + const capture = createCapture(); + const unpublication = deferred(); + const publishStream = vi.fn().mockResolvedValue(undefined); + const unpublishStream = vi.fn(() => unpublication.promise); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(capture.stream); + + const { result } = renderHook(() => useScreenShare({ publishStream, unpublishStream })); + await act(async () => { + await result.current.start(); + }); + + let firstStop!: Promise; + let secondStop!: Promise; + act(() => { + firstStop = result.current.stop(); + secondStop = result.current.stop(); + }); + await act(async () => { + await Promise.resolve(); + }); + expect(unpublishStream).toHaveBeenCalledTimes(1); + + unpublication.resolve(undefined); + await act(async () => { + await Promise.all([firstStop, secondStop]); + }); + expect(result.current.state).toBe('idle'); + }); + + it('stops capture and unpublishes on unmount', async () => { + const capture = createCapture(); + const publishStream = vi.fn().mockResolvedValue(undefined); + const unpublishStream = vi.fn().mockResolvedValue(undefined); + vi.mocked(mediaDevices.getDisplayMedia).mockResolvedValue(capture.stream); + + const { result, unmount } = renderHook(() => + useScreenShare({ publishStream, unpublishStream }) + ); + await act(async () => { + await result.current.start(); + }); + + unmount(); + + expect(capture.track.stop).toHaveBeenCalledTimes(1); + await waitFor(() => expect(unpublishStream).toHaveBeenCalledTimes(1)); + }); +}); diff --git a/apps/mobile/src/hooks/useScreenShare.ts b/apps/mobile/src/hooks/useScreenShare.ts new file mode 100644 index 00000000..243e8585 --- /dev/null +++ b/apps/mobile/src/hooks/useScreenShare.ts @@ -0,0 +1,226 @@ +import { useCallback, useEffect, useRef, useState } from 'react'; +import { mediaDevices } from 'react-native-webrtc'; +import type { MediaStream, MediaStreamTrack } from 'react-native-webrtc'; + +export type ScreenShareState = 'idle' | 'requesting' | 'publishing' | 'active' | 'stopping'; + +interface UseScreenShareOptions { + publishStream: (stream: MediaStream) => Promise; + unpublishStream: () => Promise; +} + +interface UseScreenShareReturn { + state: ScreenShareState; + isSharing: boolean; + isBusy: boolean; + error: string | null; + start: () => Promise; + stop: () => Promise; + clearError: () => void; +} + +function stopTracks(stream: MediaStream): void { + for (const track of stream.getTracks()) { + track.stop(); + } +} + +function getStartError(error: unknown): string { + if (error instanceof Error && error.message.trim()) { + return error.message; + } + return 'Unable to start screen sharing. Please check the capture permission and try again.'; +} + +/** Owns the native capture stream and serializes every screen-share transition. */ +export function useScreenShare({ + publishStream, + unpublishStream, +}: UseScreenShareOptions): UseScreenShareReturn { + const [state, setState] = useState('idle'); + const [error, setError] = useState(null); + + const stateRef = useRef('idle'); + const streamRef = useRef(null); + const generationRef = useRef(0); + const mountedRef = useRef(true); + const stopPromiseRef = useRef | null>(null); + const publicationRef = useRef | null>(null); + const endedListenersRef = useRef void>>(new Map()); + + const transition = useCallback((nextState: ScreenShareState) => { + stateRef.current = nextState; + if (mountedRef.current) { + setState(nextState); + } + }, []); + + const detachEndedListeners = useCallback(() => { + for (const [track, listener] of endedListenersRef.current) { + track.removeEventListener('ended', listener); + } + endedListenersRef.current.clear(); + }, []); + + const stop = useCallback(async (): Promise => { + if (stopPromiseRef.current) { + return stopPromiseRef.current; + } + + if (stateRef.current === 'idle' && !streamRef.current) { + return; + } + + generationRef.current += 1; + transition('stopping'); + + const stream = streamRef.current; + const publication = publicationRef.current; + streamRef.current = null; + detachEndedListeners(); + if (stream) { + stopTracks(stream); + } + + const stopPromise = (async () => { + await Promise.resolve(); + try { + if (stream) { + try { + await publication; + } catch { + // Publishing performs its own rollback; local teardown still continues. + } + await unpublishStream(); + } + } catch (unpublishError) { + console.error('[ScreenShare] Failed to unpublish capture:', unpublishError); + } finally { + transition('idle'); + stopPromiseRef.current = null; + } + })(); + + stopPromiseRef.current = stopPromise; + return stopPromise; + }, [detachEndedListeners, transition, unpublishStream]); + + const attachEndedListeners = useCallback( + (stream: MediaStream) => { + for (const track of stream.getTracks()) { + const listener = () => { + void stop(); + }; + track.addEventListener('ended', listener); + endedListenersRef.current.set(track, listener); + } + }, + [stop] + ); + + const start = useCallback(async (): Promise => { + if (stateRef.current !== 'idle' || stopPromiseRef.current) { + return false; + } + + const generation = ++generationRef.current; + transition('requesting'); + setError(null); + + let stream: MediaStream | null = null; + let publication: Promise | null = null; + try { + stream = await mediaDevices.getDisplayMedia(); + + if (generationRef.current !== generation) { + stopTracks(stream); + return false; + } + + streamRef.current = stream; + attachEndedListeners(stream); + transition('publishing'); + publication = publishStream(stream); + publicationRef.current = publication; + await publication; + + if (generationRef.current !== generation || streamRef.current !== stream) { + return false; + } + + transition('active'); + return true; + } catch (startError) { + const ownsTransition = generationRef.current === generation; + if (ownsTransition) { + if (streamRef.current === stream) { + streamRef.current = null; + } + detachEndedListeners(); + if (stream) { + stopTracks(stream); + try { + await unpublishStream(); + } catch (rollbackError) { + console.error('[ScreenShare] Failed to roll back capture:', rollbackError); + } + } + } + + if (mountedRef.current && ownsTransition) { + setError( + stream + ? 'Could not share your screen with viewers. Check your connection and try again.' + : getStartError(startError) + ); + transition('idle'); + } + return false; + } finally { + if (publicationRef.current === publication) { + publicationRef.current = null; + } + } + }, [attachEndedListeners, detachEndedListeners, publishStream, transition, unpublishStream]); + + const clearError = useCallback(() => { + setError(null); + }, []); + + useEffect(() => { + mountedRef.current = true; + return () => { + mountedRef.current = false; + generationRef.current += 1; + const stream = streamRef.current; + const publication = publicationRef.current; + streamRef.current = null; + detachEndedListeners(); + if (stream) { + stopTracks(stream); + void (async () => { + try { + await publication; + } catch { + // Publishing performs its own rollback; local teardown still continues. + } + try { + await unpublishStream(); + } catch (unpublishError) { + console.error('[ScreenShare] Failed to unpublish during cleanup:', unpublishError); + } + })(); + } + }; + }, [detachEndedListeners, unpublishStream]); + + return { + state, + isSharing: state === 'active', + isBusy: state !== 'idle' && state !== 'active', + error, + start, + stop, + clearError, + }; +} diff --git a/apps/mobile/src/hooks/useWebRTCHost.test.ts b/apps/mobile/src/hooks/useWebRTCHost.test.ts index 3183dbc1..63ee4032 100644 --- a/apps/mobile/src/hooks/useWebRTCHost.test.ts +++ b/apps/mobile/src/hooks/useWebRTCHost.test.ts @@ -1,8 +1,9 @@ import { describe, it, expect, vi, beforeEach } from 'vitest'; -import { renderHook, act } from '@testing-library/react'; +import { renderHook, act, waitFor } from '@testing-library/react'; import { useWebRTCHost } from './useWebRTCHost'; import { createEventSource } from '../lib/event-source'; import { getStoredAuth } from '../lib/secure-storage'; +import type { MediaStream, MediaStreamTrack } from 'react-native-webrtc'; vi.mock('../config', () => ({ API_BASE_URL: 'https://pairux.com', @@ -44,7 +45,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -62,7 +62,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -82,7 +81,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -102,7 +100,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -119,7 +116,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -141,7 +137,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -162,7 +157,6 @@ describe('useWebRTCHost', () => { useWebRTCHost({ sessionId: 'session-1', hostId: 'host-1', - localStream: null, }) ); @@ -173,4 +167,137 @@ describe('useWebRTCHost', () => { unmount(); expect(mockClose).toHaveBeenCalled(); }); + + it('removes only screen-share senders when unpublishing', async () => { + const { result } = renderHook(() => + useWebRTCHost({ + sessionId: 'session-1', + hostId: 'host-1', + }) + ); + + await act(async () => { + await result.current.startHosting(); + }); + const presenceJoinListener = mockAddEventListener.mock.calls.find( + ([eventName]) => eventName === 'presence-join' + )?.[1] as ((event: { data: string }) => void) | undefined; + act(() => { + presenceJoinListener?.({ + data: JSON.stringify({ presences: [{ user_id: 'viewer-1', role: 'viewer' }] }), + }); + }); + await waitFor(() => expect(result.current.viewerCount).toBe(1)); + + const viewer = result.current.viewers.get('viewer-1'); + expect(viewer).toBeDefined(); + const screenTrack = { + id: 'screen', + kind: 'video', + stop: vi.fn(), + } as unknown as MediaStreamTrack; + const screenStream = { getTracks: () => [screenTrack] } as unknown as MediaStream; + + await act(async () => { + await result.current.publishStream(screenStream); + }); + const screenSender = viewer?.peerConnection + .getSenders() + .find((sender) => sender.track === screenTrack); + expect(screenSender).toBeDefined(); + + await act(async () => { + await result.current.unpublishStream(); + }); + + expect(viewer?.peerConnection.removeTrack).toHaveBeenCalledWith(screenSender); + expect(viewer?.peerConnection.getSenders()).not.toContain(screenSender); + expect( + viewer?.peerConnection.getSenders().some((sender) => sender.track?.kind === 'audio') + ).toBe(true); + }); + + it('does not add a duplicate sender when the same stream is published twice', async () => { + const { result } = renderHook(() => + useWebRTCHost({ + sessionId: 'session-1', + hostId: 'host-1', + }) + ); + + await act(async () => { + await result.current.startHosting(); + }); + const presenceJoinListener = mockAddEventListener.mock.calls.find( + ([eventName]) => eventName === 'presence-join' + )?.[1] as ((event: { data: string }) => void) | undefined; + act(() => { + presenceJoinListener?.({ + data: JSON.stringify({ presences: [{ user_id: 'viewer-1', role: 'viewer' }] }), + }); + }); + await waitFor(() => expect(result.current.viewerCount).toBe(1)); + + const viewer = result.current.viewers.get('viewer-1'); + const screenTrack = { + id: 'screen', + kind: 'video', + stop: vi.fn(), + } as unknown as MediaStreamTrack; + const screenStream = { getTracks: () => [screenTrack] } as unknown as MediaStream; + + await act(async () => { + await result.current.publishStream(screenStream); + await result.current.publishStream(screenStream); + }); + + const screenSenders = viewer?.peerConnection + .getSenders() + .filter((sender) => sender.track === screenTrack); + expect(screenSenders).toHaveLength(1); + }); + + it('rejects publishing and rolls back its sender when signaling fails', async () => { + const { result } = renderHook(() => + useWebRTCHost({ + sessionId: 'session-1', + hostId: 'host-1', + }) + ); + + await act(async () => { + await result.current.startHosting(); + }); + const presenceJoinListener = mockAddEventListener.mock.calls.find( + ([eventName]) => eventName === 'presence-join' + )?.[1] as ((event: { data: string }) => void) | undefined; + act(() => { + presenceJoinListener?.({ + data: JSON.stringify({ presences: [{ user_id: 'viewer-1', role: 'viewer' }] }), + }); + }); + await waitFor(() => expect(result.current.viewerCount).toBe(1)); + + const viewer = result.current.viewers.get('viewer-1'); + const screenTrack = { + id: 'screen', + kind: 'video', + stop: vi.fn(), + } as unknown as MediaStreamTrack; + const screenStream = { getTracks: () => [screenTrack] } as unknown as MediaStream; + vi.mocked(fetch).mockResolvedValue({ + ok: false, + text: async () => 'signaling unavailable', + } as Response); + + await act(async () => { + await expect(result.current.publishStream(screenStream)).rejects.toThrow( + 'Failed to signal viewer viewer-1' + ); + }); + + expect(viewer?.peerConnection.getSenders().some((sender) => sender.track === screenTrack)).toBe( + false + ); + }); }); diff --git a/apps/mobile/src/hooks/useWebRTCHost.ts b/apps/mobile/src/hooks/useWebRTCHost.ts index c8dec330..5ac9e345 100644 --- a/apps/mobile/src/hooks/useWebRTCHost.ts +++ b/apps/mobile/src/hooks/useWebRTCHost.ts @@ -9,7 +9,7 @@ * - No HTMLAudioElement (audio handled by RN WebRTC) */ import { useState, useEffect, useRef, useCallback } from 'react'; -import type { MediaStreamTrack } from 'react-native-webrtc'; +import type { MediaStreamTrack, RTCRtpSender } from 'react-native-webrtc'; import { RTCPeerConnection, RTCIceCandidate, MediaStream, mediaDevices } from 'react-native-webrtc'; import type { ConnectionState, @@ -93,7 +93,6 @@ interface SignalMessage { interface UseWebRTCHostOptions { sessionId: string; hostId: string; - localStream: MediaStream | null; allowControl?: boolean; onViewerJoined?: (viewerId: string) => void; onViewerLeft?: (viewerId: string) => void; @@ -123,7 +122,6 @@ interface UseWebRTCHostReturn { export function useWebRTCHost({ sessionId, hostId, - localStream, onViewerJoined, onViewerLeft, onControlRequest, @@ -142,13 +140,14 @@ export function useWebRTCHost({ const removeViewerRef = useRef<((viewerId: string) => void) | undefined>(undefined); const authTokenRef = useRef(null); const isStartingRef = useRef(false); - const localStreamRef = useRef(localStream); + const localStreamRef = useRef(null); const hostMicStreamRef = useRef(null); + const publishedStreamSendersRef = useRef>>(new Map()); + const publishedStreamVersionRef = useRef(0); const iceServersRef = useRef(DEFAULT_ICE_SERVERS); const pendingCandidatesRef = useRef>(new Map()); - localStreamRef.current = localStream; const onControlRequestRef = useRef(onControlRequest); const onInputReceivedRef = useRef(onInputReceived); onControlRequestRef.current = onControlRequest; @@ -156,7 +155,7 @@ export function useWebRTCHost({ // Send signal via API const sendSignal = useCallback( - async (signal: SignalMessage) => { + async (signal: SignalMessage): Promise => { try { const headers: Record = { 'Content-Type': 'application/json' }; if (authTokenRef.current) { @@ -171,14 +170,43 @@ export function useWebRTCHost({ if (!response.ok) { console.error('[WebRTCHost] Failed to send signal:', await response.text()); + return false; } + return true; } catch (err) { console.error('[WebRTCHost] Error sending signal:', err); + return false; } }, [sessionId] ); + const renegotiateViewer = useCallback( + async (viewer: ViewerConnection): Promise => { + const offer = (await viewer.peerConnection.createOffer({})) as OfferAnswer; + if (offer.sdp) { + offer.sdp = tuneOpusForVoice(offer.sdp); + } + await viewer.peerConnection.setLocalDescription(offer); + + if (!offer.sdp) { + throw new Error(`Failed to create an SDP offer for viewer ${viewer.id}`); + } + + const sent = await sendSignal({ + type: 'offer', + sdp: offer.sdp, + senderId: hostId, + targetId: viewer.id, + timestamp: Date.now(), + }); + if (!sent) { + throw new Error(`Failed to signal viewer ${viewer.id}`); + } + }, + [hostId, sendSignal] + ); + // Report usage stats const reportStats = useCallback(async () => { for (const viewer of viewersRef.current.values()) { @@ -326,12 +354,15 @@ export function useWebRTCHost({ iceServers: iceServersRef.current, iceCandidatePoolSize: 10, }); + const publishedSenders = new Map(); + publishedStreamSendersRef.current.set(viewerId, publishedSenders); // Add local stream tracks const currentStream = localStreamRef.current; if (currentStream) { currentStream.getTracks().forEach((track) => { - pc.addTrack(track, currentStream); + const sender = pc.addTrack(track, currentStream); + publishedSenders.set(track.kind, sender); }); } @@ -453,6 +484,7 @@ export function useWebRTCHost({ console.log('[WebRTCHost] Removing viewer:', viewerId); viewer.peerConnection.close(); viewersRef.current.delete(viewerId); + publishedStreamSendersRef.current.delete(viewerId); pendingCandidatesRef.current.delete(viewerId); setViewers(new Map(viewersRef.current)); onViewerLeft?.(viewerId); @@ -628,6 +660,7 @@ export function useWebRTCHost({ eventSourceRef.current = eventSource; eventSource.addEventListener('connected', (event) => { + if (eventSourceRef.current !== eventSource) return; console.log('[WebRTCHost] SSE connected'); isStartingRef.current = false; setIsHosting(true); @@ -644,6 +677,7 @@ export function useWebRTCHost({ }); eventSource.addEventListener('signal', (event) => { + if (eventSourceRef.current !== eventSource) return; try { const signal = JSON.parse(event.data) as SignalMessage; void handleSignalMessage(signal); @@ -653,6 +687,7 @@ export function useWebRTCHost({ }); eventSource.addEventListener('presence-join', (event) => { + if (eventSourceRef.current !== eventSource) return; try { const { presences } = JSON.parse(event.data) as { presences: { user_id: string; role: string }[]; @@ -668,6 +703,7 @@ export function useWebRTCHost({ }); eventSource.addEventListener('presence-leave', (event) => { + if (eventSourceRef.current !== eventSource) return; try { const { presences } = JSON.parse(event.data) as { presences: { user_id: string }[]; @@ -681,6 +717,7 @@ export function useWebRTCHost({ }); eventSource.addEventListener('error', () => { + if (eventSourceRef.current !== eventSource) return; console.error('[WebRTCHost] SSE error'); isStartingRef.current = false; setError('Connection to server lost. Reconnecting...'); @@ -711,6 +748,13 @@ export function useWebRTCHost({ viewer.peerConnection.close(); }); viewersRef.current.clear(); + publishedStreamSendersRef.current.clear(); + publishedStreamVersionRef.current += 1; + const publishedStream = localStreamRef.current; + localStreamRef.current = null; + publishedStream?.getTracks().forEach((track) => { + track.stop(); + }); setViewers(new Map()); setIsHosting(false); @@ -724,77 +768,121 @@ export function useWebRTCHost({ setHasMic(false); }, []); + const removePublishedSenders = useCallback((viewer: ViewerConnection): boolean => { + const publishedSenders = publishedStreamSendersRef.current.get(viewer.id); + if (!publishedSenders || publishedSenders.size === 0) { + return false; + } + + for (const sender of publishedSenders.values()) { + viewer.peerConnection.removeTrack(sender); + } + publishedSenders.clear(); + return true; + }, []); + // Publish screen share stream const publishStream = useCallback( async (stream: MediaStream) => { localStreamRef.current = stream; + const publishVersion = ++publishedStreamVersionRef.current; - for (const viewer of viewersRef.current.values()) { - if (viewer.connectionState !== 'connected' && viewer.connectionState !== 'connecting') - continue; + try { + for (const viewer of viewersRef.current.values()) { + if ( + publishedStreamVersionRef.current !== publishVersion || + localStreamRef.current !== stream + ) { + return; + } + if (viewer.connectionState !== 'connected' && viewer.connectionState !== 'connecting') { + continue; + } - try { - stream.getTracks().forEach((track) => { - viewer.peerConnection.addTrack(track, stream); - }); + let publishedSenders = publishedStreamSendersRef.current.get(viewer.id); + if (!publishedSenders) { + publishedSenders = new Map(); + publishedStreamSendersRef.current.set(viewer.id, publishedSenders); + } - const offer = (await viewer.peerConnection.createOffer({})) as OfferAnswer; - // In-band FEC turns a lost packet into a duller syllable, not a gap. - if (offer.sdp) offer.sdp = tuneOpusForVoice(offer.sdp); - await viewer.peerConnection.setLocalDescription(offer); + const desiredTracks = new Map(stream.getTracks().map((track) => [track.kind, track])); + let negotiationNeeded = false; - if (offer.sdp) { - await sendSignal({ - type: 'offer', - sdp: offer.sdp, - senderId: hostId, - targetId: viewer.id, - timestamp: Date.now(), - }); + for (const [kind, sender] of publishedSenders) { + const desiredTrack = desiredTracks.get(kind); + if (!desiredTrack || sender.track !== desiredTrack) { + viewer.peerConnection.removeTrack(sender); + publishedSenders.delete(kind); + negotiationNeeded = true; + } + } + + for (const track of desiredTracks.values()) { + if (!publishedSenders.has(track.kind)) { + const sender = viewer.peerConnection.addTrack(track, stream); + publishedSenders.set(track.kind, sender); + negotiationNeeded = true; + } + } + + if (negotiationNeeded) { + await renegotiateViewer(viewer); + } + } + } catch (publishError) { + if ( + publishedStreamVersionRef.current === publishVersion && + localStreamRef.current === stream + ) { + localStreamRef.current = null; + publishedStreamVersionRef.current += 1; + + for (const viewer of viewersRef.current.values()) { + if (!removePublishedSenders(viewer)) continue; + try { + await renegotiateViewer(viewer); + } catch (rollbackError) { + console.error( + `[WebRTCHost] Failed to roll back stream for ${viewer.id}:`, + rollbackError + ); + } } - } catch (err) { - console.error(`[WebRTCHost] Failed to publish stream to ${viewer.id}:`, err); } + throw publishError; } }, - [hostId, sendSignal] + [removePublishedSenders, renegotiateViewer] ); // Unpublish stream const unpublishStream = useCallback(async () => { localStreamRef.current = null; + const unpublishVersion = ++publishedStreamVersionRef.current; + const failures: string[] = []; for (const viewer of viewersRef.current.values()) { - if (viewer.connectionState !== 'connected' && viewer.connectionState !== 'connecting') + if (publishedStreamVersionRef.current !== unpublishVersion) { + return; + } + if (viewer.connectionState !== 'connected' && viewer.connectionState !== 'connecting') { continue; + } try { - const senders = viewer.peerConnection.getSenders(); - for (const sender of senders) { - if (sender.track?.kind === 'video') { - viewer.peerConnection.removeTrack(sender); - } + if (removePublishedSenders(viewer)) { + await renegotiateViewer(viewer); } - - const offer = (await viewer.peerConnection.createOffer({})) as OfferAnswer; - // In-band FEC turns a lost packet into a duller syllable, not a gap. - if (offer.sdp) offer.sdp = tuneOpusForVoice(offer.sdp); - await viewer.peerConnection.setLocalDescription(offer); - - if (offer.sdp) { - await sendSignal({ - type: 'offer', - sdp: offer.sdp, - senderId: hostId, - targetId: viewer.id, - timestamp: Date.now(), - }); - } - } catch (err) { - console.error(`[WebRTCHost] Failed to unpublish stream from ${viewer.id}:`, err); + } catch (unpublishError) { + failures.push(viewer.id); + console.error(`[WebRTCHost] Failed to unpublish stream from ${viewer.id}:`, unpublishError); } } - }, [hostId, sendSignal]); + + if (failures.length > 0) { + throw new Error(`Failed to unpublish screen share from ${String(failures.length)} viewer(s)`); + } + }, [removePublishedSenders, renegotiateViewer]); // Cleanup on unmount useEffect(() => { @@ -803,23 +891,6 @@ export function useWebRTCHost({ }; }, [stopHosting]); - // Update stream when it changes - useEffect(() => { - if (!localStream || !isHosting) return; - - viewersRef.current.forEach((viewer) => { - const senders = viewer.peerConnection.getSenders(); - localStream.getTracks().forEach((track) => { - const existingSender = senders.find((s) => s.track?.kind === track.kind); - if (existingSender) { - void existingSender.replaceTrack(track); - } else { - viewer.peerConnection.addTrack(track, localStream); - } - }); - }); - }, [localStream, isHosting]); - // Grant control const grantControl = useCallback( (viewerId: string) => { @@ -885,6 +956,7 @@ export function useWebRTCHost({ viewer.peerConnection.close(); viewersRef.current.delete(viewerId); + publishedStreamSendersRef.current.delete(viewerId); setViewers(new Map(viewersRef.current)); onViewerLeft?.(viewerId); }, diff --git a/apps/mobile/src/lib/host-session.test.ts b/apps/mobile/src/lib/host-session.test.ts new file mode 100644 index 00000000..872b92eb --- /dev/null +++ b/apps/mobile/src/lib/host-session.test.ts @@ -0,0 +1,41 @@ +import { describe, expect, it, vi } from 'vitest'; +import { endHostSession } from './host-session'; + +describe('endHostSession', () => { + it('keeps the local host active when the server cannot end the session', async () => { + const stopScreenShare = vi.fn().mockResolvedValue(undefined); + const endSession = vi.fn().mockResolvedValue({ error: 'Network error' }); + const stopHosting = vi.fn(); + const goBack = vi.fn(); + + await expect( + endHostSession({ stopScreenShare, endSession, stopHosting, goBack }) + ).rejects.toThrow('Network error'); + + expect(stopScreenShare).toHaveBeenCalledTimes(1); + expect(endSession).toHaveBeenCalledTimes(1); + expect(stopHosting).not.toHaveBeenCalled(); + expect(goBack).not.toHaveBeenCalled(); + }); + + it('tears down and leaves only after the server ends the session', async () => { + const calls: string[] = []; + const stopScreenShare = vi.fn(async () => { + calls.push('screen'); + }); + const endSession = vi.fn(async () => { + calls.push('server'); + return { data: { id: 'session-1' } }; + }); + const stopHosting = vi.fn(() => { + calls.push('host'); + }); + const goBack = vi.fn(() => { + calls.push('back'); + }); + + await endHostSession({ stopScreenShare, endSession, stopHosting, goBack }); + + expect(calls).toEqual(['screen', 'server', 'host', 'back']); + }); +}); diff --git a/apps/mobile/src/lib/host-session.ts b/apps/mobile/src/lib/host-session.ts new file mode 100644 index 00000000..e5b61197 --- /dev/null +++ b/apps/mobile/src/lib/host-session.ts @@ -0,0 +1,26 @@ +import type { ApiResponse } from './api'; + +interface EndHostSessionOptions { + stopScreenShare: () => Promise; + endSession: () => Promise>; + stopHosting: () => void; + goBack: () => void; +} + +/** Ends the server session before tearing down the local host connection. */ +export async function endHostSession({ + stopScreenShare, + endSession, + stopHosting, + goBack, +}: EndHostSessionOptions): Promise { + await stopScreenShare(); + + const result = await endSession(); + if (result.error) { + throw new Error(result.error); + } + + stopHosting(); + goBack(); +} diff --git a/apps/mobile/src/test/setup.ts b/apps/mobile/src/test/setup.ts index 332f3876..ff7b49a7 100644 --- a/apps/mobile/src/test/setup.ts +++ b/apps/mobile/src/test/setup.ts @@ -80,6 +80,11 @@ class MockRTCPeerConnection { _transceivers: unknown[] = []; _remoteStreams = new Map(); _pendingTrackEvents: unknown[] = []; + _senders: { + track: { kind: string } | null; + getParameters: () => { encodings: Record[] }; + setParameters: ReturnType; + }[] = []; createOffer = vi.fn().mockResolvedValue({ type: 'offer', sdp: 'mock-sdp' }); createAnswer = vi.fn().mockResolvedValue({ type: 'answer', sdp: 'mock-answer-sdp' }); @@ -90,13 +95,19 @@ class MockRTCPeerConnection { this.remoteDescription = desc; }); addIceCandidate = vi.fn().mockResolvedValue(undefined); - addTrack = vi.fn().mockReturnValue({ - getParameters: () => ({ encodings: [{}] }), - setParameters: vi.fn().mockResolvedValue(undefined), - track: null, + addTrack = vi.fn((track: { kind: string }) => { + const sender = { + getParameters: () => ({ encodings: [{}] }), + setParameters: vi.fn().mockResolvedValue(undefined), + track, + }; + this._senders.push(sender); + return sender; }); - removeTrack = vi.fn(); - getSenders = vi.fn().mockReturnValue([]); + removeTrack = vi.fn((sender: (typeof this._senders)[number]) => { + this._senders = this._senders.filter((candidate) => candidate !== sender); + }); + getSenders = vi.fn(() => [...this._senders]); getReceivers = vi.fn().mockReturnValue([]); getTransceivers = vi.fn().mockReturnValue([]); getStats = vi.fn().mockResolvedValue(new Map());