mirror of
https://github.com/facefusion/facefusion.git
synced 2026-09-15 20:15:28 +02:00
convert rtc store (#1242)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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',
|
||||
{
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 =
|
||||
{
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user