tests: added LLMAssistantAggregator unit tests
This commit is contained in:
@@ -13,10 +13,17 @@ from pipecat.frames.frames import (
|
|||||||
FunctionCallResultFrame,
|
FunctionCallResultFrame,
|
||||||
FunctionCallsStartedFrame,
|
FunctionCallsStartedFrame,
|
||||||
InterruptionFrame,
|
InterruptionFrame,
|
||||||
|
LLMContextAssistantTimestampFrame,
|
||||||
LLMContextFrame,
|
LLMContextFrame,
|
||||||
|
LLMFullResponseEndFrame,
|
||||||
|
LLMFullResponseStartFrame,
|
||||||
LLMMessagesAppendFrame,
|
LLMMessagesAppendFrame,
|
||||||
LLMMessagesUpdateFrame,
|
LLMMessagesUpdateFrame,
|
||||||
LLMRunFrame,
|
LLMRunFrame,
|
||||||
|
LLMTextFrame,
|
||||||
|
LLMThoughtEndFrame,
|
||||||
|
LLMThoughtStartFrame,
|
||||||
|
LLMThoughtTextFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
UserStartedSpeakingFrame,
|
UserStartedSpeakingFrame,
|
||||||
UserStoppedSpeakingFrame,
|
UserStoppedSpeakingFrame,
|
||||||
@@ -26,6 +33,9 @@ from pipecat.frames.frames import (
|
|||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.processors.aggregators.llm_context import LLMContext
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
from pipecat.processors.aggregators.llm_response_universal import (
|
from pipecat.processors.aggregators.llm_response_universal import (
|
||||||
|
AssistantThoughtMessage,
|
||||||
|
AssistantTurnStoppedMessage,
|
||||||
|
LLMAssistantAggregator,
|
||||||
LLMUserAggregator,
|
LLMUserAggregator,
|
||||||
LLMUserAggregatorParams,
|
LLMUserAggregatorParams,
|
||||||
)
|
)
|
||||||
@@ -333,3 +343,172 @@ class TestLLMUserAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
# The user mute strategies should have muted the user.
|
# The user mute strategies should have muted the user.
|
||||||
self.assertFalse(user_turn)
|
self.assertFalse(user_turn)
|
||||||
|
|
||||||
|
|
||||||
|
class TestLLMAssistantAggregator(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def test_empty(self):
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
aggregator = LLMAssistantAggregator(context)
|
||||||
|
|
||||||
|
should_start = None
|
||||||
|
should_stop = None
|
||||||
|
stop_message = None
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_started")
|
||||||
|
async def on_assistant_turn_started(aggregator):
|
||||||
|
nonlocal should_start
|
||||||
|
should_start = True
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_stopped")
|
||||||
|
async def on_assistant_turn_stopped(aggregator, message: AssistantTurnStoppedMessage):
|
||||||
|
nonlocal should_stop, stop_message
|
||||||
|
should_stop = True
|
||||||
|
stop_message = message
|
||||||
|
|
||||||
|
frames_to_send = [LLMFullResponseStartFrame(), LLMFullResponseEndFrame()]
|
||||||
|
await run_test(aggregator, frames_to_send=frames_to_send)
|
||||||
|
self.assertTrue(should_start)
|
||||||
|
self.assertIsNone(should_stop)
|
||||||
|
self.assertIsNone(stop_message)
|
||||||
|
|
||||||
|
async def test_simple(self):
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
aggregator = LLMAssistantAggregator(context)
|
||||||
|
|
||||||
|
should_start = None
|
||||||
|
should_stop = None
|
||||||
|
stop_message = None
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_started")
|
||||||
|
async def on_assistant_turn_started(aggregator):
|
||||||
|
nonlocal should_start
|
||||||
|
should_start = True
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_stopped")
|
||||||
|
async def on_assistant_turn_stopped(aggregator, message: AssistantTurnStoppedMessage):
|
||||||
|
nonlocal should_stop, stop_message
|
||||||
|
should_stop = True
|
||||||
|
stop_message = message
|
||||||
|
|
||||||
|
frames_to_send = [
|
||||||
|
LLMFullResponseStartFrame(),
|
||||||
|
LLMTextFrame("Hello from Pipecat!"),
|
||||||
|
LLMFullResponseEndFrame(),
|
||||||
|
]
|
||||||
|
expected_down_frames = [LLMContextFrame, LLMContextAssistantTimestampFrame]
|
||||||
|
await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
self.assertTrue(should_start)
|
||||||
|
self.assertTrue(should_stop)
|
||||||
|
self.assertEqual(stop_message.content, "Hello from Pipecat!")
|
||||||
|
|
||||||
|
async def test_multiple(self):
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
aggregator = LLMAssistantAggregator(context)
|
||||||
|
|
||||||
|
should_start = None
|
||||||
|
should_stop = None
|
||||||
|
stop_message = None
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_started")
|
||||||
|
async def on_assistant_turn_started(aggregator):
|
||||||
|
nonlocal should_start
|
||||||
|
should_start = True
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_stopped")
|
||||||
|
async def on_assistant_turn_stopped(aggregator, message: AssistantTurnStoppedMessage):
|
||||||
|
nonlocal should_stop, stop_message
|
||||||
|
should_stop = True
|
||||||
|
stop_message = message
|
||||||
|
|
||||||
|
frames_to_send = [
|
||||||
|
LLMFullResponseStartFrame(),
|
||||||
|
LLMTextFrame("Hello "),
|
||||||
|
LLMTextFrame("from "),
|
||||||
|
LLMTextFrame("Pipecat!"),
|
||||||
|
LLMFullResponseEndFrame(),
|
||||||
|
]
|
||||||
|
expected_down_frames = [LLMContextFrame, LLMContextAssistantTimestampFrame]
|
||||||
|
await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
self.assertTrue(should_start)
|
||||||
|
self.assertTrue(should_stop)
|
||||||
|
self.assertEqual(stop_message.content, "Hello from Pipecat!")
|
||||||
|
|
||||||
|
async def test_interruption(self):
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
aggregator = LLMAssistantAggregator(context)
|
||||||
|
|
||||||
|
should_start = 0
|
||||||
|
should_stop = 0
|
||||||
|
stop_messages = []
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_started")
|
||||||
|
async def on_assistant_turn_started(aggregator):
|
||||||
|
nonlocal should_start
|
||||||
|
should_start += 1
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_stopped")
|
||||||
|
async def on_assistant_turn_stopped(aggregator, message: AssistantTurnStoppedMessage):
|
||||||
|
nonlocal should_stop, stop_messages
|
||||||
|
should_stop += 1
|
||||||
|
stop_messages.append(message)
|
||||||
|
|
||||||
|
frames_to_send = [
|
||||||
|
LLMFullResponseStartFrame(),
|
||||||
|
LLMTextFrame("Hello "),
|
||||||
|
SleepFrame(),
|
||||||
|
InterruptionFrame(),
|
||||||
|
LLMFullResponseStartFrame(),
|
||||||
|
LLMTextFrame("Hello "),
|
||||||
|
LLMTextFrame("there!"),
|
||||||
|
LLMFullResponseEndFrame(),
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
LLMContextFrame,
|
||||||
|
LLMContextAssistantTimestampFrame,
|
||||||
|
InterruptionFrame,
|
||||||
|
LLMContextFrame,
|
||||||
|
LLMContextAssistantTimestampFrame,
|
||||||
|
]
|
||||||
|
await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
self.assertEqual(should_start, 2)
|
||||||
|
self.assertEqual(should_stop, 2)
|
||||||
|
self.assertEqual(stop_messages[0].content, "Hello")
|
||||||
|
self.assertEqual(stop_messages[1].content, "Hello there!")
|
||||||
|
|
||||||
|
async def test_thought(self):
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
aggregator = LLMAssistantAggregator(context)
|
||||||
|
|
||||||
|
thought_message = None
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_thought")
|
||||||
|
async def on_assistant_thought(aggregator, message: AssistantThoughtMessage):
|
||||||
|
nonlocal thought_message
|
||||||
|
thought_message = message
|
||||||
|
|
||||||
|
frames_to_send = [
|
||||||
|
LLMFullResponseStartFrame(),
|
||||||
|
LLMThoughtStartFrame(),
|
||||||
|
LLMThoughtTextFrame(text="I'm thinking!"),
|
||||||
|
LLMThoughtEndFrame(),
|
||||||
|
LLMFullResponseEndFrame(),
|
||||||
|
]
|
||||||
|
await run_test(aggregator, frames_to_send=frames_to_send)
|
||||||
|
self.assertEqual(thought_message.content, "I'm thinking!")
|
||||||
|
|||||||
Reference in New Issue
Block a user