Merge pull request #230 from pipecat-ai/aleix/processor-names

processor names
This commit is contained in:
Aleix Conchillo Flaqué
2024-06-13 03:16:07 +08:00
committed by GitHub
11 changed files with 64 additions and 52 deletions

View File

@@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased] ## [Unreleased]
### Added
- Allow specifying frame processors' name through a new `name` constructor
argument.
### Changed ### Changed
- `daily_rest.DailyRoomProperties` now allows extra unknown parameters. - `daily_rest.DailyRoomProperties` now allows extra unknown parameters.

View File

@@ -4,11 +4,9 @@
# SPDX-License-Identifier: BSD 2-Clause License # SPDX-License-Identifier: BSD 2-Clause License
# #
from itertools import chain
from typing import Callable, Coroutine, List from typing import Callable, Coroutine, List
from pipecat.frames.frames import Frame, MetricsFrame, StartFrame from pipecat.frames.frames import Frame
from pipecat.pipeline.base_pipeline import BasePipeline from pipecat.pipeline.base_pipeline import BasePipeline
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
@@ -81,9 +79,6 @@ class Pipeline(BasePipeline):
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
if isinstance(frame, StartFrame) and self.metrics_enabled:
await self._send_initial_metrics()
if direction == FrameDirection.DOWNSTREAM: if direction == FrameDirection.DOWNSTREAM:
await self._source.process_frame(frame, FrameDirection.DOWNSTREAM) await self._source.process_frame(frame, FrameDirection.DOWNSTREAM)
elif direction == FrameDirection.UPSTREAM: elif direction == FrameDirection.UPSTREAM:
@@ -98,9 +93,3 @@ class Pipeline(BasePipeline):
for curr in self._processors[1:]: for curr in self._processors[1:]:
prev.link(curr) prev.link(curr)
prev = curr prev = curr
async def _send_initial_metrics(self):
processors = self.processors_with_metrics()
ttfb = dict(zip([p.name for p in processors], [0] * len(processors)))
frame = MetricsFrame(ttfb=ttfb)
await self._source.process_frame(frame, FrameDirection.DOWNSTREAM)

View File

@@ -10,7 +10,8 @@ from typing import AsyncIterable, Iterable
from pydantic import BaseModel from pydantic import BaseModel
from pipecat.frames.frames import CancelFrame, EndFrame, ErrorFrame, Frame, StartFrame, StopTaskFrame from pipecat.frames.frames import CancelFrame, EndFrame, ErrorFrame, Frame, MetricsFrame, StartFrame, StopTaskFrame
from pipecat.pipeline.base_pipeline import BasePipeline
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from pipecat.utils.utils import obj_count, obj_id from pipecat.utils.utils import obj_count, obj_id
@@ -40,7 +41,7 @@ class Source(FrameProcessor):
class PipelineTask: class PipelineTask:
def __init__(self, pipeline: FrameProcessor, params: PipelineParams = PipelineParams()): def __init__(self, pipeline: BasePipeline, params: PipelineParams = PipelineParams()):
self.id: int = obj_id() self.id: int = obj_id()
self.name: str = f"{self.__class__.__name__}#{obj_count(self)}" self.name: str = f"{self.__class__.__name__}#{obj_count(self)}"
@@ -89,12 +90,18 @@ class PipelineTask:
else: else:
raise Exception("Frames must be an iterable or async iterable") raise Exception("Frames must be an iterable or async iterable")
def _initial_metrics_frame(self) -> MetricsFrame:
processors = self._pipeline.processors_with_metrics()
ttfb = dict(zip([p.name for p in processors], [0] * len(processors)))
return MetricsFrame(ttfb=ttfb)
async def _process_down_queue(self): async def _process_down_queue(self):
start_frame = StartFrame( start_frame = StartFrame(
allow_interruptions=self._params.allow_interruptions, allow_interruptions=self._params.allow_interruptions,
enable_metrics=self._params.enable_metrics, enable_metrics=self._params.enable_metrics,
) )
await self._source.process_frame(start_frame, FrameDirection.DOWNSTREAM) await self._source.process_frame(start_frame, FrameDirection.DOWNSTREAM)
await self._source.process_frame(self._initial_metrics_frame(), FrameDirection.DOWNSTREAM)
running = True running = True
should_cleanup = True should_cleanup = True

View File

@@ -22,9 +22,9 @@ class FrameDirection(Enum):
class FrameProcessor: class FrameProcessor:
def __init__(self, loop: asyncio.AbstractEventLoop | None = None): def __init__(self, name: str | None = None, loop: asyncio.AbstractEventLoop | None = None):
self.id: int = obj_id() self.id: int = obj_id()
self.name = f"{self.__class__.__name__}#{obj_count(self)}" self.name = name or 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: asyncio.AbstractEventLoop = loop or asyncio.get_running_loop() self._loop: asyncio.AbstractEventLoop = loop or asyncio.get_running_loop()

View File

