refactor to a single peer, refuse via 409 on handshake (#1233)

This commit is contained in:
Henry Ruhs
2026-09-11 11:49:11 +02:00
committed by GitHub
parent cdb8513d11
commit 40c554b59d
8 changed files with 80 additions and 75 deletions
+11 -8
View File
@@ -1,9 +1,9 @@
from starlette.requests import Request from starlette.requests import Request
from starlette.responses import Response 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 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.api_helper import get_sec_websocket_protocol
from facefusion.apis.session_helper import extract_access_token from facefusion.apis.session_helper import extract_access_token
from facefusion.apis.stream_manager import destroy_stream, process_image, process_video 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 'Location': request.url_for('delete_stream').path
} }
content_type = request.headers.get('content-type')
access_token = extract_access_token(request.scope) access_token = extract_access_token(request.scope)
session_id = session_manager.find_session_id(access_token) session_id = session_manager.find_session_id(access_token)
session_context.set_session_id(session_id) session_context.set_session_id(session_id)
if session_id and content_type == 'application/sdp': if session_id:
sdp_offer = await request.body() if not rtc_store.has_peer(session_id):
sdp_answer = process_video(session_id, sdp_offer.decode()) sdp_offer = await request.body()
sdp_answer = process_video(session_id, sdp_offer.decode())
if sdp_answer: if sdp_answer:
return Response(sdp_answer, status_code = HTTP_201_CREATED, media_type = 'application/sdp', headers = headers) 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) return Response(status_code = HTTP_404_NOT_FOUND)
+6 -6
View File
@@ -91,8 +91,7 @@ def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpA
codec = audio_codec codec = audio_codec
) )
rtc_store.init_peers(session_id) rtc_store.set_peer(session_id, rtc_peer)
rtc_store.get_peers(session_id).append(rtc_peer)
content_store.clear() content_store.clear()
threading.Thread(target = run_peer_loop, args = (session_id, rtc_peer), daemon = True).start() 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_receiver_thread.join()
video_encoder_thread.join() video_encoder_thread.join()
video_executor.shutdown(wait = True) video_executor.shutdown(wait = True)
rtc_store.delete_peers(session_id)
rtc_store.delete_peer(session_id)
def destroy_stream(session_id : SessionId) -> bool: def destroy_stream(session_id : SessionId) -> bool:
if rtc_store.has_peers(session_id): if rtc_store.has_peer(session_id):
rtc_store.delete_peers(session_id) rtc_store.delete_peer(session_id)
return not rtc_store.has_peers(session_id) return not rtc_store.has_peer(session_id)
return False return False
+5 -9
View File
@@ -1,5 +1,5 @@
import ctypes import ctypes
from typing import List, Optional from typing import Optional
from facefusion.libraries import datachannel as datachannel_module 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 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 return None
def delete_peers(rtc_peers : List[RtcPeer]) -> None: def delete_peer(rtc_peer : RtcPeer) -> None:
datachannel_library = datachannel_module.create_static_library() datachannel_library = datachannel_module.create_static_library()
peer_connection = rtc_peer.get('peer_connection')
for rtc_peer in rtc_peers: if peer_connection:
peer_connection = rtc_peer.get('peer_connection') datachannel_library.rtcDeletePeerConnection(peer_connection)
if peer_connection:
datachannel_library.rtcDeletePeerConnection(peer_connection)
return None
def add_audio_track(peer_connection : PeerConnection, media_direction : MediaDirection, audio_codec : AudioCodec, payload_type : int) -> RtcAudioTrack: def add_audio_track(peer_connection : PeerConnection, media_direction : MediaDirection, audio_codec : AudioCodec, payload_type : int) -> RtcAudioTrack:
+10 -16
View File
@@ -1,4 +1,4 @@
from typing import List from typing import Optional
from facefusion import rtc from facefusion import rtc
from facefusion.types import RtcPeer, RtcStore, SessionId from facefusion.types import RtcPeer, RtcStore, SessionId
@@ -6,27 +6,21 @@ from facefusion.types import RtcPeer, RtcStore, SessionId
RTC_STORE : RtcStore = {} RTC_STORE : RtcStore = {}
def init_peers(session_id : SessionId) -> None: def has_peer(session_id : SessionId) -> bool:
RTC_STORE[session_id] = [] return session_id in RTC_STORE
def has_peers(session_id : SessionId) -> bool: def get_peer(session_id : SessionId) -> Optional[RtcPeer]:
return bool(RTC_STORE.get(session_id))
def get_peers(session_id : SessionId) -> List[RtcPeer]:
return RTC_STORE.get(session_id) 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: if session_id in RTC_STORE:
rtc_peers = get_peers(session_id) rtc.delete_peer(RTC_STORE.pop(session_id))
if rtc_peers:
rtc.delete_peers(rtc_peers)
del RTC_STORE[session_id]
return None
def clear() -> None: def clear() -> None:
+1 -1
View File
@@ -368,7 +368,7 @@ RtcPeer = TypedDict('RtcPeer',
'sender_bitrate': ctypes.c_uint, 'sender_bitrate': ctypes.c_uint,
'receiver_bitrate': ctypes.c_uint 'receiver_bitrate': ctypes.c_uint
}) })
RtcStore : TypeAlias = Dict[SessionId, List[RtcPeer]] RtcStore : TypeAlias = Dict[SessionId, RtcPeer]
ContentSet = TypedDict('ContentSet', ContentSet = TypedDict('ContentSet',
{ {
+19 -2
View File
@@ -156,7 +156,16 @@ def test_delete_stream_video(test_client : TestClient) -> None:
'Content-Type': 'application/sdp' '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 = 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 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
+14 -16
View File
@@ -114,18 +114,18 @@ def test_process_video(video_codec : VideoCodec, session_id : str) -> None:
assert 'a=recvonly' in sdp_answer assert 'a=recvonly' in sdp_answer
assert 'a=sendonly' in sdp_answer assert 'a=sendonly' in sdp_answer
for peer in rtc_store.get_peers(session_id): rtc_peer = rtc_store.get_peer(session_id)
sender_bitrate = peer.get('sender_bitrate') sender_bitrate = rtc_peer.get('sender_bitrate')
receiver_bitrate = peer.get('receiver_bitrate') receiver_bitrate = rtc_peer.get('receiver_bitrate')
assert sender_bitrate.value == 0 assert sender_bitrate.value == 0
assert receiver_bitrate.value == 8000 assert receiver_bitrate.value == 8000
rtc.handle_sender_bitrate(0, 8000000, ctypes.addressof(sender_bitrate)) rtc.handle_sender_bitrate(0, 8000000, ctypes.addressof(sender_bitrate))
assert sender_bitrate.value == 8000 assert sender_bitrate.value == 8000
rtc.adapt_receiver_bitrate(peer, 4000) rtc.adapt_receiver_bitrate(rtc_peer, 4000)
assert receiver_bitrate.value == 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') ]) @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) 'receiver_bitrate': ctypes.c_uint(0)
} }
rtc_store.init_peers(session_id) rtc_store.set_peer(session_id, rtc_peer)
rtc_store.get_peers(session_id).append(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.receive_video_frames'):
with patch('facefusion.apis.stream_manager.run_video_encode_loop'): 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.start()
thread.join(timeout = 5.0) 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: def test_destroy_stream() -> None:
@@ -177,10 +176,9 @@ def test_destroy_stream() -> None:
} }
session_id = 'test-destroy-stream' session_id = 'test-destroy-stream'
rtc_store.init_peers(session_id) rtc_store.set_peer(session_id, rtc_peer)
rtc_store.get_peers(session_id).append(rtc_peer)
assert destroy_stream(session_id) is True 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 assert destroy_stream(session_id) is False
+14 -17
View File
@@ -1,11 +1,10 @@
import ctypes import ctypes
from typing import List
import pytest import pytest
from facefusion import state_manager from facefusion import state_manager
from facefusion.libraries import datachannel as datachannel_module, opus as opus_module, vpx as vpx_module 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 from facefusion.types import RtcPeer, VideoCodec
@@ -116,25 +115,23 @@ def test_send_audio() -> None:
datachannel_library.rtcDeletePeerConnection(peer_connection) datachannel_library.rtcDeletePeerConnection(peer_connection)
def test_delete_peers() -> None: def test_delete_peer() -> None:
datachannel_library = datachannel_module.create_static_library() datachannel_library = datachannel_module.create_static_library()
peer_connection = create_peer_connection() peer_connection = create_peer_connection()
rtc_peers : List[RtcPeer] =\ rtc_peer : RtcPeer =\
[ {
'peer_connection': peer_connection,
'video':
{ {
'peer_connection': peer_connection, 'sender_track': 0,
'video': 'receiver_track': 0,
{ 'codec': 'vp8'
'sender_track': 0, },
'receiver_track': 0, 'sender_bitrate': ctypes.c_uint(0),
'codec': 'vp8' 'receiver_bitrate': ctypes.c_uint(0)
}, }
'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 assert datachannel_library.rtcDeletePeerConnection(peer_connection) == -1