tests: add OpenAI context aggregator tests

This commit is contained in:
Aleix Conchillo Flaqué
2025-02-12 23:29:45 -08:00
parent 7c815121ea
commit b602e78625

View File

@@ -11,6 +11,7 @@ from pipecat.frames.frames import (
InterimTranscriptionFrame, InterimTranscriptionFrame,
LLMFullResponseEndFrame, LLMFullResponseEndFrame,
LLMFullResponseStartFrame, LLMFullResponseStartFrame,
StartInterruptionFrame,
TextFrame, TextFrame,
TranscriptionFrame, TranscriptionFrame,
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
@@ -24,6 +25,7 @@ from pipecat.processors.aggregators.openai_llm_context import (
OpenAILLMContext, OpenAILLMContext,
OpenAILLMContextFrame, OpenAILLMContextFrame,
) )
from pipecat.services.openai import OpenAIUserContextAggregator
from pipecat.tests.utils import SleepFrame, run_test from pipecat.tests.utils import SleepFrame, run_test
AGGREGATION_TIMEOUT = 0.1 AGGREGATION_TIMEOUT = 0.1
@@ -32,10 +34,16 @@ BOT_INTERRUPTION_TIMEOUT = 0.2
BOT_INTERRUPTION_SLEEP = 0.25 BOT_INTERRUPTION_SLEEP = 0.25
class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase): class BaseTestUserContextAggregator:
CONTEXT_CLASS = None # To be set in subclasses
AGGREGATOR_CLASS = None # To be set in subclasses
async def test_se(self): async def test_se(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [UserStartedSpeakingFrame(), UserStoppedSpeakingFrame()] frames_to_send = [UserStartedSpeakingFrame(), UserStoppedSpeakingFrame()]
expected_down_frames = [UserStartedSpeakingFrame, UserStoppedSpeakingFrame] expected_down_frames = [UserStartedSpeakingFrame, UserStoppedSpeakingFrame]
await run_test( await run_test(
@@ -45,8 +53,11 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
) )
async def test_ste(self): async def test_ste(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
TranscriptionFrame(text="Hello!", user_id="cat", timestamp=""), TranscriptionFrame(text="Hello!", user_id="cat", timestamp=""),
@@ -58,16 +69,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello!" assert context.messages[0]["content"] == "Hello!"
async def test_site(self): async def test_site(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
InterimTranscriptionFrame(text="Hello", user_id="cat", timestamp=""), InterimTranscriptionFrame(text="Hello", user_id="cat", timestamp=""),
@@ -80,16 +94,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat!" assert context.messages[0]["content"] == "Hello Pipecat!"
async def test_st1iest2e(self): async def test_st1iest2e(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
TranscriptionFrame(text="Hello Pipecat! ", user_id="cat", timestamp=""), TranscriptionFrame(text="Hello Pipecat! ", user_id="cat", timestamp=""),
@@ -108,16 +125,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat! How are you?" assert context.messages[0]["content"] == "Hello Pipecat! How are you?"
async def test_siet(self): async def test_siet(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
InterimTranscriptionFrame(text="How ", user_id="cat", timestamp=""), InterimTranscriptionFrame(text="How ", user_id="cat", timestamp=""),
@@ -131,16 +151,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "How are you?" assert context.messages[0]["content"] == "How are you?"
async def test_sieit(self): async def test_sieit(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
InterimTranscriptionFrame(text="How ", user_id="cat", timestamp=""), InterimTranscriptionFrame(text="How ", user_id="cat", timestamp=""),
@@ -155,16 +178,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "How are you?" assert context.messages[0]["content"] == "How are you?"
async def test_set(self): async def test_set(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
UserStoppedSpeakingFrame(), UserStoppedSpeakingFrame(),
@@ -176,16 +202,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "How are you?" assert context.messages[0]["content"] == "How are you?"
async def test_seit(self): async def test_seit(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
UserStoppedSpeakingFrame(), UserStoppedSpeakingFrame(),
@@ -198,16 +227,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "How are you?" assert context.messages[0]["content"] == "How are you?"
async def test_st1et2(self): async def test_st1et2(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
TranscriptionFrame(text="Hello Pipecat!", user_id="cat", timestamp=""), TranscriptionFrame(text="Hello Pipecat!", user_id="cat", timestamp=""),
@@ -222,17 +254,20 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
OpenAILLMContextFrame, OpenAILLMContextFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat!" assert context.messages[0]["content"] == "Hello Pipecat!"
assert received_down[-1].context.messages[1]["content"] == "How are you?" assert context.messages[1]["content"] == "How are you?"
async def test_set1t2(self): async def test_set1t2(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
UserStoppedSpeakingFrame(), UserStoppedSpeakingFrame(),
@@ -245,16 +280,19 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat! How are you?" assert context.messages[0]["content"] == "Hello Pipecat! How are you?"
async def test_siet1it2(self): async def test_siet1it2(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
UserStartedSpeakingFrame(), UserStartedSpeakingFrame(),
InterimTranscriptionFrame(text="Hello ", user_id="cat", timestamp=""), InterimTranscriptionFrame(text="Hello ", user_id="cat", timestamp=""),
@@ -270,33 +308,39 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat! How are you?" assert context.messages[0]["content"] == "Hello Pipecat! How are you?"
async def test_t(self): async def test_t(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
TranscriptionFrame(text="Hello!", user_id="cat", timestamp=""), TranscriptionFrame(text="Hello!", user_id="cat", timestamp=""),
SleepFrame(sleep=AGGREGATION_SLEEP), SleepFrame(sleep=AGGREGATION_SLEEP),
] ]
expected_down_frames = [OpenAILLMContextFrame] expected_down_frames = [OpenAILLMContextFrame]
expected_up_frames = [BotInterruptionFrame] expected_up_frames = [BotInterruptionFrame]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
expected_up_frames=expected_up_frames, expected_up_frames=expected_up_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello!" assert context.messages[0]["content"] == "Hello!"
async def test_it(self): async def test_it(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator(context, aggregation_timeout=AGGREGATION_TIMEOUT) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, aggregation_timeout=AGGREGATION_TIMEOUT)
frames_to_send = [ frames_to_send = [
InterimTranscriptionFrame(text="Hello ", user_id="cat", timestamp=""), InterimTranscriptionFrame(text="Hello ", user_id="cat", timestamp=""),
TranscriptionFrame(text="Hello Pipecat!", user_id="cat", timestamp=""), TranscriptionFrame(text="Hello Pipecat!", user_id="cat", timestamp=""),
@@ -304,17 +348,20 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
] ]
expected_down_frames = [OpenAILLMContextFrame] expected_down_frames = [OpenAILLMContextFrame]
expected_up_frames = [BotInterruptionFrame] expected_up_frames = [BotInterruptionFrame]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
expected_up_frames=expected_up_frames, expected_up_frames=expected_up_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat!" assert context.messages[0]["content"] == "Hello Pipecat!"
async def test_sie_delay_it(self): async def test_sie_delay_it(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMUserContextAggregator( assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(
context, context,
aggregation_timeout=AGGREGATION_TIMEOUT, aggregation_timeout=AGGREGATION_TIMEOUT,
bot_interruption_timeout=BOT_INTERRUPTION_TIMEOUT, bot_interruption_timeout=BOT_INTERRUPTION_TIMEOUT,
@@ -335,19 +382,25 @@ class TestLLMUserContextAggreagator(unittest.IsolatedAsyncioTestCase):
OpenAILLMContextFrame, OpenAILLMContextFrame,
] ]
expected_up_frames = [BotInterruptionFrame] expected_up_frames = [BotInterruptionFrame]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
expected_up_frames=expected_up_frames, expected_up_frames=expected_up_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "How are you?" assert context.messages[0]["content"] == "How are you?"
class TestLLMAssistantContextAggreagator(unittest.IsolatedAsyncioTestCase): class BaseTestAssistantContextAggreagator:
CONTEXT_CLASS = None # To be set in subclasses
AGGREGATOR_CLASS = None # To be set in subclasses
async def test_empty(self): async def test_empty(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMAssistantContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [LLMFullResponseStartFrame(), LLMFullResponseEndFrame()] frames_to_send = [LLMFullResponseStartFrame(), LLMFullResponseEndFrame()]
expected_down_frames = [] expected_down_frames = []
await run_test( await run_test(
@@ -356,25 +409,31 @@ class TestLLMAssistantContextAggreagator(unittest.IsolatedAsyncioTestCase):
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
async def test_single(self): async def test_single_text(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMAssistantContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
LLMFullResponseStartFrame(), LLMFullResponseStartFrame(),
TextFrame(text="Hello Pipecat!"), TextFrame(text="Hello Pipecat!"),
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [OpenAILLMContextFrame] expected_down_frames = [OpenAILLMContextFrame]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat!" assert context.messages[0]["content"] == "Hello Pipecat!"
async def test_multiple(self): async def test_multiple_text(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMAssistantContextAggregator(context, expect_stripped_words=False) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, expect_stripped_words=False)
frames_to_send = [ frames_to_send = [
LLMFullResponseStartFrame(), LLMFullResponseStartFrame(),
TextFrame(text="Hello "), TextFrame(text="Hello "),
@@ -384,16 +443,19 @@ class TestLLMAssistantContextAggreagator(unittest.IsolatedAsyncioTestCase):
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [OpenAILLMContextFrame] expected_down_frames = [OpenAILLMContextFrame]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat. How are you?" assert context.messages[0]["content"] == "Hello Pipecat. How are you?"
async def test_multiple_stripped(self): async def test_multiple_text_stripped(self):
context = OpenAILLMContext() assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
aggregator = LLMAssistantContextAggregator(context) assert self.AGGREGATOR_CLASS is not None, "AGGREGATOR_CLASS must be set in a subclass"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
LLMFullResponseStartFrame(), LLMFullResponseStartFrame(),
TextFrame(text="Hello"), TextFrame(text="Hello"),
@@ -403,9 +465,101 @@ class TestLLMAssistantContextAggreagator(unittest.IsolatedAsyncioTestCase):
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [OpenAILLMContextFrame] expected_down_frames = [OpenAILLMContextFrame]
(received_down, _) = await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames, expected_down_frames=expected_down_frames,
) )
assert received_down[-1].context.messages[0]["content"] == "Hello Pipecat. How are you?" assert context.messages[0]["content"] == "Hello Pipecat. How are you?"
async def test_multiple_llm_responses(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"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, expect_stripped_words=False)
frames_to_send = [
LLMFullResponseStartFrame(),
TextFrame(text="Hello "),
TextFrame(text="Pipecat."),
LLMFullResponseEndFrame(),
LLMFullResponseStartFrame(),
TextFrame(text="How are "),
TextFrame(text="you?"),
LLMFullResponseEndFrame(),
]
expected_down_frames = [OpenAILLMContextFrame, OpenAILLMContextFrame]
await run_test(
aggregator,
frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames,
)
assert context.messages[0]["content"] == "Hello Pipecat."
assert context.messages[1]["content"] == "How are you?"
async def test_multiple_llm_responses_interruption(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"
context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context, expect_stripped_words=False)
frames_to_send = [
LLMFullResponseStartFrame(),
TextFrame(text="Hello "),
TextFrame(text="Pipecat."),
LLMFullResponseEndFrame(),
SleepFrame(AGGREGATION_SLEEP),
StartInterruptionFrame(),
LLMFullResponseStartFrame(),
TextFrame(text="How are "),
TextFrame(text="you?"),
LLMFullResponseEndFrame(),
]
expected_down_frames = [
OpenAILLMContextFrame,
StartInterruptionFrame,
OpenAILLMContextFrame,
]
await run_test(
aggregator,
frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames,
)
assert context.messages[0]["content"] == "Hello Pipecat."
assert context.messages[1]["content"] == "How are you?"
#
# LLMUserContextAggregator, LLMAssistantContextAggregator
#
class TestLLMUserContextAggregator(BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase):
CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = LLMUserContextAggregator
class TestLLMAssistantContextAggregator(
BaseTestAssistantContextAggreagator, unittest.IsolatedAsyncioTestCase
):
CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = LLMAssistantContextAggregator
#
# OpenAI
#
class TestOpenAIUserContextAggregator(
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
):
CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = OpenAIUserContextAggregator
class TestOpenAIAssistantContextAggregator(
BaseTestAssistantContextAggreagator, unittest.IsolatedAsyncioTestCase
):
CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = LLMAssistantContextAggregator