use get_event_loop() and move event handlers to BaseTransport
This commit is contained in:
@@ -46,7 +46,7 @@ class PipelineRunner:
|
|||||||
return self._running
|
return self._running
|
||||||
|
|
||||||
def _setup_sigint(self):
|
def _setup_sigint(self):
|
||||||
loop = asyncio.get_running_loop()
|
loop = asyncio.get_event_loop()
|
||||||
loop.add_signal_handler(
|
loop.add_signal_handler(
|
||||||
signal.SIGINT,
|
signal.SIGINT,
|
||||||
lambda *args: asyncio.create_task(self._sigint_handler())
|
lambda *args: asyncio.create_task(self._sigint_handler())
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ class FrameProcessor:
|
|||||||
self.name = f"{self.__class__.__name__}#{obj_count(self)}"
|
self.name = f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
self._prev: "FrameProcessor" | None = None
|
self._prev: "FrameProcessor" | None = None
|
||||||
self._next: "FrameProcessor" | None = None
|
self._next: "FrameProcessor" | None = None
|
||||||
self._loop: AbstractEventLoop = asyncio.get_running_loop()
|
self._loop: AbstractEventLoop = asyncio.get_event_loop()
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -4,7 +4,12 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import inspect
|
||||||
|
import types
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
from pydantic import ConfigDict
|
from pydantic import ConfigDict
|
||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
@@ -12,6 +17,8 @@ from pydantic.main import BaseModel
|
|||||||
from pipecat.processors.frame_processor import FrameProcessor
|
from pipecat.processors.frame_processor import FrameProcessor
|
||||||
from pipecat.vad.vad_analyzer import VADAnalyzer
|
from pipecat.vad.vad_analyzer import VADAnalyzer
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
|
||||||
class TransportParams(BaseModel):
|
class TransportParams(BaseModel):
|
||||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||||
@@ -36,6 +43,10 @@ class TransportParams(BaseModel):
|
|||||||
|
|
||||||
class BaseTransport(ABC):
|
class BaseTransport(ABC):
|
||||||
|
|
||||||
|
def __init__(self, loop: asyncio.AbstractEventLoop):
|
||||||
|
self._loop = loop
|
||||||
|
self._event_handlers: dict = {}
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> FrameProcessor:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
@@ -43,3 +54,42 @@ class BaseTransport(ABC):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def output(self) -> FrameProcessor:
|
def output(self) -> FrameProcessor:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def event_handler(self, event_name: str):
|
||||||
|
def decorator(handler):
|
||||||
|
self._add_event_handler(event_name, handler)
|
||||||
|
return handler
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
def _register_event_handler(self, event_name: str):
|
||||||
|
methods = inspect.getmembers(self, predicate=inspect.ismethod)
|
||||||
|
if event_name not in [method[0] for method in methods]:
|
||||||
|
raise Exception(f"Event handler {event_name} not found")
|
||||||
|
|
||||||
|
self._event_handlers[event_name] = [getattr(self, event_name)]
|
||||||
|
|
||||||
|
patch_method = types.MethodType(partial(self._patch_method, event_name), self)
|
||||||
|
setattr(self, event_name, patch_method)
|
||||||
|
|
||||||
|
def _add_event_handler(self, event_name: str, handler):
|
||||||
|
if event_name not in self._event_handlers:
|
||||||
|
raise Exception(f"Event handler {event_name} not registered")
|
||||||
|
self._event_handlers[event_name].append(types.MethodType(handler, self))
|
||||||
|
|
||||||
|
def _patch_method(self, event_name, *args, **kwargs):
|
||||||
|
try:
|
||||||
|
for handler in self._event_handlers[event_name]:
|
||||||
|
if inspect.iscoroutinefunction(handler):
|
||||||
|
# Beware, if handler() calls another event handler it
|
||||||
|
# will deadlock. You shouldn't do that anyways.
|
||||||
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
|
handler(*args[1:], **kwargs), self._loop)
|
||||||
|
|
||||||
|
# wait for the coroutine to finish. This will also
|
||||||
|
# raise any exceptions raised by the coroutine.
|
||||||
|
future.result()
|
||||||
|
else:
|
||||||
|
handler(*args[1:], **kwargs)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Exception in event handler {event_name}: {e}")
|
||||||
|
raise e
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ class WebsocketServerParams(TransportParams):
|
|||||||
|
|
||||||
|
|
||||||
class WebsocketServerCallbacks(BaseModel):
|
class WebsocketServerCallbacks(BaseModel):
|
||||||
on_connection: Callable[[websockets.WebSocketServerProtocol], Awaitable[None]]
|
on_connection: Callable[[websockets.WebSocketServerProtocol], None]
|
||||||
|
|
||||||
|
|
||||||
class WebsocketServerInputTransport(FrameProcessor):
|
class WebsocketServerInputTransport(FrameProcessor):
|
||||||
@@ -87,7 +87,7 @@ class WebsocketServerInputTransport(FrameProcessor):
|
|||||||
self._websocket = websocket
|
self._websocket = websocket
|
||||||
|
|
||||||
# Notify
|
# Notify
|
||||||
await self._callbacks.on_connection(websocket)
|
self._callbacks.on_connection(websocket)
|
||||||
|
|
||||||
# Handle incoming messages
|
# Handle incoming messages
|
||||||
async for message in websocket:
|
async for message in websocket:
|
||||||
@@ -148,9 +148,11 @@ class WebsocketServerOutputTransport(FrameProcessor):
|
|||||||
self._audio_buffer = bytes()
|
self._audio_buffer = bytes()
|
||||||
self._in_tts_audio = False
|
self._in_tts_audio = False
|
||||||
|
|
||||||
async def set_client_connection(self, websocket: websockets.WebSocketServerProtocol):
|
def set_client_connection(self, websocket: websockets.WebSocketServerProtocol):
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
await self._websocket.close()
|
loop = self.get_event_loop()
|
||||||
|
future = asyncio.run_coroutine_threadsafe(self._websocket.close(), loop)
|
||||||
|
future.result()
|
||||||
logger.warning("Only one client allowed, using new connection")
|
logger.warning("Only one client allowed, using new connection")
|
||||||
self._websocket = websocket
|
self._websocket = websocket
|
||||||
|
|
||||||
@@ -196,9 +198,13 @@ class WebsocketServerOutputTransport(FrameProcessor):
|
|||||||
|
|
||||||
class WebsocketServerTransport(BaseTransport):
|
class WebsocketServerTransport(BaseTransport):
|
||||||
|
|
||||||
def __init__(self, host: str = "localhost", port: int = 8765,
|
def __init__(
|
||||||
params: WebsocketServerParams = WebsocketServerParams()):
|
self,
|
||||||
super().__init__()
|
host: str = "localhost",
|
||||||
|
port: int = 8765,
|
||||||
|
params: WebsocketServerParams = WebsocketServerParams(),
|
||||||
|
loop: asyncio.AbstractEventLoop = asyncio.get_event_loop()):
|
||||||
|
super().__init__(loop)
|
||||||
self._host = host
|
self._host = host
|
||||||
self._port = port
|
self._port = port
|
||||||
self._params = params
|
self._params = params
|
||||||
@@ -210,6 +216,10 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
self._output: WebsocketServerOutputTransport | None = None
|
self._output: WebsocketServerOutputTransport | None = None
|
||||||
self._websocket: websockets.WebSocketServerProtocol | None = None
|
self._websocket: websockets.WebSocketServerProtocol | None = None
|
||||||
|
|
||||||
|
# Register supported handlers. The user will only be able to register
|
||||||
|
# these handlers.
|
||||||
|
self._register_event_handler("on_client_connected")
|
||||||
|
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> FrameProcessor:
|
||||||
if not self._input:
|
if not self._input:
|
||||||
self._input = WebsocketServerInputTransport(
|
self._input = WebsocketServerInputTransport(
|
||||||
@@ -221,8 +231,19 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
self._output = WebsocketServerOutputTransport(self._params)
|
self._output = WebsocketServerOutputTransport(self._params)
|
||||||
return self._output
|
return self._output
|
||||||
|
|
||||||
async def _on_connection(self, websocket):
|
def _on_connection(self, websocket):
|
||||||
if self._output:
|
if self._output:
|
||||||
await self._output.set_client_connection(websocket)
|
print("000 AAAAAAAAAAAAAAAAAAA")
|
||||||
|
self._output.set_client_connection(websocket)
|
||||||
|
print("111 AAAAAAAAAAAAAAAAAAA")
|
||||||
|
self.on_client_connected(websocket)
|
||||||
|
print("222 AAAAAAAAAAAAAAAAAAA")
|
||||||
else:
|
else:
|
||||||
logger.error("A WebsocketServerTransport output is missing in the pipeline")
|
logger.error("A WebsocketServerTransport output is missing in the pipeline")
|
||||||
|
|
||||||
|
#
|
||||||
|
# Decorators (event handlers)
|
||||||
|
#
|
||||||
|
|
||||||
|
def on_client_connected(self, client):
|
||||||
|
pass
|
||||||
|
|||||||
@@ -6,15 +6,12 @@
|
|||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import asyncio
|
import asyncio
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
|
||||||
import inspect
|
|
||||||
import queue
|
import queue
|
||||||
import time
|
import time
|
||||||
import types
|
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from functools import partial
|
|
||||||
from typing import Any, Callable, Mapping
|
from typing import Any, Callable, Mapping
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
from daily import (
|
from daily import (
|
||||||
CallClient,
|
CallClient,
|
||||||
@@ -139,7 +136,8 @@ class DailyTransportClient(EventHandler):
|
|||||||
token: str | None,
|
token: str | None,
|
||||||
bot_name: str,
|
bot_name: str,
|
||||||
params: DailyParams,
|
params: DailyParams,
|
||||||
callbacks: DailyCallbacks):
|
callbacks: DailyCallbacks,
|
||||||
|
loop: asyncio.AbstractEventLoop):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
if not self._daily_initialized:
|
if not self._daily_initialized:
|
||||||
@@ -151,6 +149,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
self._bot_name: str = bot_name
|
self._bot_name: str = bot_name
|
||||||
self._params: DailyParams = params
|
self._params: DailyParams = params
|
||||||
self._callbacks = callbacks
|
self._callbacks = callbacks
|
||||||
|
self._loop = loop
|
||||||
|
|
||||||
self._participant_id: str = ""
|
self._participant_id: str = ""
|
||||||
self._video_renderers = {}
|
self._video_renderers = {}
|
||||||
@@ -212,8 +211,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
|
|
||||||
self._joining = True
|
self._joining = True
|
||||||
|
|
||||||
loop = asyncio.get_running_loop()
|
await self._loop.run_in_executor(self._executor, self._join)
|
||||||
await loop.run_in_executor(self._executor, self._join)
|
|
||||||
|
|
||||||
def _join(self):
|
def _join(self):
|
||||||
logger.info(f"Joining {self._room_url}")
|
logger.info(f"Joining {self._room_url}")
|
||||||
@@ -304,8 +302,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
self._joined = False
|
self._joined = False
|
||||||
self._leaving = True
|
self._leaving = True
|
||||||
|
|
||||||
loop = asyncio.get_running_loop()
|
await self._loop.run_in_executor(self._executor, self._leave)
|
||||||
await loop.run_in_executor(self._executor, self._leave)
|
|
||||||
|
|
||||||
def _leave(self):
|
def _leave(self):
|
||||||
logger.info(f"Leaving {self._room_url}")
|
logger.info(f"Leaving {self._room_url}")
|
||||||
@@ -335,8 +332,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
self._callbacks.on_error(error_msg)
|
self._callbacks.on_error(error_msg)
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
loop = asyncio.get_running_loop()
|
await self._loop.run_in_executor(self._executor, self._cleanup)
|
||||||
await loop.run_in_executor(self._executor, self._cleanup)
|
|
||||||
|
|
||||||
def _cleanup(self):
|
def _cleanup(self):
|
||||||
if self._client:
|
if self._client:
|
||||||
@@ -485,8 +481,7 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
# This will set _running=True
|
# This will set _running=True
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
# Create camera in thread (runs if _running is true).
|
# Create camera in thread (runs if _running is true).
|
||||||
loop = asyncio.get_running_loop()
|
self._camera_in_thread = self._loop.run_in_executor(
|
||||||
self._camera_in_thread = loop.run_in_executor(
|
|
||||||
self._in_executor, self._camera_in_thread_handler)
|
self._in_executor, self._camera_in_thread_handler)
|
||||||
|
|
||||||
async def stop(self):
|
async def stop(self):
|
||||||
@@ -642,7 +637,15 @@ class DailyOutputTransport(BaseOutputTransport):
|
|||||||
|
|
||||||
class DailyTransport(BaseTransport):
|
class DailyTransport(BaseTransport):
|
||||||
|
|
||||||
def __init__(self, room_url: str, token: str | None, bot_name: str, params: DailyParams):
|
def __init__(
|
||||||
|
self,
|
||||||
|
room_url: str,
|
||||||
|
token: str | None,
|
||||||
|
bot_name: str,
|
||||||
|
params: DailyParams,
|
||||||
|
loop: asyncio.AbstractEventLoop = asyncio.get_event_loop()):
|
||||||
|
super().__init__(loop)
|
||||||
|
|
||||||
callbacks = DailyCallbacks(
|
callbacks = DailyCallbacks(
|
||||||
on_joined=self._on_joined,
|
on_joined=self._on_joined,
|
||||||
on_left=self._on_left,
|
on_left=self._on_left,
|
||||||
@@ -660,12 +663,9 @@ class DailyTransport(BaseTransport):
|
|||||||
)
|
)
|
||||||
self._params = params
|
self._params = params
|
||||||
|
|
||||||
self._client = DailyTransportClient(room_url, token, bot_name, params, callbacks)
|
self._client = DailyTransportClient(room_url, token, bot_name, params, callbacks, loop)
|
||||||
self._input: DailyInputTransport | None = None
|
self._input: DailyInputTransport | None = None
|
||||||
self._output: DailyOutputTransport | None = None
|
self._output: DailyOutputTransport | None = None
|
||||||
self._loop = asyncio.get_running_loop()
|
|
||||||
|
|
||||||
self._event_handlers: dict = {}
|
|
||||||
|
|
||||||
# Register supported handlers. The user will only be able to register
|
# Register supported handlers. The user will only be able to register
|
||||||
# these handlers.
|
# these handlers.
|
||||||
@@ -868,45 +868,3 @@ class DailyTransport(BaseTransport):
|
|||||||
|
|
||||||
def on_participant_left(self, participant, reason):
|
def on_participant_left(self, participant, reason):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def event_handler(self, event_name: str):
|
|
||||||
def decorator(handler):
|
|
||||||
self._add_event_handler(event_name, handler)
|
|
||||||
return handler
|
|
||||||
return decorator
|
|
||||||
|
|
||||||
def _register_event_handler(self, event_name: str):
|
|
||||||
methods = inspect.getmembers(self, predicate=inspect.ismethod)
|
|
||||||
if event_name not in [method[0] for method in methods]:
|
|
||||||
raise Exception(f"Event handler {event_name} not found")
|
|
||||||
|
|
||||||
self._event_handlers[event_name] = [getattr(self, event_name)]
|
|
||||||
|
|
||||||
patch_method = types.MethodType(partial(self._patch_method, event_name), self)
|
|
||||||
setattr(self, event_name, patch_method)
|
|
||||||
|
|
||||||
def _add_event_handler(self, event_name: str, handler):
|
|
||||||
if event_name not in self._event_handlers:
|
|
||||||
raise Exception(f"Event handler {event_name} not registered")
|
|
||||||
self._event_handlers[event_name].append(types.MethodType(handler, self))
|
|
||||||
|
|
||||||
def _patch_method(self, event_name, *args, **kwargs):
|
|
||||||
try:
|
|
||||||
for handler in self._event_handlers[event_name]:
|
|
||||||
if inspect.iscoroutinefunction(handler):
|
|
||||||
# Beware, if handler() calls another event handler it
|
|
||||||
# will deadlock. You shouldn't do that anyways.
|
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
|
||||||
handler(*args[1:], **kwargs), self._loop)
|
|
||||||
|
|
||||||
# wait for the coroutine to finish. This will also
|
|
||||||
# raise any exceptions raised by the coroutine.
|
|
||||||
future.result()
|
|
||||||
else:
|
|
||||||
handler(*args[1:], **kwargs)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Exception in event handler {event_name}: {e}")
|
|
||||||
raise e
|
|
||||||
|
|
||||||
# def start_recording(self):
|
|
||||||
# self.client.start_recording()
|
|
||||||
|
|||||||
Reference in New Issue
Block a user