Merge pull request #2639 from pipecat-ai/aleix/min-words-interruption-unit-test

MinWordsInterruptionStrategy unit test
This commit is contained in:
Aleix Conchillo Flaqué
2025-09-11 18:52:39 -07:00
committed by GitHub
3 changed files with 109 additions and 5 deletions

View File

@@ -128,7 +128,7 @@ async def run_test(
expected_up_frames: Optional[Sequence[type]] = None, expected_up_frames: Optional[Sequence[type]] = None,
ignore_start: bool = True, ignore_start: bool = True,
observers: Optional[List[BaseObserver]] = None, observers: Optional[List[BaseObserver]] = None,
start_metadata: Optional[Dict[str, Any]] = None, pipeline_params: Optional[PipelineParams] = None,
send_end_frame: bool = True, send_end_frame: bool = True,
) -> Tuple[Sequence[Frame], Sequence[Frame]]: ) -> Tuple[Sequence[Frame], Sequence[Frame]]:
"""Run a test pipeline with the specified processor and validate frame flow. """Run a test pipeline with the specified processor and validate frame flow.
@@ -144,7 +144,7 @@ async def run_test(
expected_up_frames: Expected frame types flowing upstream (optional). expected_up_frames: Expected frame types flowing upstream (optional).
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.
start_metadata: Optional metadata to include with the StartFrame. pipeline_params: Optional pipeline parameters.
send_end_frame: Whether to send an EndFrame at the end of the test. send_end_frame: Whether to send an EndFrame at the end of the test.
Returns: Returns:
@@ -154,7 +154,7 @@ async def run_test(
AssertionError: If the received frames don't match the expected frame types. AssertionError: If the received frames don't match the expected frame types.
""" """
observers = observers or [] observers = observers or []
start_metadata = start_metadata or {} pipeline_params = pipeline_params or PipelineParams()
received_up = asyncio.Queue() received_up = asyncio.Queue()
received_down = asyncio.Queue() received_down = asyncio.Queue()
@@ -173,7 +173,7 @@ async def run_test(
task = PipelineTask( task = PipelineTask(
pipeline, pipeline,
params=PipelineParams(start_metadata=start_metadata), params=pipeline_params,
observers=observers, observers=observers,
cancel_on_idle_timeout=False, cancel_on_idle_timeout=False,
) )

View File

@@ -8,16 +8,20 @@ import json
import unittest import unittest
from typing import Any from typing import Any
from pipecat.audio.interruptions.min_words_interruption_strategy import MinWordsInterruptionStrategy
from pipecat.audio.turn.smart_turn.base_smart_turn import SmartTurnParams from pipecat.audio.turn.smart_turn.base_smart_turn import SmartTurnParams
from pipecat.audio.vad.vad_analyzer import VADParams from pipecat.audio.vad.vad_analyzer import VADParams
from pipecat.frames.frames import ( from pipecat.frames.frames import (
BotStartedSpeakingFrame,
EmulateUserStartedSpeakingFrame, EmulateUserStartedSpeakingFrame,
EmulateUserStoppedSpeakingFrame, EmulateUserStoppedSpeakingFrame,
Frame,
FunctionCallInProgressFrame, FunctionCallInProgressFrame,
FunctionCallResultFrame, FunctionCallResultFrame,
FunctionCallResultProperties, FunctionCallResultProperties,
InterimTranscriptionFrame, InterimTranscriptionFrame,
InterruptionFrame, InterruptionFrame,
InterruptionTaskFrame,
LLMFullResponseEndFrame, LLMFullResponseEndFrame,
LLMFullResponseStartFrame, LLMFullResponseStartFrame,
OpenAILLMContextAssistantTimestampFrame, OpenAILLMContextAssistantTimestampFrame,
@@ -27,6 +31,8 @@ from pipecat.frames.frames import (
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
) )
from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.task import PipelineParams
from pipecat.processors.aggregators.llm_response import ( from pipecat.processors.aggregators.llm_response import (
LLMAssistantAggregatorParams, LLMAssistantAggregatorParams,
LLMUserAggregatorParams, LLMUserAggregatorParams,
@@ -36,6 +42,7 @@ from pipecat.processors.aggregators.openai_llm_context import (
OpenAILLMContext, OpenAILLMContext,
OpenAILLMContextFrame, OpenAILLMContextFrame,
) )
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from pipecat.services.anthropic.llm import ( from pipecat.services.anthropic.llm import (
AnthropicAssistantContextAggregator, AnthropicAssistantContextAggregator,
AnthropicLLMContext, AnthropicLLMContext,
@@ -481,6 +488,103 @@ class BaseTestUserContextAggregator:
) )
self.check_message_content(context, 0, "How are you?") self.check_message_content(context, 0, "How are you?")
async def test_min_words_interruption_strategy_one_word(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
class ContextProcessor(FrameProcessor):
def __init__(self):
super().__init__()
self.context_received = False
async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if isinstance(frame, OpenAILLMContextFrame):
self.context_received = True
await self.push_frame(frame, direction)
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
context_processor = ContextProcessor()
pipeline = Pipeline([aggregator, context_processor])
frames_to_send = [
BotStartedSpeakingFrame(),
UserStartedSpeakingFrame(),
TranscriptionFrame(text="Can", user_id="cat", timestamp=""),
SleepFrame(),
UserStoppedSpeakingFrame(),
]
expected_down_frames = [
BotStartedSpeakingFrame,
UserStartedSpeakingFrame,
UserStoppedSpeakingFrame,
]
await run_test(
pipeline,
frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames,
pipeline_params=PipelineParams(
interruption_strategies=[MinWordsInterruptionStrategy(min_words=2)]
),
)
assert not context_processor.context_received
async def test_min_words_interruption_strategy_two_words(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
class ContextProcessor(FrameProcessor):
def __init__(self):
super().__init__()
self.context_received = False
async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if isinstance(frame, OpenAILLMContextFrame):
self.context_received = True
elif isinstance(frame, InterruptionFrame):
self.context_received = False
await self.push_frame(frame, direction)
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
context_processor = ContextProcessor()
pipeline = Pipeline([aggregator, context_processor])
frames_to_send = [
BotStartedSpeakingFrame(),
UserStartedSpeakingFrame(),
TranscriptionFrame(text="Can you", user_id="cat", timestamp=""),
SleepFrame(),
UserStoppedSpeakingFrame(),
]
expected_up_frames = [InterruptionTaskFrame]
expected_down_frames = [
BotStartedSpeakingFrame,
UserStartedSpeakingFrame,
InterruptionFrame,
UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES,
]
await run_test(
pipeline,
frames_to_send=frames_to_send,
expected_up_frames=expected_up_frames,
expected_down_frames=expected_down_frames,
pipeline_params=PipelineParams(
interruption_strategies=[MinWordsInterruptionStrategy(min_words=2)]
),
)
self.check_message_content(context, 0, "Can you")
# If the context is not received or it has been cleared by the
# interruption then we have an issue.
assert context_processor.context_received
class BaseTestAssistantContextAggreagator: class BaseTestAssistantContextAggreagator:
CONTEXT_CLASS = None # To be set in subclasses CONTEXT_CLASS = None # To be set in subclasses

View File

@@ -65,7 +65,7 @@ class TestPipeline(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,
ignore_start=False, ignore_start=False,
start_metadata={"foo": "bar"}, pipeline_params=PipelineParams(start_metadata={"foo": "bar"}),
) )
assert "foo" in received_down[-1].metadata assert "foo" in received_down[-1].metadata