Merge pull request #2639 from pipecat-ai/aleix/min-words-interruption-unit-test
MinWordsInterruptionStrategy unit test
This commit is contained in:
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user