diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 49e97c14..eed09610 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -4,7 +4,7 @@ from starlette.requests import Request from starlette.responses import JSONResponse from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZED, HTTP_404_NOT_FOUND -from facefusion import content_store, face_store, inference_manager, process_manager, session_context, session_manager, state_manager, translator, video_manager +from facefusion import content_store, face_store, inference_manager, process_manager, rtc_store, session_context, session_manager, state_manager, translator, video_manager from facefusion.apis import asset_store from facefusion.apis.session_helper import validate_api_key from facefusion.apis.stream_manager import destroy_stream @@ -26,6 +26,7 @@ async def create_session(request : Request) -> JSONResponse: inference_manager.init() video_manager.init() process_manager.init() + rtc_store.init() return JSONResponse( { @@ -82,7 +83,7 @@ async def destroy_session(request : Request) -> JSONResponse: 'message': translator.get('directory_not_removed', 'facefusion.apis') }, status_code = HTTP_404_NOT_FOUND) - destroy_stream(session_id) + destroy_stream() asset_store.delete_assets(session_id) session_manager.clear_session(session_id) diff --git a/facefusion/apis/endpoints/stream.py b/facefusion/apis/endpoints/stream.py index 6fbc239b..4a81a3a5 100644 --- a/facefusion/apis/endpoints/stream.py +++ b/facefusion/apis/endpoints/stream.py @@ -3,7 +3,7 @@ from starlette.responses import Response from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_404_NOT_FOUND, HTTP_409_CONFLICT from starlette.websockets import WebSocket, WebSocketState -from facefusion import rtc_store, session_context +from facefusion import rtc_store from facefusion.apis.api_helper import get_sec_websocket_protocol from facefusion.apis.stream_manager import destroy_stream, process_image, process_video @@ -23,11 +23,9 @@ async def post_stream(request : Request) -> Response: { 'Location': request.url_for('delete_stream').path } - session_id = session_context.get_session_id() - - if not rtc_store.has_peer(session_id): + if not rtc_store.has_peer(): sdp_offer = await request.body() - sdp_answer = process_video(session_id, sdp_offer.decode()) + sdp_answer = process_video(sdp_offer.decode()) if sdp_answer: return Response(sdp_answer, status_code = HTTP_201_CREATED, media_type = 'application/sdp', headers = headers) @@ -39,9 +37,7 @@ async def post_stream(request : Request) -> Response: async def delete_stream(request : Request) -> Response: - session_id = session_context.get_session_id() - - if destroy_stream(session_id): + if destroy_stream(): return Response(status_code = HTTP_200_OK) return Response(status_code = HTTP_404_NOT_FOUND) diff --git a/facefusion/apis/stream_manager.py b/facefusion/apis/stream_manager.py index f94571c7..20e77ee4 100644 --- a/facefusion/apis/stream_manager.py +++ b/facefusion/apis/stream_manager.py @@ -13,7 +13,7 @@ from facefusion.apis.stream_audio import receive_audio_frames, run_audio_encode_ from facefusion.apis.stream_video import receive_video_frames, run_video_encode_loop from facefusion.content_analyser import analyse_frame from facefusion.libraries import datachannel as datachannel_module -from facefusion.types import AudioCodec, AudioFrame, BufferPack, PeerConnection, RtcPeer, RtcPeerAudio, SdpAnswer, SdpOffer, SessionId, Time, VideoCodec, VisionFrame +from facefusion.types import AudioCodec, AudioFrame, BufferPack, PeerConnection, RtcPeer, RtcPeerAudio, SdpAnswer, SdpOffer, Time, VideoCodec, VisionFrame from facefusion.vision import from_buffer, is_vision_frame, obscure_frame, read_static_images, to_buffer @@ -41,7 +41,7 @@ async def receive_vision_frames(websocket : WebSocket) -> AsyncIterator[VisionFr websocket_event = await websocket.receive() -def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpAnswer]: +def process_video(sdp_offer : SdpOffer) -> Optional[SdpAnswer]: video_codec : VideoCodec = 'vp8' if rtc.get_payload_type(sdp_offer, 'vp9'): @@ -93,11 +93,11 @@ def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpA ) content_store.clear() - rtc_store.set_peer(session_id, rtc_peer) + rtc_store.set_peer(rtc_peer) threading.Thread( target = copy_context().run, - args = (run_peer_loop, session_id, rtc_peer), + args = (run_peer_loop, rtc_peer), daemon = True ).start() @@ -108,7 +108,7 @@ def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpA return None -def run_peer_loop(session_id : SessionId, rtc_peer : RtcPeer) -> None: +def run_peer_loop(rtc_peer : RtcPeer) -> None: execution_thread_count = state_manager.get_item('execution_thread_count') video_queue : Queue[Tuple[Time, Future[BufferPack]]] = Queue(maxsize = execution_thread_count) audio_queue : Queue[Tuple[Time, AudioFrame]] = Queue(maxsize = execution_thread_count * 10) @@ -147,12 +147,12 @@ def run_peer_loop(session_id : SessionId, rtc_peer : RtcPeer) -> None: video_encoder_thread.join() video_executor.shutdown(wait = True) - rtc_store.delete_peer(session_id) + rtc_store.delete_peer() -def destroy_stream(session_id : SessionId) -> bool: - if rtc_store.has_peer(session_id): - rtc_store.delete_peer(session_id) - return not rtc_store.has_peer(session_id) +def destroy_stream() -> bool: + if rtc_store.has_peer(): + rtc_store.delete_peer() + return not rtc_store.has_peer() return False diff --git a/facefusion/rtc_store.py b/facefusion/rtc_store.py index 49bff11b..abf79eeb 100644 --- a/facefusion/rtc_store.py +++ b/facefusion/rtc_store.py @@ -1,27 +1,38 @@ from typing import Optional -from facefusion import rtc -from facefusion.types import RtcPeer, RtcStore, SessionId +from facefusion import rtc, store_creator +from facefusion.session_manager import resolve_owner_id +from facefusion.types import RtcPeer, Store -RTC_STORE : RtcStore = {} +RTC_STORE : Store = store_creator.create_store(None) -def has_peer(session_id : SessionId) -> bool: - return session_id in RTC_STORE +def init() -> None: + owner_id = resolve_owner_id() + store_creator.init_content(RTC_STORE, owner_id) -def get_peer(session_id : SessionId) -> Optional[RtcPeer]: - return RTC_STORE.get(session_id) +def has_peer() -> bool: + owner_id = resolve_owner_id() + + return bool(store_creator.get_content(RTC_STORE, owner_id)) -def set_peer(session_id : SessionId, rtc_peer : RtcPeer) -> None: - RTC_STORE[session_id] = rtc_peer +def get_peer() -> Optional[RtcPeer]: + owner_id = resolve_owner_id() + + return store_creator.get_content(RTC_STORE, owner_id) -def delete_peer(session_id : SessionId) -> None: - if session_id in RTC_STORE: - rtc.delete_peer(RTC_STORE.pop(session_id)) +def set_peer(rtc_peer : RtcPeer) -> None: + owner_id = resolve_owner_id() + store_creator.set_content(RTC_STORE, owner_id, rtc_peer) -def clear() -> None: - RTC_STORE.clear() +def delete_peer() -> None: + owner_id = resolve_owner_id() + rtc_peer = store_creator.get_content(RTC_STORE, owner_id) + + if rtc_peer: + rtc.delete_peer(rtc_peer) + store_creator.init_content(RTC_STORE, owner_id) diff --git a/facefusion/types.py b/facefusion/types.py index 10c673a9..64a6edf8 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -376,7 +376,6 @@ RtcPeer = TypedDict('RtcPeer', 'sender_bitrate': ctypes.c_uint, 'receiver_bitrate': ctypes.c_uint }) -RtcStore : TypeAlias = Dict[SessionId, RtcPeer] ContentSet = TypedDict('ContentSet', { diff --git a/tests/test_api_session.py b/tests/test_api_session.py index 5bdacfd4..1b421406 100644 --- a/tests/test_api_session.py +++ b/tests/test_api_session.py @@ -8,7 +8,7 @@ from unittest.mock import patch import pytest from starlette.testclient import TestClient -from facefusion import metadata, process_manager, rtc, rtc_store, session_manager, state_manager +from facefusion import metadata, process_manager, rtc, rtc_store, session_context, session_manager, state_manager from facefusion.apis import asset_store from facefusion.apis.core import create_api from facefusion.download import conditional_download @@ -255,7 +255,8 @@ def test_destroy_session(test_client : TestClient) -> None: 'sender_bitrate': ctypes.c_uint(0), 'receiver_bitrate': ctypes.c_uint(0) } - rtc_store.set_peer(session_id, rtc_peer) + session_context.set_session_id(session_id) + rtc_store.set_peer(rtc_peer) delete_session_response = test_client.delete('/session', headers = { @@ -264,7 +265,7 @@ def test_destroy_session(test_client : TestClient) -> None: assert session_manager.find_session_id(access_token) is None assert asset_store.get_assets(session_id) is None - assert rtc_store.has_peer(session_id) is False + assert rtc_store.has_peer() is False assert delete_session_response.status_code == 200 for asset_path in asset_paths: diff --git a/tests/test_api_stream.py b/tests/test_api_stream.py index a12de75f..5222992d 100644 --- a/tests/test_api_stream.py +++ b/tests/test_api_stream.py @@ -5,7 +5,7 @@ from unittest.mock import patch import pytest from starlette.testclient import TestClient -from facefusion import metadata, rtc, rtc_store, session_manager, state_manager +from facefusion import metadata, rtc, rtc_store, session_context, session_manager, state_manager from facefusion.apis import asset_store from facefusion.apis.core import create_api, pre_check from facefusion.download import conditional_download @@ -36,7 +36,7 @@ def before_all() -> None: def before_each() -> None: session_manager.SESSIONS.clear() asset_store.clear() - rtc_store.clear() + rtc_store.delete_peer() @pytest.fixture(scope = 'module') @@ -157,7 +157,9 @@ def test_delete_stream_video(test_client : TestClient) -> None: 'Content-Type': 'application/sdp' }) - assert rtc_store.has_peer(session_id) is True + session_context.set_session_id(session_id) + + assert rtc_store.has_peer() is True post_response = test_client.post('/stream', content = sdp_offer, headers = { @@ -166,7 +168,9 @@ def test_delete_stream_video(test_client : TestClient) -> None: }) assert post_response.status_code == 409 - assert rtc_store.has_peer(session_id) is True + session_context.set_session_id(session_id) + + assert rtc_store.has_peer() is True delete_response = test_client.delete('/stream', headers = { @@ -174,7 +178,7 @@ def test_delete_stream_video(test_client : TestClient) -> None: }) assert delete_response.status_code == 200 - assert rtc_store.has_peer(session_id) is False + assert rtc_store.has_peer() is False post_response = test_client.post('/stream', content = 'invalid', headers = { diff --git a/tests/test_api_stream_audio.py b/tests/test_api_stream_audio.py index 2a97dd10..d8bcec30 100644 --- a/tests/test_api_stream_audio.py +++ b/tests/test_api_stream_audio.py @@ -35,7 +35,7 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: - rtc_store.clear() + rtc_store.delete_peer() def dispatch_frame(buffer : Buffer, track : int, frame_handler : FrameHandler) -> threading.Event: diff --git a/tests/test_api_stream_manager.py b/tests/test_api_stream_manager.py index c49b459a..36cb9daa 100644 --- a/tests/test_api_stream_manager.py +++ b/tests/test_api_stream_manager.py @@ -36,7 +36,7 @@ def before_all() -> None: def before_each() -> Iterator[None]: local_id = resolve_local_id() - rtc_store.clear() + rtc_store.delete_peer() yield @@ -117,14 +117,15 @@ def test_process_video(video_codec : VideoCodec, session_id : str) -> None: datachannel_module.create_static_library().rtcDeletePeerConnection(peer_connection) with patch('facefusion.apis.stream_manager.threading.Thread'): - sdp_answer = process_video(session_id, sdp_offer) + set_session_id(session_id) + sdp_answer = process_video(sdp_offer) assert sdp_answer assert 'm=video' in sdp_answer assert 'a=recvonly' in sdp_answer assert 'a=sendonly' in sdp_answer - rtc_peer = rtc_store.get_peer(session_id) + rtc_peer = rtc_store.get_peer() sender_bitrate = rtc_peer.get('sender_bitrate') receiver_bitrate = rtc_peer.get('receiver_bitrate') @@ -156,22 +157,22 @@ def test_run_peer_loop(video_codec : VideoCodec, payload_type : int, session_id 'receiver_bitrate': ctypes.c_uint(0) } - rtc_store.set_peer(session_id, rtc_peer) store_creator.set_content(state_manager.STATE_SET, session_id, state_manager.get_state()) set_session_id(session_id) + rtc_store.set_peer(rtc_peer) - assert rtc_store.has_peer(session_id) is True + assert rtc_store.has_peer() is True with patch('facefusion.thread_helper.ThreadPoolExecutor') as thread_pool_executor_mock: with patch('facefusion.apis.stream_manager.receive_video_frames'): with patch('facefusion.apis.stream_manager.run_video_encode_loop'): - thread = threading.Thread(target = copy_context().run, args = (run_peer_loop, session_id, rtc_peer), daemon = True) + thread = threading.Thread(target = copy_context().run, args = (run_peer_loop, rtc_peer), daemon = True) thread.start() thread.join(timeout = 5.0) thread_pool_executor_mock.assert_called_once_with(max_workers = 8, initializer = set_session_id, initargs = tuple([ session_id ])) - assert rtc_store.has_peer(session_id) is False + assert rtc_store.has_peer() is False def test_destroy_stream() -> None: @@ -189,11 +190,11 @@ def test_destroy_stream() -> None: 'sender_bitrate': ctypes.c_uint(0), 'receiver_bitrate': ctypes.c_uint(0) } - session_id = 'test-destroy-stream' - rtc_store.set_peer(session_id, rtc_peer) + set_session_id('test-destroy-stream') + rtc_store.set_peer(rtc_peer) - assert destroy_stream(session_id) is True - assert rtc_store.get_peer(session_id) is None + assert destroy_stream() is True + assert rtc_store.get_peer() is None - assert destroy_stream(session_id) is False + assert destroy_stream() is False diff --git a/tests/test_api_stream_video.py b/tests/test_api_stream_video.py index 5d0db57d..9adcf292 100644 --- a/tests/test_api_stream_video.py +++ b/tests/test_api_stream_video.py @@ -41,7 +41,7 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: - rtc_store.clear() + rtc_store.delete_peer() def dispatch_frame(buffer : Buffer, track : int, frame_handler : FrameHandler) -> threading.Event: