convert rtc store (#1242)

This commit is contained in:
Henry Ruhs
2026-09-12 16:55:18 +02:00
committed by GitHub
parent 1052b6f532
commit 97a4a406c0
10 changed files with 70 additions and 57 deletions
+3 -2
View File
@@ -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)
+4 -8
View File
@@ -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)
+10 -10
View File
@@ -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
+25 -14
View File
@@ -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)
-1
View File
@@ -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',
{
+4 -3
View File
@@ -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:
+9 -5
View File
@@ -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 =
{
+1 -1
View File
@@ -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:
+13 -12
View File
@@ -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
+1 -1
View File
@@ -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: