From 40c554b59d4585ce98fb930a921781cc6db17bf8 Mon Sep 17 00:00:00 2001 From: Henry Ruhs Date: Fri, 11 Sep 2026 11:49:11 +0200 Subject: [PATCH] refactor to a single peer, refuse via 409 on handshake (#1233) --- facefusion/apis/endpoints/stream.py | 19 ++++++++++-------- facefusion/apis/stream_manager.py | 12 +++++------ facefusion/rtc.py | 14 +++++-------- facefusion/rtc_store.py | 26 ++++++++++-------------- facefusion/types.py | 2 +- tests/test_api_stream.py | 21 +++++++++++++++++-- tests/test_api_stream_manager.py | 30 +++++++++++++--------------- tests/test_rtc.py | 31 +++++++++++++---------------- 8 files changed, 80 insertions(+), 75 deletions(-) diff --git a/facefusion/apis/endpoints/stream.py b/facefusion/apis/endpoints/stream.py index 6f84d001..b3bfef14 100644 --- a/facefusion/apis/endpoints/stream.py +++ b/facefusion/apis/endpoints/stream.py @@ -1,9 +1,9 @@ from starlette.requests import Request from starlette.responses import Response -from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_404_NOT_FOUND +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 session_context, session_manager +from facefusion import rtc_store, session_context, session_manager from facefusion.apis.api_helper import get_sec_websocket_protocol from facefusion.apis.session_helper import extract_access_token from facefusion.apis.stream_manager import destroy_stream, process_image, process_video @@ -27,17 +27,20 @@ async def post_stream(request : Request) -> Response: { 'Location': request.url_for('delete_stream').path } - content_type = request.headers.get('content-type') access_token = extract_access_token(request.scope) session_id = session_manager.find_session_id(access_token) session_context.set_session_id(session_id) - if session_id and content_type == 'application/sdp': - sdp_offer = await request.body() - sdp_answer = process_video(session_id, sdp_offer.decode()) + if session_id: + if not rtc_store.has_peer(session_id): + sdp_offer = await request.body() + sdp_answer = process_video(session_id, sdp_offer.decode()) - if sdp_answer: - return Response(sdp_answer, status_code = HTTP_201_CREATED, media_type = 'application/sdp', headers = headers) + if sdp_answer: + return Response(sdp_answer, status_code = HTTP_201_CREATED, media_type = 'application/sdp', headers = headers) + + else: + return Response(status_code = HTTP_409_CONFLICT) return Response(status_code = HTTP_404_NOT_FOUND) diff --git a/facefusion/apis/stream_manager.py b/facefusion/apis/stream_manager.py index 14824452..3a2157fc 100644 --- a/facefusion/apis/stream_manager.py +++ b/facefusion/apis/stream_manager.py @@ -91,8 +91,7 @@ def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpA codec = audio_codec ) - rtc_store.init_peers(session_id) - rtc_store.get_peers(session_id).append(rtc_peer) + rtc_store.set_peer(session_id, rtc_peer) content_store.clear() threading.Thread(target = run_peer_loop, args = (session_id, rtc_peer), daemon = True).start() @@ -126,12 +125,13 @@ def run_peer_loop(session_id : SessionId, rtc_peer : RtcPeer) -> None: video_receiver_thread.join() video_encoder_thread.join() video_executor.shutdown(wait = True) - rtc_store.delete_peers(session_id) + + rtc_store.delete_peer(session_id) def destroy_stream(session_id : SessionId) -> bool: - if rtc_store.has_peers(session_id): - rtc_store.delete_peers(session_id) - return not rtc_store.has_peers(session_id) + if rtc_store.has_peer(session_id): + rtc_store.delete_peer(session_id) + return not rtc_store.has_peer(session_id) return False diff --git a/facefusion/rtc.py b/facefusion/rtc.py index 8af13e90..6b45c81c 100644 --- a/facefusion/rtc.py +++ b/facefusion/rtc.py @@ -1,5 +1,5 @@ import ctypes -from typing import List, Optional +from typing import Optional from facefusion.libraries import datachannel as datachannel_module from facefusion.types import AudioCodec, BitRate, Buffer, MediaDirection, PeerConnection, RtcAudioTrack, RtcPeer, RtcTrackInit, RtcVideoTrack, SdpAnswer, SdpOffer, Time, VideoCodec @@ -75,16 +75,12 @@ def send_audio(rtc_peer : RtcPeer, audio_buffer : Buffer, audio_timestamp : int) return None -def delete_peers(rtc_peers : List[RtcPeer]) -> None: +def delete_peer(rtc_peer : RtcPeer) -> None: datachannel_library = datachannel_module.create_static_library() + peer_connection = rtc_peer.get('peer_connection') - for rtc_peer in rtc_peers: - peer_connection = rtc_peer.get('peer_connection') - - if peer_connection: - datachannel_library.rtcDeletePeerConnection(peer_connection) - - return None + if peer_connection: + datachannel_library.rtcDeletePeerConnection(peer_connection) def add_audio_track(peer_connection : PeerConnection, media_direction : MediaDirection, audio_codec : AudioCodec, payload_type : int) -> RtcAudioTrack: diff --git a/facefusion/rtc_store.py b/facefusion/rtc_store.py index 65ec6725..49bff11b 100644 --- a/facefusion/rtc_store.py +++ b/facefusion/rtc_store.py @@ -1,4 +1,4 @@ -from typing import List +from typing import Optional from facefusion import rtc from facefusion.types import RtcPeer, RtcStore, SessionId @@ -6,27 +6,21 @@ from facefusion.types import RtcPeer, RtcStore, SessionId RTC_STORE : RtcStore = {} -def init_peers(session_id : SessionId) -> None: - RTC_STORE[session_id] = [] +def has_peer(session_id : SessionId) -> bool: + return session_id in RTC_STORE -def has_peers(session_id : SessionId) -> bool: - return bool(RTC_STORE.get(session_id)) - - -def get_peers(session_id : SessionId) -> List[RtcPeer]: +def get_peer(session_id : SessionId) -> Optional[RtcPeer]: return RTC_STORE.get(session_id) -def delete_peers(session_id : SessionId) -> None: +def set_peer(session_id : SessionId, rtc_peer : RtcPeer) -> None: + RTC_STORE[session_id] = rtc_peer + + +def delete_peer(session_id : SessionId) -> None: if session_id in RTC_STORE: - rtc_peers = get_peers(session_id) - - if rtc_peers: - rtc.delete_peers(rtc_peers) - del RTC_STORE[session_id] - - return None + rtc.delete_peer(RTC_STORE.pop(session_id)) def clear() -> None: diff --git a/facefusion/types.py b/facefusion/types.py index 9febe356..16fa9de4 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -368,7 +368,7 @@ RtcPeer = TypedDict('RtcPeer', 'sender_bitrate': ctypes.c_uint, 'receiver_bitrate': ctypes.c_uint }) -RtcStore : TypeAlias = Dict[SessionId, List[RtcPeer]] +RtcStore : TypeAlias = Dict[SessionId, RtcPeer] ContentSet = TypedDict('ContentSet', { diff --git a/tests/test_api_stream.py b/tests/test_api_stream.py index d3eb962e..13c9da1f 100644 --- a/tests/test_api_stream.py +++ b/tests/test_api_stream.py @@ -156,7 +156,16 @@ def test_delete_stream_video(test_client : TestClient) -> None: 'Content-Type': 'application/sdp' }) - assert rtc_store.get_peers(session_id) + assert rtc_store.has_peer(session_id) is True + + post_response = test_client.post('/stream', content = sdp_offer, headers = + { + 'Authorization': 'Bearer ' + access_token, + 'Content-Type': 'application/sdp' + }) + + assert post_response.status_code == 409 + assert rtc_store.has_peer(session_id) is True delete_response = test_client.delete('/stream', headers = { @@ -164,4 +173,12 @@ def test_delete_stream_video(test_client : TestClient) -> None: }) assert delete_response.status_code == 200 - assert rtc_store.get_peers(session_id) is None + assert rtc_store.has_peer(session_id) is False + + post_response = test_client.post('/stream', content = 'invalid', headers = + { + 'Authorization': 'Bearer ' + access_token, + 'Content-Type': 'application/sdp' + }) + + assert post_response.status_code == 404 diff --git a/tests/test_api_stream_manager.py b/tests/test_api_stream_manager.py index f1a4a06e..e90be621 100644 --- a/tests/test_api_stream_manager.py +++ b/tests/test_api_stream_manager.py @@ -114,18 +114,18 @@ def test_process_video(video_codec : VideoCodec, session_id : str) -> None: assert 'a=recvonly' in sdp_answer assert 'a=sendonly' in sdp_answer - for peer in rtc_store.get_peers(session_id): - sender_bitrate = peer.get('sender_bitrate') - receiver_bitrate = peer.get('receiver_bitrate') + rtc_peer = rtc_store.get_peer(session_id) + sender_bitrate = rtc_peer.get('sender_bitrate') + receiver_bitrate = rtc_peer.get('receiver_bitrate') - assert sender_bitrate.value == 0 - assert receiver_bitrate.value == 8000 + assert sender_bitrate.value == 0 + assert receiver_bitrate.value == 8000 - rtc.handle_sender_bitrate(0, 8000000, ctypes.addressof(sender_bitrate)) - assert sender_bitrate.value == 8000 + rtc.handle_sender_bitrate(0, 8000000, ctypes.addressof(sender_bitrate)) + assert sender_bitrate.value == 8000 - rtc.adapt_receiver_bitrate(peer, 4000) - assert receiver_bitrate.value == 4000 + rtc.adapt_receiver_bitrate(rtc_peer, 4000) + assert receiver_bitrate.value == 4000 @pytest.mark.parametrize('video_codec, payload_type, session_id', [ ('av1', 35, 'test-run-peer-loop-av1'), ('vp8', 96, 'test-run-peer-loop-vp8') ]) @@ -146,10 +146,9 @@ def test_run_peer_loop(video_codec : VideoCodec, payload_type : int, session_id 'receiver_bitrate': ctypes.c_uint(0) } - rtc_store.init_peers(session_id) - rtc_store.get_peers(session_id).append(rtc_peer) + rtc_store.set_peer(session_id, rtc_peer) - assert rtc_store.has_peers(session_id) is True + assert rtc_store.has_peer(session_id) is True with patch('facefusion.apis.stream_manager.receive_video_frames'): with patch('facefusion.apis.stream_manager.run_video_encode_loop'): @@ -157,7 +156,7 @@ def test_run_peer_loop(video_codec : VideoCodec, payload_type : int, session_id thread.start() thread.join(timeout = 5.0) - assert rtc_store.has_peers(session_id) is False + assert rtc_store.has_peer(session_id) is False def test_destroy_stream() -> None: @@ -177,10 +176,9 @@ def test_destroy_stream() -> None: } session_id = 'test-destroy-stream' - rtc_store.init_peers(session_id) - rtc_store.get_peers(session_id).append(rtc_peer) + rtc_store.set_peer(session_id, rtc_peer) assert destroy_stream(session_id) is True - assert rtc_store.get_peers(session_id) is None + assert rtc_store.get_peer(session_id) is None assert destroy_stream(session_id) is False diff --git a/tests/test_rtc.py b/tests/test_rtc.py index 35e11ad2..40c0c372 100644 --- a/tests/test_rtc.py +++ b/tests/test_rtc.py @@ -1,11 +1,10 @@ import ctypes -from typing import List import pytest from facefusion import state_manager from facefusion.libraries import datachannel as datachannel_module, opus as opus_module, vpx as vpx_module -from facefusion.rtc import adapt_receiver_bitrate, add_audio_track, add_video_track, create_peer_connection, create_sdp_answer, create_sdp_offer, delete_peers, get_payload_type, handle_sender_bitrate, send_audio, send_video, set_remote_description, wire_sender_bitrate +from facefusion.rtc import adapt_receiver_bitrate, add_audio_track, add_video_track, create_peer_connection, create_sdp_answer, create_sdp_offer, delete_peer, get_payload_type, handle_sender_bitrate, send_audio, send_video, set_remote_description, wire_sender_bitrate from facefusion.types import RtcPeer, VideoCodec @@ -116,25 +115,23 @@ def test_send_audio() -> None: datachannel_library.rtcDeletePeerConnection(peer_connection) -def test_delete_peers() -> None: +def test_delete_peer() -> None: datachannel_library = datachannel_module.create_static_library() peer_connection = create_peer_connection() - rtc_peers : List[RtcPeer] =\ - [ + rtc_peer : RtcPeer =\ + { + 'peer_connection': peer_connection, + 'video': { - 'peer_connection': peer_connection, - 'video': - { - 'sender_track': 0, - 'receiver_track': 0, - 'codec': 'vp8' - }, - 'sender_bitrate': ctypes.c_uint(0), - 'receiver_bitrate': ctypes.c_uint(0) - } - ] + 'sender_track': 0, + 'receiver_track': 0, + 'codec': 'vp8' + }, + 'sender_bitrate': ctypes.c_uint(0), + 'receiver_bitrate': ctypes.c_uint(0) + } - delete_peers(rtc_peers) + delete_peer(rtc_peer) assert datachannel_library.rtcDeletePeerConnection(peer_connection) == -1