@@ -28,8 +28,8 @@ from pipecat.utils.utils import exp_smoothing
class AIService(FrameProcessor): class AIService(FrameProcessor):
def __init__(self): def __init__(self, **kwargs):
super().__init__() super().__init__(**kwargs)
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
pass pass
@@ -61,8 +61,8 @@ class AIService(FrameProcessor):
class LLMService(AIService): class LLMService(AIService):
"""This class is a no-op but serves as a base class for LLM services.""" """This class is a no-op but serves as a base class for LLM services."""
def __init__(self): def __init__(self, **kwargs):
super().__init__() super().__init__(**kwargs)
self._callbacks = {} self._callbacks = {}
self._start_callbacks = {} self._start_callbacks = {}
@@ -91,8 +91,8 @@ class LLMService(AIService):
class TTSService(AIService): class TTSService(AIService):
def __init__(self, aggregate_sentences: bool = True): def __init__(self, aggregate_sentences: bool = True, **kwargs):
super().__init__() super().__init__(**kwargs)
self._aggregate_sentences: bool = aggregate_sentences self._aggregate_sentences: bool = aggregate_sentences
self._current_sentence: str = "" self._current_sentence: str = ""
@@ -146,8 +146,9 @@ class STTService(AIService):
max_silence_secs: float = 0.3, max_silence_secs: float = 0.3,
max_buffer_secs: float = 1.5, max_buffer_secs: float = 1.5,
sample_rate: int = 16000, sample_rate: int = 16000,
num_channels: int = 1): num_channels: int = 1,
super().__init__() **kwargs):
super().__init__(**kwargs)
self._min_volume = min_volume self._min_volume = min_volume
self._max_silence_secs = max_silence_secs self._max_silence_secs = max_silence_secs
self._max_buffer_secs = max_buffer_secs self._max_buffer_secs = max_buffer_secs
@@ -216,8 +217,8 @@ class STTService(AIService):
class ImageGenService(AIService): class ImageGenService(AIService):
def __init__(self): def __init__(self, **kwargs):
super().__init__() super().__init__(**kwargs)
# Renders the image. Returns an Image object. # Renders the image. Returns an Image object.
@abstractmethod @abstractmethod
@@ -237,8 +238,8 @@ class ImageGenService(AIService):
class VisionService(AIService): class VisionService(AIService):
"""VisionService is a base class for vision services.""" """VisionService is a base class for vision services."""
def __init__(self): def __init__(self, **kwargs):
super().__init__() super().__init__(**kwargs)
self._describe_text = None self._describe_text = None
@abstractmethod @abstractmethod

View File

