mirror of
https://github.com/facefusion/facefusion.git
synced 2026-09-15 12:05:27 +02:00
handle garbage and empty buffer for websocket (#1232)
This commit is contained in:
@@ -5,8 +5,6 @@ from concurrent.futures import Future, ThreadPoolExecutor
|
||||
from queue import Queue
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import cv2
|
||||
import numpy
|
||||
from starlette.websockets import WebSocket
|
||||
|
||||
from facefusion import content_store, rtc, rtc_store, state_manager, streamer
|
||||
@@ -15,13 +13,11 @@ from facefusion.apis.stream_video import receive_video_frames, run_video_encode_
|
||||
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.vision import is_vision_frame, obscure_frame, read_static_images, to_buffer
|
||||
from facefusion.vision import from_buffer, is_vision_frame, obscure_frame, read_static_images, to_buffer
|
||||
|
||||
|
||||
async def process_image(websocket : WebSocket) -> None:
|
||||
capture_vision_frame = await anext(receive_vision_frames(websocket), None)
|
||||
|
||||
if is_vision_frame(capture_vision_frame):
|
||||
async for capture_vision_frame in receive_vision_frames(websocket):
|
||||
source_vision_frames = read_static_images(state_manager.get_item('source_paths'))
|
||||
output_vision_frame = streamer.process_stream_frame(source_vision_frames, capture_vision_frame)
|
||||
|
||||
@@ -36,8 +32,7 @@ async def receive_vision_frames(websocket : WebSocket) -> AsyncIterator[VisionFr
|
||||
websocket_event = await websocket.receive()
|
||||
|
||||
while websocket_event.get('type') == 'websocket.receive':
|
||||
vision_buffer = websocket_event.get('bytes') or bytes()
|
||||
vision_frame = cv2.imdecode(numpy.frombuffer(vision_buffer, numpy.uint8), cv2.IMREAD_COLOR)
|
||||
vision_frame = from_buffer(websocket_event.get('bytes'))
|
||||
|
||||
if is_vision_frame(vision_frame):
|
||||
yield vision_frame
|
||||
|
||||
@@ -316,6 +316,12 @@ def create_empty_vision_frame() -> VisionFrame:
|
||||
return numpy.zeros((1, 1, 3)).astype(numpy.uint8)
|
||||
|
||||
|
||||
def from_buffer(vision_buffer : Buffer) -> Optional[VisionFrame]:
|
||||
if vision_buffer:
|
||||
return cv2.imdecode(numpy.frombuffer(vision_buffer, numpy.uint8), cv2.IMREAD_COLOR)
|
||||
return None
|
||||
|
||||
|
||||
def to_buffer(vision_frame : VisionFrame) -> Buffer:
|
||||
is_success, vision_buffer = cv2.imencode('.jpg', vision_frame)
|
||||
|
||||
|
||||
@@ -42,6 +42,9 @@ async def test_process_image() -> None:
|
||||
{
|
||||
'type': 'websocket.receive',
|
||||
'bytes': image_buffer
|
||||
},
|
||||
{
|
||||
'type': 'websocket.disconnect'
|
||||
}
|
||||
]
|
||||
|
||||
@@ -68,14 +71,25 @@ async def test_receive_vision_frames() -> None:
|
||||
'type': 'websocket.receive',
|
||||
'bytes': 'invalid'.encode()
|
||||
},
|
||||
{
|
||||
'type': 'websocket.receive',
|
||||
'bytes': bytes()
|
||||
},
|
||||
{
|
||||
'type': 'websocket.receive',
|
||||
'bytes': image_buffer
|
||||
},
|
||||
{
|
||||
'type': 'websocket.disconnect'
|
||||
}
|
||||
]
|
||||
vision_frames = []
|
||||
|
||||
vision_frames = receive_vision_frames(websocket_mock)
|
||||
async for vision_frame in receive_vision_frames(websocket_mock):
|
||||
vision_frames.append(vision_frame)
|
||||
|
||||
assert create_hash((await anext(vision_frames)).tobytes()) == '5ed32ca0'
|
||||
assert len(vision_frames) == 2
|
||||
assert create_hash(vision_frames[0].tobytes()) == '5ed32ca0'
|
||||
|
||||
|
||||
@pytest.mark.parametrize('video_codec, session_id', [ ('av1', 'test-process-video-av1'), ('vp8', 'test-process-video-vp8') ])
|
||||
|
||||
Reference in New Issue
Block a user