Merge pull request #3519 from pipecat-ai/aleix/embedded-rtvi-processor
automatically add RTVI to the pipeline
This commit is contained in:
1
changelog/3519.added.2.md
Normal file
1
changelog/3519.added.2.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Added `RTVIProcessor.create_rtvi_observer()` factory method for creating RTVI observers.
|
||||||
1
changelog/3519.added.3.md
Normal file
1
changelog/3519.added.3.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Added `FrameProcessor.broadcast_frame_instance(frame)` method to broadcast a frame instance by extracting its fields and creating new instances for each direction.
|
||||||
1
changelog/3519.added.md
Normal file
1
changelog/3519.added.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- `PipelineTask` now automatically adds `RTVIProcessor` and registers `RTVIObserver` when `enable_rtvi=True` (default), simplifying pipeline setup.
|
||||||
1
changelog/3519.fixed.2.md
Normal file
1
changelog/3519.fixed.2.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Fixed `FrameProcessor.broadcast_frame()` to deep copy kwargs, preventing shared mutable references between the downstream and upstream frame instances.
|
||||||
1
changelog/3519.fixed.md
Normal file
1
changelog/3519.fixed.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Transports now properly broadcast `InputTransportMessageFrame` frames both upstream and downstream instead of only pushing downstream.
|
||||||
@@ -23,7 +23,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
LLMContextAggregatorPair,
|
LLMContextAggregatorPair,
|
||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIObserver, RTVIProcessor
|
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
@@ -93,12 +92,9 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
rtvi = RTVIProcessor()
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(),
|
transport.input(),
|
||||||
rtvi,
|
|
||||||
stt,
|
stt,
|
||||||
user_aggregator,
|
user_aggregator,
|
||||||
llm,
|
llm,
|
||||||
@@ -115,7 +111,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
enable_usage_metrics=True,
|
enable_usage_metrics=True,
|
||||||
),
|
),
|
||||||
observers=[
|
observers=[
|
||||||
RTVIObserver(rtvi),
|
|
||||||
DebugLogObserver(
|
DebugLogObserver(
|
||||||
frame_types={
|
frame_types={
|
||||||
TTSTextFrame: (BaseOutputTransport, FrameEndpoint.SOURCE),
|
TTSTextFrame: (BaseOutputTransport, FrameEndpoint.SOURCE),
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
LLMContextAggregatorPair,
|
LLMContextAggregatorPair,
|
||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIConfig, RTVIObserver, RTVIProcessor
|
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
@@ -88,12 +87,9 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
rtvi = RTVIProcessor(config=RTVIConfig(config=[]))
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(),
|
transport.input(),
|
||||||
rtvi,
|
|
||||||
stt,
|
stt,
|
||||||
user_aggregator,
|
user_aggregator,
|
||||||
llm,
|
llm,
|
||||||
@@ -110,7 +106,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
enable_usage_metrics=True,
|
enable_usage_metrics=True,
|
||||||
),
|
),
|
||||||
observers=[
|
observers=[
|
||||||
RTVIObserver(rtvi),
|
|
||||||
DebugLogObserver(
|
DebugLogObserver(
|
||||||
frame_types={
|
frame_types={
|
||||||
TTSTextFrame: (BaseOutputTransport, FrameEndpoint.SOURCE),
|
TTSTextFrame: (BaseOutputTransport, FrameEndpoint.SOURCE),
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
LLMContextAggregatorPair,
|
LLMContextAggregatorPair,
|
||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIConfig, RTVIObserver, RTVIProcessor
|
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
@@ -90,12 +89,9 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
rtvi = RTVIProcessor(config=RTVIConfig(config=[]))
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(), # Transport user input
|
transport.input(), # Transport user input
|
||||||
rtvi,
|
|
||||||
stt,
|
stt,
|
||||||
user_aggregator, # User responses
|
user_aggregator, # User responses
|
||||||
llm, # LLM
|
llm, # LLM
|
||||||
@@ -114,7 +110,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
),
|
),
|
||||||
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
||||||
observers=[
|
observers=[
|
||||||
RTVIObserver(rtvi),
|
|
||||||
DebugLogObserver(
|
DebugLogObserver(
|
||||||
frame_types={
|
frame_types={
|
||||||
TTSTextFrame: (BaseOutputTransport, FrameEndpoint.SOURCE),
|
TTSTextFrame: (BaseOutputTransport, FrameEndpoint.SOURCE),
|
||||||
@@ -123,10 +118,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
@rtvi.event_handler("on_client_ready")
|
|
||||||
async def on_client_ready(rtvi):
|
|
||||||
await rtvi.set_bot_ready()
|
|
||||||
|
|
||||||
@transport.event_handler("on_client_connected")
|
@transport.event_handler("on_client_connected")
|
||||||
async def on_client_connected(transport, client):
|
async def on_client_connected(transport, client):
|
||||||
logger.info(f"Client connected")
|
logger.info(f"Client connected")
|
||||||
|
|||||||
@@ -59,7 +59,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
LLMContextAggregatorPair,
|
LLMContextAggregatorPair,
|
||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIConfig, RTVIObserver, RTVIProcessor
|
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
@@ -255,12 +254,10 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
),
|
),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
rtvi = RTVIProcessor(config=RTVIConfig(config=[]))
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(),
|
transport.input(),
|
||||||
rtvi,
|
|
||||||
stt,
|
stt,
|
||||||
user_aggregator,
|
user_aggregator,
|
||||||
memory,
|
memory,
|
||||||
@@ -278,12 +275,10 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
enable_usage_metrics=True,
|
enable_usage_metrics=True,
|
||||||
),
|
),
|
||||||
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
||||||
observers=[RTVIObserver(rtvi)],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@rtvi.event_handler("on_client_ready")
|
@task.rtvi.event_handler("on_client_ready")
|
||||||
async def on_client_ready(rtvi):
|
async def on_client_ready(rtvi):
|
||||||
await rtvi.set_bot_ready()
|
|
||||||
# Get personalized greeting based on user memories. Can pass agent_id and run_id as per requirement of the application to manage short term memory or agent specific memory.
|
# Get personalized greeting based on user memories. Can pass agent_id and run_id as per requirement of the application to manage short term memory or agent specific memory.
|
||||||
greeting = await get_initial_greeting(
|
greeting = await get_initial_greeting(
|
||||||
memory_client=memory.memory_client, user_id=USER_ID, agent_id=None, run_id=None
|
memory_client=memory.memory_client, user_id=USER_ID, agent_id=None, run_id=None
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
LLMContextAggregatorPair,
|
LLMContextAggregatorPair,
|
||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIObserver, RTVIProcessor
|
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.cartesia.tts import CartesiaTTSService
|
from pipecat.services.cartesia.tts import CartesiaTTSService
|
||||||
@@ -87,8 +86,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
rtvi = RTVIProcessor()
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(), # Transport user input
|
transport.input(), # Transport user input
|
||||||
@@ -108,13 +105,11 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
enable_metrics=True,
|
enable_metrics=True,
|
||||||
enable_usage_metrics=True,
|
enable_usage_metrics=True,
|
||||||
),
|
),
|
||||||
observers=[RTVIObserver(rtvi)],
|
|
||||||
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
||||||
)
|
)
|
||||||
|
|
||||||
@rtvi.event_handler("on_client_ready")
|
@task.rtvi.event_handler("on_client_ready")
|
||||||
async def on_client_ready(rtvi):
|
async def on_client_ready(rtvi):
|
||||||
await rtvi.set_bot_ready()
|
|
||||||
# Kick off the conversation
|
# Kick off the conversation
|
||||||
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
await task.queue_frames([LLMRunFrame()])
|
await task.queue_frames([LLMRunFrame()])
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
#
|
#
|
||||||
# Copyright (c) 2025, Daily
|
# Copyright (c) 2024-2026, Daily
|
||||||
#
|
#
|
||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
@@ -22,7 +22,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIObserver, RTVIProcessor
|
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.google.gemini_live.llm import GeminiLiveLLMService
|
from pipecat.services.google.gemini_live.llm import GeminiLiveLLMService
|
||||||
@@ -125,14 +124,10 @@ async def run_bot(pipecat_transport):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# RTVI events for Pipecat client UI
|
|
||||||
rtvi = RTVIProcessor()
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
pipecat_transport.input(),
|
pipecat_transport.input(),
|
||||||
user_aggregator,
|
user_aggregator,
|
||||||
rtvi,
|
|
||||||
llm, # LLM
|
llm, # LLM
|
||||||
EdgeDetectionProcessor(
|
EdgeDetectionProcessor(
|
||||||
pipecat_transport._params.video_out_width,
|
pipecat_transport._params.video_out_width,
|
||||||
@@ -149,13 +144,11 @@ async def run_bot(pipecat_transport):
|
|||||||
enable_metrics=True,
|
enable_metrics=True,
|
||||||
enable_usage_metrics=True,
|
enable_usage_metrics=True,
|
||||||
),
|
),
|
||||||
observers=[RTVIObserver(rtvi)],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@rtvi.event_handler("on_client_ready")
|
@task.rtvi.event_handler("on_client_ready")
|
||||||
async def on_client_ready(rtvi):
|
async def on_client_ready(rtvi):
|
||||||
logger.info("Pipecat client ready.")
|
logger.info("Pipecat client ready.")
|
||||||
await rtvi.set_bot_ready()
|
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
await task.queue_frames([LLMRunFrame()])
|
await task.queue_frames([LLMRunFrame()])
|
||||||
|
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ from pipecat.pipeline.pipeline import Pipeline, PipelineSink, PipelineSource
|
|||||||
from pipecat.pipeline.task_observer import TaskObserver
|
from pipecat.pipeline.task_observer import TaskObserver
|
||||||
from pipecat.processors.aggregators.llm_response import LLMUserContextAggregator
|
from pipecat.processors.aggregators.llm_response import LLMUserContextAggregator
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor, FrameProcessorSetup
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor, FrameProcessorSetup
|
||||||
|
from pipecat.processors.frameworks.rtvi import RTVIObserverParams, RTVIProcessor
|
||||||
from pipecat.utils.asyncio.task_manager import BaseTaskManager, TaskManager, TaskManagerParams
|
from pipecat.utils.asyncio.task_manager import BaseTaskManager, TaskManager, TaskManagerParams
|
||||||
from pipecat.utils.tracing.setup import is_tracing_available
|
from pipecat.utils.tracing.setup import is_tracing_available
|
||||||
from pipecat.utils.tracing.turn_trace_observer import TurnTraceObserver
|
from pipecat.utils.tracing.turn_trace_observer import TurnTraceObserver
|
||||||
@@ -225,9 +226,12 @@ class PipelineTask(BasePipelineTask):
|
|||||||
conversation_id: Optional[str] = None,
|
conversation_id: Optional[str] = None,
|
||||||
enable_tracing: bool = False,
|
enable_tracing: bool = False,
|
||||||
enable_turn_tracking: bool = True,
|
enable_turn_tracking: bool = True,
|
||||||
|
enable_rtvi: bool = True,
|
||||||
idle_timeout_frames: Tuple[Type[Frame], ...] = (BotSpeakingFrame, UserSpeakingFrame),
|
idle_timeout_frames: Tuple[Type[Frame], ...] = (BotSpeakingFrame, UserSpeakingFrame),
|
||||||
idle_timeout_secs: Optional[float] = IDLE_TIMEOUT_SECS,
|
idle_timeout_secs: Optional[float] = IDLE_TIMEOUT_SECS,
|
||||||
observers: Optional[List[BaseObserver]] = None,
|
observers: Optional[List[BaseObserver]] = None,
|
||||||
|
rtvi_processor: Optional[RTVIProcessor] = None,
|
||||||
|
rtvi_observer_params: Optional[RTVIObserverParams] = None,
|
||||||
task_manager: Optional[BaseTaskManager] = None,
|
task_manager: Optional[BaseTaskManager] = None,
|
||||||
):
|
):
|
||||||
"""Initialize the PipelineTask.
|
"""Initialize the PipelineTask.
|
||||||
@@ -244,6 +248,7 @@ class PipelineTask(BasePipelineTask):
|
|||||||
check_dangling_tasks: Whether to check for processors' tasks finishing properly.
|
check_dangling_tasks: Whether to check for processors' tasks finishing properly.
|
||||||
clock: Clock implementation for timing operations.
|
clock: Clock implementation for timing operations.
|
||||||
conversation_id: Optional custom ID for the conversation.
|
conversation_id: Optional custom ID for the conversation.
|
||||||
|
enable_rtvi: Whether to automatically add RTVI support to the pipeline.
|
||||||
enable_tracing: Whether to enable tracing.
|
enable_tracing: Whether to enable tracing.
|
||||||
enable_turn_tracking: Whether to enable turn tracking.
|
enable_turn_tracking: Whether to enable turn tracking.
|
||||||
idle_timeout_frames: A tuple with the frames that should trigger an idle
|
idle_timeout_frames: A tuple with the frames that should trigger an idle
|
||||||
@@ -252,6 +257,8 @@ class PipelineTask(BasePipelineTask):
|
|||||||
None. If a pipeline is idle the pipeline task will be cancelled
|
None. If a pipeline is idle the pipeline task will be cancelled
|
||||||
automatically.
|
automatically.
|
||||||
observers: List of observers for monitoring pipeline execution.
|
observers: List of observers for monitoring pipeline execution.
|
||||||
|
rtvi_observer_params: The RTVI observer parameter to use if RTVI is enabled.
|
||||||
|
rtvi_processor: The RTVI processor to add if RTVI is enabled.
|
||||||
task_manager: Optional task manager for handling asyncio tasks.
|
task_manager: Optional task manager for handling asyncio tasks.
|
||||||
"""
|
"""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -306,6 +313,16 @@ class PipelineTask(BasePipelineTask):
|
|||||||
self._heartbeat_push_task: Optional[asyncio.Task] = None
|
self._heartbeat_push_task: Optional[asyncio.Task] = None
|
||||||
self._heartbeat_monitor_task: Optional[asyncio.Task] = None
|
self._heartbeat_monitor_task: Optional[asyncio.Task] = None
|
||||||
|
|
||||||
|
# RTVI support
|
||||||
|
self._rtvi = None
|
||||||
|
if enable_rtvi:
|
||||||
|
self._rtvi = rtvi_processor or RTVIProcessor()
|
||||||
|
observers.append(self._rtvi.create_rtvi_observer(params=rtvi_observer_params))
|
||||||
|
|
||||||
|
@self.rtvi.event_handler("on_client_ready")
|
||||||
|
async def on_client_ready(rtvi: RTVIProcessor):
|
||||||
|
await rtvi.set_bot_ready()
|
||||||
|
|
||||||
# This is the idle event. When selected frames are pushed from any
|
# This is the idle event. When selected frames are pushed from any
|
||||||
# processor we consider the pipeline is not idle. We use an observer
|
# processor we consider the pipeline is not idle. We use an observer
|
||||||
# which will be listening any part of the pipeline.
|
# which will be listening any part of the pipeline.
|
||||||
@@ -335,7 +352,8 @@ class PipelineTask(BasePipelineTask):
|
|||||||
# allows us to receive and react to downstream frames.
|
# allows us to receive and react to downstream frames.
|
||||||
source = PipelineSource(self._source_push_frame, name=f"{self}::Source")
|
source = PipelineSource(self._source_push_frame, name=f"{self}::Source")
|
||||||
sink = PipelineSink(self._sink_push_frame, name=f"{self}::Sink")
|
sink = PipelineSink(self._sink_push_frame, name=f"{self}::Sink")
|
||||||
self._pipeline = Pipeline([pipeline], source=source, sink=sink)
|
processors = [self._rtvi, pipeline] if self._rtvi else [pipeline]
|
||||||
|
self._pipeline = Pipeline(processors, source=source, sink=sink)
|
||||||
|
|
||||||
# The task observer acts as a proxy to the provided observers. This way,
|
# The task observer acts as a proxy to the provided observers. This way,
|
||||||
# we only need to pass a single observer (using the StartFrame) which
|
# we only need to pass a single observer (using the StartFrame) which
|
||||||
@@ -398,6 +416,17 @@ class PipelineTask(BasePipelineTask):
|
|||||||
"""
|
"""
|
||||||
return self._turn_trace_observer
|
return self._turn_trace_observer
|
||||||
|
|
||||||
|
@property
|
||||||
|
def rtvi(self) -> RTVIProcessor:
|
||||||
|
"""Get the RTVI processor if RTVI is enabled.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The RTVI processor added to the pipeline when RTVI is enabled.
|
||||||
|
"""
|
||||||
|
if not self._rtvi:
|
||||||
|
raise Exception(f"{self} RTVI is not enabled.")
|
||||||
|
return self._rtvi
|
||||||
|
|
||||||
def event_handler(self, event_name: str):
|
def event_handler(self, event_name: str):
|
||||||
"""Decorator for registering event handlers.
|
"""Decorator for registering event handlers.
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,9 @@ management, and frame flow control mechanisms.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import dataclasses
|
||||||
import traceback
|
import traceback
|
||||||
|
from copy import deepcopy
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import (
|
from typing import (
|
||||||
@@ -779,8 +781,40 @@ class FrameProcessor(BaseObject):
|
|||||||
frame_cls: The class of the frame to be broadcasted.
|
frame_cls: The class of the frame to be broadcasted.
|
||||||
**kwargs: Keyword arguments to be passed to the frame's constructor.
|
**kwargs: Keyword arguments to be passed to the frame's constructor.
|
||||||
"""
|
"""
|
||||||
await self.push_frame(frame_cls(**kwargs))
|
await self.push_frame(frame_cls(**deepcopy(kwargs)))
|
||||||
await self.push_frame(frame_cls(**kwargs), FrameDirection.UPSTREAM)
|
await self.push_frame(frame_cls(**deepcopy(kwargs)), FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
|
async def broadcast_frame_instance(self, frame: Frame):
|
||||||
|
"""Broadcasts a frame instance upstream and downstream.
|
||||||
|
|
||||||
|
This method creates two new frame instances copying all fields from the
|
||||||
|
original frame except `id` and `name`, which get fresh values.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
frame: The frame instance to broadcast.
|
||||||
|
|
||||||
|
Note:
|
||||||
|
Prefer using `broadcast_frame()` when possible, as it is more
|
||||||
|
efficient. This method should only be used when you are not the
|
||||||
|
creator of the frame and need to broadcast an existing instance.
|
||||||
|
"""
|
||||||
|
frame_cls = type(frame)
|
||||||
|
init_fields = {f.name: getattr(frame, f.name) for f in dataclasses.fields(frame) if f.init}
|
||||||
|
extra_fields = {
|
||||||
|
f.name: getattr(frame, f.name)
|
||||||
|
for f in dataclasses.fields(frame)
|
||||||
|
if not f.init and f.name not in ("id", "name")
|
||||||
|
}
|
||||||
|
|
||||||
|
new_frame = frame_cls(**deepcopy(init_fields))
|
||||||
|
for k, v in deepcopy(extra_fields).items():
|
||||||
|
setattr(new_frame, k, v)
|
||||||
|
await self.push_frame(new_frame)
|
||||||
|
|
||||||
|
new_frame = frame_cls(**deepcopy(init_fields))
|
||||||
|
for k, v in deepcopy(extra_fields).items():
|
||||||
|
setattr(new_frame, k, v)
|
||||||
|
await self.push_frame(new_frame, FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
async def __start(self, frame: StartFrame):
|
async def __start(self, frame: StartFrame):
|
||||||
"""Handle the start frame to initialize processor state.
|
"""Handle the start frame to initialize processor state.
|
||||||
|
|||||||
@@ -1100,13 +1100,11 @@ class RTVIObserver(BaseObserver):
|
|||||||
|
|
||||||
if (
|
if (
|
||||||
isinstance(frame, (UserStartedSpeakingFrame, UserStoppedSpeakingFrame))
|
isinstance(frame, (UserStartedSpeakingFrame, UserStoppedSpeakingFrame))
|
||||||
and (direction == FrameDirection.DOWNSTREAM)
|
|
||||||
and self._params.user_speaking_enabled
|
and self._params.user_speaking_enabled
|
||||||
):
|
):
|
||||||
await self._handle_interruptions(frame)
|
await self._handle_interruptions(frame)
|
||||||
elif (
|
elif (
|
||||||
isinstance(frame, (BotStartedSpeakingFrame, BotStoppedSpeakingFrame))
|
isinstance(frame, (BotStartedSpeakingFrame, BotStoppedSpeakingFrame))
|
||||||
and (direction == FrameDirection.UPSTREAM)
|
|
||||||
and self._params.bot_speaking_enabled
|
and self._params.bot_speaking_enabled
|
||||||
):
|
):
|
||||||
await self._handle_bot_speaking(frame)
|
await self._handle_bot_speaking(frame)
|
||||||
@@ -1413,6 +1411,18 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
|
|
||||||
self._registered_services[service.name] = service
|
self._registered_services[service.name] = service
|
||||||
|
|
||||||
|
def create_rtvi_observer(self, *, params: Optional[RTVIObserverParams] = None, **kwargs):
|
||||||
|
"""Creates a new RTVI Observer.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
params: Settings to enable/disable specific messages.
|
||||||
|
**kwargs: Additional arguments passed to the observer.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A new RTVI observer.
|
||||||
|
"""
|
||||||
|
return RTVIObserver(self, params=params, **kwargs)
|
||||||
|
|
||||||
async def set_client_ready(self):
|
async def set_client_ready(self):
|
||||||
"""Mark the client as ready and trigger the ready event."""
|
"""Mark the client as ready and trigger the ready event."""
|
||||||
self._client_ready = True
|
self._client_ready = True
|
||||||
|
|||||||
@@ -126,7 +126,7 @@ class ProtobufFrameSerializer(FrameSerializer):
|
|||||||
if "pts" in args_dict:
|
if "pts" in args_dict:
|
||||||
del args_dict["pts"]
|
del args_dict["pts"]
|
||||||
|
|
||||||
# Special handling for MessageFrame -> OutputTransportMessageUrgentFrame
|
# Special handling for MessageFrame -> InputTransportMessageFrame
|
||||||
if class_name == MessageFrame:
|
if class_name == MessageFrame:
|
||||||
try:
|
try:
|
||||||
msg = json.loads(args_dict["data"])
|
msg = json.loads(args_dict["data"])
|
||||||
|
|||||||
@@ -4,7 +4,7 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
"""Google RTVI integration models and observer implementation.
|
"""Google RTVI processor and observer implementation.
|
||||||
|
|
||||||
This module provides integration with Google's services through the RTVI framework,
|
This module provides integration with Google's services through the RTVI framework,
|
||||||
including models for search responses and an observer for handling Google-specific
|
including models for search responses and an observer for handling Google-specific
|
||||||
@@ -16,7 +16,7 @@ from typing import List, Literal, Optional
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from pipecat.observers.base_observer import FramePushed
|
from pipecat.observers.base_observer import FramePushed
|
||||||
from pipecat.processors.frameworks.rtvi import RTVIObserver, RTVIProcessor
|
from pipecat.processors.frameworks.rtvi import RTVIObserver, RTVIObserverParams, RTVIProcessor
|
||||||
from pipecat.services.google.frames import LLMSearchOrigin, LLMSearchResponseFrame
|
from pipecat.services.google.frames import LLMSearchOrigin, LLMSearchResponseFrame
|
||||||
|
|
||||||
|
|
||||||
@@ -86,4 +86,23 @@ class GoogleRTVIObserver(RTVIObserver):
|
|||||||
rendered_content=frame.rendered_content,
|
rendered_content=frame.rendered_content,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
await self.push_transport_message_urgent(message)
|
await self.send_rtvi_message(message)
|
||||||
|
|
||||||
|
|
||||||
|
class GoogleRTVIProcessor(RTVIProcessor):
|
||||||
|
"""RTVI processor for Google service integration.
|
||||||
|
|
||||||
|
Creates a specific Google RTVI Observer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def create_rtvi_observer(self, *, params: Optional[RTVIObserverParams] = None, **kwargs):
|
||||||
|
"""Creates a new RTVI Observer.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
params: Settings to enable/disable specific messages.
|
||||||
|
**kwargs: Additional arguments passed to the observer.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A new RTVI observer.
|
||||||
|
"""
|
||||||
|
return GoogleRTVIObserver(self)
|
||||||
|
|||||||
@@ -123,9 +123,10 @@ class QueuedFrameProcessor(FrameProcessor):
|
|||||||
async def run_test(
|
async def run_test(
|
||||||
processor: FrameProcessor,
|
processor: FrameProcessor,
|
||||||
*,
|
*,
|
||||||
frames_to_send: Sequence[Frame],
|
enable_rtvi: bool = False,
|
||||||
expected_down_frames: Optional[Sequence[type]] = None,
|
expected_down_frames: Optional[Sequence[type]] = None,
|
||||||
expected_up_frames: Optional[Sequence[type]] = None,
|
expected_up_frames: Optional[Sequence[type]] = None,
|
||||||
|
frames_to_send: Sequence[Frame],
|
||||||
ignore_start: bool = True,
|
ignore_start: bool = True,
|
||||||
observers: Optional[List[BaseObserver]] = None,
|
observers: Optional[List[BaseObserver]] = None,
|
||||||
pipeline_params: Optional[PipelineParams] = None,
|
pipeline_params: Optional[PipelineParams] = None,
|
||||||
@@ -139,9 +140,10 @@ async def run_test(
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
processor: The frame processor to test.
|
processor: The frame processor to test.
|
||||||
frames_to_send: Sequence of frames to send through the processor.
|
enable_rtvi: Whether RTVI should be enabled in this test.
|
||||||
expected_down_frames: Expected frame types flowing downstream (optional).
|
expected_down_frames: Expected frame types flowing downstream (optional).
|
||||||
expected_up_frames: Expected frame types flowing upstream (optional).
|
expected_up_frames: Expected frame types flowing upstream (optional).
|
||||||
|
frames_to_send: Sequence of frames to send through the processor.
|
||||||
ignore_start: Whether to ignore StartFrames in frame validation.
|
ignore_start: Whether to ignore StartFrames in frame validation.
|
||||||
observers: Optional list of observers to attach to the pipeline.
|
observers: Optional list of observers to attach to the pipeline.
|
||||||
pipeline_params: Optional pipeline parameters.
|
pipeline_params: Optional pipeline parameters.
|
||||||
@@ -173,9 +175,10 @@ async def run_test(
|
|||||||
|
|
||||||
task = PipelineTask(
|
task = PipelineTask(
|
||||||
pipeline,
|
pipeline,
|
||||||
params=pipeline_params,
|
|
||||||
observers=observers,
|
|
||||||
cancel_on_idle_timeout=False,
|
cancel_on_idle_timeout=False,
|
||||||
|
enable_rtvi=enable_rtvi,
|
||||||
|
observers=observers,
|
||||||
|
params=pipeline_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def push_frames():
|
async def push_frames():
|
||||||
|
|||||||
@@ -1728,8 +1728,9 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
message: The message data to send.
|
message: The message data to send.
|
||||||
sender: ID of the message sender.
|
sender: ID of the message sender.
|
||||||
"""
|
"""
|
||||||
frame = DailyInputTransportMessageFrame(message=message, participant_id=sender)
|
await self.broadcast_frame_class(
|
||||||
await self.push_frame(frame)
|
DailyInputTransportMessageFrame, message=message, participant_id=sender
|
||||||
|
)
|
||||||
|
|
||||||
#
|
#
|
||||||
# Audio in
|
# Audio in
|
||||||
|
|||||||
@@ -698,8 +698,7 @@ class SmallWebRTCInputTransport(BaseInputTransport):
|
|||||||
message: The application message to process.
|
message: The application message to process.
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Received app message inside SmallWebRTCInputTransport {message}")
|
logger.debug(f"Received app message inside SmallWebRTCInputTransport {message}")
|
||||||
frame = InputTransportMessageFrame(message=message)
|
await self.broadcast_frame_class(InputTransportMessageFrame, message=message)
|
||||||
await self.push_frame(frame)
|
|
||||||
|
|
||||||
# Add this method similar to DailyInputTransport.request_participant_image
|
# Add this method similar to DailyInputTransport.request_participant_image
|
||||||
async def request_participant_image(self, frame: UserImageRequestFrame):
|
async def request_participant_image(self, frame: UserImageRequestFrame):
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from pipecat.frames.frames import (
|
|||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InputAudioRawFrame,
|
InputAudioRawFrame,
|
||||||
|
InputTransportMessageFrame,
|
||||||
OutputAudioRawFrame,
|
OutputAudioRawFrame,
|
||||||
OutputTransportMessageFrame,
|
OutputTransportMessageFrame,
|
||||||
OutputTransportMessageUrgentFrame,
|
OutputTransportMessageUrgentFrame,
|
||||||
@@ -298,6 +299,8 @@ class WebsocketClientInputTransport(BaseInputTransport):
|
|||||||
return
|
return
|
||||||
if isinstance(frame, InputAudioRawFrame) and self._params.audio_in_enabled:
|
if isinstance(frame, InputAudioRawFrame) and self._params.audio_in_enabled:
|
||||||
await self.push_audio_frame(frame)
|
await self.push_audio_frame(frame)
|
||||||
|
elif isinstance(frame, InputTransportMessageFrame):
|
||||||
|
await self.broadcast_frame(frame)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ from pipecat.frames.frames import (
|
|||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InputAudioRawFrame,
|
InputAudioRawFrame,
|
||||||
|
InputTransportMessageFrame,
|
||||||
InterruptionFrame,
|
InterruptionFrame,
|
||||||
OutputAudioRawFrame,
|
OutputAudioRawFrame,
|
||||||
OutputTransportMessageFrame,
|
OutputTransportMessageFrame,
|
||||||
@@ -311,6 +312,8 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
|
|||||||
|
|
||||||
if isinstance(frame, InputAudioRawFrame):
|
if isinstance(frame, InputAudioRawFrame):
|
||||||
await self.push_audio_frame(frame)
|
await self.push_audio_frame(frame)
|
||||||
|
elif isinstance(frame, InputTransportMessageFrame):
|
||||||
|
await self.broadcast_frame(frame)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -25,6 +25,8 @@ from pipecat.frames.frames import (
|
|||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InputAudioRawFrame,
|
InputAudioRawFrame,
|
||||||
|
InputTransportMessageFrame,
|
||||||
|
InputTransportMessageUrgentFrame,
|
||||||
InterruptionFrame,
|
InterruptionFrame,
|
||||||
OutputAudioRawFrame,
|
OutputAudioRawFrame,
|
||||||
OutputTransportMessageFrame,
|
OutputTransportMessageFrame,
|
||||||
@@ -214,6 +216,8 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
|
|
||||||
if isinstance(frame, InputAudioRawFrame):
|
if isinstance(frame, InputAudioRawFrame):
|
||||||
await self.push_audio_frame(frame)
|
await self.push_audio_frame(frame)
|
||||||
|
elif isinstance(frame, InputTransportMessageFrame):
|
||||||
|
await self.broadcast_frame(frame)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -6,7 +6,8 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import unittest
|
import unittest
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, field
|
||||||
|
from typing import List
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
DataFrame,
|
DataFrame,
|
||||||
@@ -24,6 +25,15 @@ from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
|||||||
from pipecat.tests.utils import SleepFrame, run_test
|
from pipecat.tests.utils import SleepFrame, run_test
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class BroadcastTestFrame(DataFrame):
|
||||||
|
"""Test frame with init fields for broadcast testing."""
|
||||||
|
|
||||||
|
text: str = ""
|
||||||
|
value: int = 0
|
||||||
|
items: List[str] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
class TestFrameProcessor(unittest.IsolatedAsyncioTestCase):
|
class TestFrameProcessor(unittest.IsolatedAsyncioTestCase):
|
||||||
async def test_before_after_events(self):
|
async def test_before_after_events(self):
|
||||||
identity = IdentityFilter()
|
identity = IdentityFilter()
|
||||||
@@ -186,3 +196,157 @@ class TestFrameProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send=frames_to_send,
|
frames_to_send=frames_to_send,
|
||||||
expected_down_frames=expected_down_frames,
|
expected_down_frames=expected_down_frames,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def test_broadcast_frame(self):
|
||||||
|
"""Test that broadcast_frame creates two separate frames with fresh IDs."""
|
||||||
|
downstream_frames: List[Frame] = []
|
||||||
|
upstream_frames: List[Frame] = []
|
||||||
|
|
||||||
|
class BroadcastTestProcessor(FrameProcessor):
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
if isinstance(frame, TextFrame):
|
||||||
|
await self.broadcast_frame(
|
||||||
|
BroadcastTestFrame, text="hello", value=42, items=["a", "b"]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
class CaptureProcessor(FrameProcessor):
|
||||||
|
def __init__(self, capture_list: List[Frame], direction: FrameDirection):
|
||||||
|
super().__init__()
|
||||||
|
self._capture_list = capture_list
|
||||||
|
self._capture_direction = direction
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
if direction == self._capture_direction and isinstance(frame, BroadcastTestFrame):
|
||||||
|
self._capture_list.append(frame)
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
up_capture = CaptureProcessor(upstream_frames, FrameDirection.UPSTREAM)
|
||||||
|
broadcaster = BroadcastTestProcessor()
|
||||||
|
down_capture = CaptureProcessor(downstream_frames, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
|
pipeline = Pipeline([up_capture, broadcaster, down_capture])
|
||||||
|
|
||||||
|
frames_to_send = [TextFrame(text="trigger")]
|
||||||
|
expected_down_frames = [BroadcastTestFrame]
|
||||||
|
expected_up_frames = [BroadcastTestFrame]
|
||||||
|
|
||||||
|
await run_test(
|
||||||
|
pipeline,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
expected_up_frames=expected_up_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify we got one frame in each direction
|
||||||
|
self.assertEqual(len(downstream_frames), 1)
|
||||||
|
self.assertEqual(len(upstream_frames), 1)
|
||||||
|
|
||||||
|
down_frame = downstream_frames[0]
|
||||||
|
up_frame = upstream_frames[0]
|
||||||
|
|
||||||
|
# Verify the frames have different IDs (they are separate instances)
|
||||||
|
self.assertNotEqual(down_frame.id, up_frame.id)
|
||||||
|
|
||||||
|
# Verify the frames have the correct field values
|
||||||
|
self.assertEqual(down_frame.text, "hello")
|
||||||
|
self.assertEqual(down_frame.value, 42)
|
||||||
|
self.assertEqual(down_frame.items, ["a", "b"])
|
||||||
|
self.assertEqual(up_frame.text, "hello")
|
||||||
|
self.assertEqual(up_frame.value, 42)
|
||||||
|
self.assertEqual(up_frame.items, ["a", "b"])
|
||||||
|
|
||||||
|
# Verify the items lists are separate instances (not shared references)
|
||||||
|
self.assertIsNot(down_frame.items, up_frame.items)
|
||||||
|
|
||||||
|
async def test_broadcast_frame_instance(self):
|
||||||
|
"""Test that broadcast_frame_instance copies all fields except id and name."""
|
||||||
|
downstream_frames: List[Frame] = []
|
||||||
|
upstream_frames: List[Frame] = []
|
||||||
|
original_frame: List[Frame] = []
|
||||||
|
|
||||||
|
class BroadcastInstanceTestProcessor(FrameProcessor):
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
if isinstance(frame, BroadcastTestFrame):
|
||||||
|
# Set some non-init fields on the frame
|
||||||
|
frame.pts = 12345
|
||||||
|
frame.metadata = {"key": "value", "nested": {"a": 1}}
|
||||||
|
original_frame.append(frame)
|
||||||
|
await self.broadcast_frame_instance(frame)
|
||||||
|
else:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
class CaptureProcessor(FrameProcessor):
|
||||||
|
def __init__(self, capture_list: List[Frame], direction: FrameDirection):
|
||||||
|
super().__init__()
|
||||||
|
self._capture_list = capture_list
|
||||||
|
self._capture_direction = direction
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
if direction == self._capture_direction and isinstance(frame, BroadcastTestFrame):
|
||||||
|
self._capture_list.append(frame)
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
up_capture = CaptureProcessor(upstream_frames, FrameDirection.UPSTREAM)
|
||||||
|
broadcaster = BroadcastInstanceTestProcessor()
|
||||||
|
down_capture = CaptureProcessor(downstream_frames, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
|
pipeline = Pipeline([up_capture, broadcaster, down_capture])
|
||||||
|
|
||||||
|
# Create a frame with mutable fields to test deep copying
|
||||||
|
test_frame = BroadcastTestFrame(text="test", value=99, items=["x", "y", "z"])
|
||||||
|
|
||||||
|
frames_to_send = [test_frame]
|
||||||
|
expected_down_frames = [BroadcastTestFrame]
|
||||||
|
expected_up_frames = [BroadcastTestFrame]
|
||||||
|
|
||||||
|
await run_test(
|
||||||
|
pipeline,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
expected_up_frames=expected_up_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify we got one frame in each direction
|
||||||
|
self.assertEqual(len(downstream_frames), 1)
|
||||||
|
self.assertEqual(len(upstream_frames), 1)
|
||||||
|
self.assertEqual(len(original_frame), 1)
|
||||||
|
|
||||||
|
orig = original_frame[0]
|
||||||
|
down_frame = downstream_frames[0]
|
||||||
|
up_frame = upstream_frames[0]
|
||||||
|
|
||||||
|
# Verify the frames have different IDs and names (fresh values)
|
||||||
|
self.assertNotEqual(down_frame.id, orig.id)
|
||||||
|
self.assertNotEqual(up_frame.id, orig.id)
|
||||||
|
self.assertNotEqual(down_frame.id, up_frame.id)
|
||||||
|
self.assertNotEqual(down_frame.name, orig.name)
|
||||||
|
self.assertNotEqual(up_frame.name, orig.name)
|
||||||
|
|
||||||
|
# Verify init fields are copied correctly
|
||||||
|
self.assertEqual(down_frame.text, "test")
|
||||||
|
self.assertEqual(down_frame.value, 99)
|
||||||
|
self.assertEqual(down_frame.items, ["x", "y", "z"])
|
||||||
|
self.assertEqual(up_frame.text, "test")
|
||||||
|
self.assertEqual(up_frame.value, 99)
|
||||||
|
self.assertEqual(up_frame.items, ["x", "y", "z"])
|
||||||
|
|
||||||
|
# Verify non-init fields (except id/name) are copied
|
||||||
|
self.assertEqual(down_frame.pts, 12345)
|
||||||
|
self.assertEqual(down_frame.metadata, {"key": "value", "nested": {"a": 1}})
|
||||||
|
self.assertEqual(up_frame.pts, 12345)
|
||||||
|
self.assertEqual(up_frame.metadata, {"key": "value", "nested": {"a": 1}})
|
||||||
|
|
||||||
|
# Verify mutable fields are deep copied (not shared references)
|
||||||
|
self.assertIsNot(down_frame.items, orig.items)
|
||||||
|
self.assertIsNot(up_frame.items, orig.items)
|
||||||
|
self.assertIsNot(down_frame.items, up_frame.items)
|
||||||
|
self.assertIsNot(down_frame.metadata, orig.metadata)
|
||||||
|
self.assertIsNot(up_frame.metadata, orig.metadata)
|
||||||
|
self.assertIsNot(down_frame.metadata, up_frame.metadata)
|
||||||
|
self.assertIsNot(down_frame.metadata["nested"], up_frame.metadata["nested"])
|
||||||
|
|||||||
Reference in New Issue
Block a user