@@ -41,7 +41,6 @@ from pipecat.services.ai_services import (
try: try:
from openai import AsyncOpenAI, AsyncStream, BadRequestError from openai import AsyncOpenAI, AsyncStream, BadRequestError
from openai.types.chat import ( from openai.types.chat import (
ChatCompletion,
ChatCompletionChunk, ChatCompletionChunk,
ChatCompletionFunctionMessageParam, ChatCompletionFunctionMessageParam,
ChatCompletionMessageParam, ChatCompletionMessageParam,
@@ -68,8 +67,8 @@ class BaseOpenAILLMService(LLMService):
calls from the LLM. calls from the LLM.
""" """
def __init__(self, model: str, api_key=None, base_url=None): def __init__(self, model: str, api_key=None, base_url=None, **kwargs):
super().__init__() super().__init__(**kwargs)
self._model: str = model self._model: str = model
self._client = self.create_client(api_key=api_key, base_url=base_url) self._client = self.create_client(api_key=api_key, base_url=base_url)

View File

@@ -28,8 +28,8 @@ from loguru import logger
class BaseInputTransport(FrameProcessor): class BaseInputTransport(FrameProcessor):
def __init__(self, params: TransportParams): def __init__(self, params: TransportParams, **kwargs):
super().__init__() super().__init__(**kwargs)
self._params = params self._params = params

View File

@@ -37,8 +37,8 @@ from loguru import logger
class BaseOutputTransport(FrameProcessor): class BaseOutputTransport(FrameProcessor):
def __init__(self, params: TransportParams): def __init__(self, params: TransportParams, **kwargs):
super().__init__() super().__init__(**kwargs)
self._params = params self._params = params
@@ -135,6 +135,9 @@ class BaseOutputTransport(FrameProcessor):
elif isinstance(frame, StartInterruptionFrame) or isinstance(frame, StopInterruptionFrame): elif isinstance(frame, StartInterruptionFrame) or isinstance(frame, StopInterruptionFrame):
await self._handle_interruptions(frame) await self._handle_interruptions(frame)
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
elif isinstance(frame, MetricsFrame):
self.send_metrics(frame)
await self.push_frame(frame, direction)
elif isinstance(frame, SystemFrame): elif isinstance(frame, SystemFrame):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
elif isinstance(frame, AudioRawFrame): elif isinstance(frame, AudioRawFrame):
@@ -182,8 +185,6 @@ class BaseOutputTransport(FrameProcessor):
self._set_camera_images(frame.images) self._set_camera_images(frame.images)
elif isinstance(frame, TransportMessageFrame): elif isinstance(frame, TransportMessageFrame):
self.send_message(frame) self.send_message(frame)
elif isinstance(frame, MetricsFrame):
self.send_metrics(frame)
else: else:
future = asyncio.run_coroutine_threadsafe( future = asyncio.run_coroutine_threadsafe(
self._internal_push_frame(frame), self.get_event_loop()) self._internal_push_frame(frame), self.get_event_loop())

View File

@@ -41,7 +41,12 @@ class TransportParams(BaseModel):
class BaseTransport(ABC): class BaseTransport(ABC):
def __init__(self, loop: asyncio.AbstractEventLoop | None): def __init__(self,
input_name: str | None = None,
output_name: str | None = None,
loop: asyncio.AbstractEventLoop | None = None):
self._input_name = input_name
self._output_name = output_name
self._loop = loop or asyncio.get_running_loop() self._loop = loop or asyncio.get_running_loop()
self._event_handlers: dict = {} self._event_handlers: dict = {}

View File

@@ -42,8 +42,9 @@ class WebsocketServerInputTransport(BaseInputTransport):
host: str, host: str,
port: int, port: int,
params: WebsocketServerParams, params: WebsocketServerParams,
callbacks: WebsocketServerCallbacks): callbacks: WebsocketServerCallbacks,
super().__init__(params) **kwargs):
super().__init__(params, **kwargs)
self._host = host self._host = host
self._port = port self._port = port
@@ -98,8 +99,8 @@ class WebsocketServerInputTransport(BaseInputTransport):
class WebsocketServerOutputTransport(BaseOutputTransport): class WebsocketServerOutputTransport(BaseOutputTransport):
def __init__(self, params: WebsocketServerParams): def __init__(self, params: WebsocketServerParams, **kwargs):
super().__init__(params) super().__init__(params, **kwargs)
self._params = params self._params = params
@@ -153,8 +154,10 @@ class WebsocketServerTransport(BaseTransport):
host: str = "localhost", host: str = "localhost",
port: int = 8765, port: int = 8765,
params: WebsocketServerParams = WebsocketServerParams(), params: WebsocketServerParams = WebsocketServerParams(),
input_name: str | None = None,
output_name: str | None = None,
loop: asyncio.AbstractEventLoop | None = None): loop: asyncio.AbstractEventLoop | None = None):
super().__init__(loop) super().__init__(input_name=input_name, output_name=output_name, loop=loop)
self._host = host self._host = host
self._port = port self._port = port
self._params = params self._params = params
@@ -175,12 +178,12 @@ class WebsocketServerTransport(BaseTransport):
def input(self) -> FrameProcessor: def input(self) -> FrameProcessor:
if not self._input: if not self._input:
self._input = WebsocketServerInputTransport( self._input = WebsocketServerInputTransport(
self._host, self._port, self._params, self._callbacks) self._host, self._port, self._params, self._callbacks, name=self._input_name)
return self._input return self._input
def output(self) -> FrameProcessor: def output(self) -> FrameProcessor:
if not self._output: if not self._output:
self._output = WebsocketServerOutputTransport(self._params) self._output = WebsocketServerOutputTransport(self._params, name=self._output_name)
return self._output return self._output
async def _on_client_connected(self, websocket): async def _on_client_connected(self, websocket):

View File

@@ -473,8 +473,8 @@ class DailyTransportClient(EventHandler):
class DailyInputTransport(BaseInputTransport): class DailyInputTransport(BaseInputTransport):
def __init__(self, client: DailyTransportClient, params: DailyParams): def __init__(self, client: DailyTransportClient, params: DailyParams, **kwargs):
super().__init__(params) super().__init__(params, **kwargs)
self._client = client self._client = client
@@ -609,8 +609,8 @@ class DailyInputTransport(BaseInputTransport):
class DailyOutputTransport(BaseOutputTransport): class DailyOutputTransport(BaseOutputTransport):
def __init__(self, client: DailyTransportClient, params: DailyParams): def __init__(self, client: DailyTransportClient, params: DailyParams, **kwargs):
super().__init__(params) super().__init__(params, **kwargs)
self._client = client self._client = client
@@ -662,8 +662,10 @@ class DailyTransport(BaseTransport):
token: str | None, token: str | None,
bot_name: str, bot_name: str,
params: DailyParams, params: DailyParams,
input_name: str | None = None,
output_name: str | None = None,
loop: asyncio.AbstractEventLoop | None = None): loop: asyncio.AbstractEventLoop | None = None):
super().__init__(loop) super().__init__(input_name=input_name, output_name=output_name, loop=loop)
callbacks = DailyCallbacks( callbacks = DailyCallbacks(
on_joined=self._on_joined, on_joined=self._on_joined,
@@ -708,12 +710,12 @@ class DailyTransport(BaseTransport):
def input(self) -> FrameProcessor: def input(self) -> FrameProcessor:
if not self._input: if not self._input:
self._input = DailyInputTransport(self._client, self._params) self._input = DailyInputTransport(self._client, self._params, name=self._input_name)
return self._input return self._input
def output(self) -> FrameProcessor: def output(self) -> FrameProcessor:
if not self._output: if not self._output:
self._output = DailyOutputTransport(self._client, self._params) self._output = DailyOutputTransport(self._client, self._params, name=self._output_name)
return self._output return self._output
# #