handle garbage and empty buffer for websocket (#1232)

This commit is contained in:
Henry Ruhs
2026-09-11 11:06:04 +02:00
committed by GitHub
parent 743ebb8622
commit 41ee571fc0
3 changed files with 25 additions and 10 deletions
+3 -8
View File
@@ -5,8 +5,6 @@ from concurrent.futures import Future, ThreadPoolExecutor
from queue import Queue from queue import Queue
from typing import Optional, Tuple from typing import Optional, Tuple
import cv2
import numpy
from starlette.websockets import WebSocket from starlette.websockets import WebSocket
from facefusion import content_store, rtc, rtc_store, state_manager, streamer 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.content_analyser import analyse_frame
from facefusion.libraries import datachannel as datachannel_module 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, 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: async def process_image(websocket : WebSocket) -> None:
capture_vision_frame = await anext(receive_vision_frames(websocket), None) async for capture_vision_frame in receive_vision_frames(websocket):
if is_vision_frame(capture_vision_frame):
source_vision_frames = read_static_images(state_manager.get_item('source_paths')) 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) 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() websocket_event = await websocket.receive()
while websocket_event.get('type') == 'websocket.receive': while websocket_event.get('type') == 'websocket.receive':
vision_buffer = websocket_event.get('bytes') or bytes() vision_frame = from_buffer(websocket_event.get('bytes'))
vision_frame = cv2.imdecode(numpy.frombuffer(vision_buffer, numpy.uint8), cv2.IMREAD_COLOR)
if is_vision_frame(vision_frame): if is_vision_frame(vision_frame):
yield vision_frame yield vision_frame
+6
View File
@@ -316,6 +316,12 @@ def create_empty_vision_frame() -> VisionFrame:
return numpy.zeros((1, 1, 3)).astype(numpy.uint8) 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: def to_buffer(vision_frame : VisionFrame) -> Buffer:
is_success, vision_buffer = cv2.imencode('.jpg', vision_frame) is_success, vision_buffer = cv2.imencode('.jpg', vision_frame)
+16 -2
View File
@@ -42,6 +42,9 @@ async def test_process_image() -> None:
{ {
'type': 'websocket.receive', 'type': 'websocket.receive',
'bytes': image_buffer 'bytes': image_buffer
},
{
'type': 'websocket.disconnect'
} }
] ]
@@ -68,14 +71,25 @@ async def test_receive_vision_frames() -> None:
'type': 'websocket.receive', 'type': 'websocket.receive',
'bytes': 'invalid'.encode() 'bytes': 'invalid'.encode()
}, },
{
'type': 'websocket.receive',
'bytes': bytes()
},
{
'type': 'websocket.receive',
'bytes': image_buffer
},
{ {
'type': 'websocket.disconnect' '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') ]) @pytest.mark.parametrize('video_codec, session_id', [ ('av1', 'test-process-video-av1'), ('vp8', 'test-process-video-vp8') ])