tests: add interruption strategies context ordering tests

This commit is contained in:
Aleix Conchillo Flaqué
2025-10-12 09:53:41 -07:00
parent 234aae3091
commit e4212fb3c0

View File

@@ -4,6 +4,7 @@
# SPDX-License-Identifier: BSD 2-Clause License # SPDX-License-Identifier: BSD 2-Clause License
# #
import asyncio
import json import json
import unittest import unittest
from typing import Any, Optional from typing import Any, Optional
@@ -30,17 +31,23 @@ from pipecat.frames.frames import (
SpeechControlParamsFrame, SpeechControlParamsFrame,
TextFrame, TextFrame,
TranscriptionFrame, TranscriptionFrame,
TTSTextFrame,
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
) )
from pipecat.pipeline.pipeline import Pipeline from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.task import PipelineParams from pipecat.pipeline.task import PipelineParams
from pipecat.processors.aggregators.llm_context import LLMContext
from pipecat.processors.aggregators.llm_response import ( from pipecat.processors.aggregators.llm_response import (
LLMAssistantAggregatorParams, LLMAssistantAggregatorParams,
LLMAssistantContextAggregator,
LLMUserAggregatorParams, LLMUserAggregatorParams,
LLMUserContextAggregator, LLMUserContextAggregator,
) )
from pipecat.processors.aggregators.llm_response_universal import LLMAssistantAggregator from pipecat.processors.aggregators.llm_response_universal import (
LLMAssistantAggregator,
LLMUserAggregator,
)
from pipecat.processors.aggregators.openai_llm_context import ( from pipecat.processors.aggregators.openai_llm_context import (
OpenAILLMContext, OpenAILLMContext,
OpenAILLMContextFrame, OpenAILLMContextFrame,
@@ -73,8 +80,11 @@ AGGREGATION_SLEEP = 0.15
class BaseTestUserContextAggregator: class BaseTestUserContextAggregator:
CONTEXT_CLASS = None # To be set in subclasses CONTEXT_CLASS = None # To be set in subclasses
AGGREGATOR_CLASS = None # To be set in subclasses CONTEXT_FRAME_CLASS = None # To be set in subclasses
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame] USER_AGGREGATOR_CLASS = None # To be set in subclasses
USER_EXPECTED_CONTEXT_FRAMES = None
ASSISTANT_AGGREGATOR_CLASS = None # To be set in subclasses
ASSISTANT_EXPECTED_CONTEXT_FRAMES = None # To be set in subclasses
def check_message_content(self, context: OpenAILLMContext, index: int, content: str): def check_message_content(self, context: OpenAILLMContext, index: int, content: str):
assert context.messages[index]["content"] == content assert context.messages[index]["content"] == content
@@ -86,10 +96,12 @@ class BaseTestUserContextAggregator:
async def test_se(self): async def test_se(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.USER_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(
@@ -100,10 +112,12 @@ class BaseTestUserContextAggregator:
async def test_ste(self): async def test_ste(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.USER_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=""),
@@ -112,7 +126,7 @@ class BaseTestUserContextAggregator:
] ]
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
] ]
await run_test( await run_test(
@@ -124,10 +138,12 @@ class BaseTestUserContextAggregator:
async def test_site(self): async def test_site(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.USER_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=""),
@@ -137,7 +153,7 @@ class BaseTestUserContextAggregator:
] ]
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
] ]
await run_test( await run_test(
@@ -149,10 +165,12 @@ class BaseTestUserContextAggregator:
async def test_st1iest2e(self): async def test_st1iest2e(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.USER_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=""),
@@ -168,7 +186,7 @@ class BaseTestUserContextAggregator:
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
] ]
await run_test( await run_test(
@@ -180,10 +198,12 @@ class BaseTestUserContextAggregator:
async def test_siet(self): async def test_siet(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -197,7 +217,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -208,10 +228,12 @@ class BaseTestUserContextAggregator:
async def test_sieit(self): async def test_sieit(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -226,7 +248,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -237,10 +259,12 @@ class BaseTestUserContextAggregator:
async def test_set(self): async def test_set(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -252,7 +276,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -263,10 +287,12 @@ class BaseTestUserContextAggregator:
async def test_seit(self): async def test_seit(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -279,7 +305,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -290,10 +316,12 @@ class BaseTestUserContextAggregator:
async def test_st1et2(self): async def test_st1et2(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -308,9 +336,9 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
SpeechControlParamsFrame, SpeechControlParamsFrame,
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -322,10 +350,12 @@ class BaseTestUserContextAggregator:
async def test_set1t2(self): async def test_set1t2(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -338,7 +368,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -349,10 +379,12 @@ class BaseTestUserContextAggregator:
async def test_siet1it2(self): async def test_siet1it2(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -368,7 +400,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -379,10 +411,12 @@ class BaseTestUserContextAggregator:
async def test_t(self): async def test_t(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context context
) # No aggregation timeout; this tests VAD emulation ) # No aggregation timeout; this tests VAD emulation
@@ -393,7 +427,7 @@ class BaseTestUserContextAggregator:
] ]
expected_down_frames = [ expected_down_frames = [
SpeechControlParamsFrame, SpeechControlParamsFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
expected_up_frames = [EmulateUserStartedSpeakingFrame, EmulateUserStoppedSpeakingFrame] expected_up_frames = [EmulateUserStartedSpeakingFrame, EmulateUserStoppedSpeakingFrame]
@@ -407,10 +441,12 @@ class BaseTestUserContextAggregator:
async def test_t_with_turn_analyzer(self): async def test_t_with_turn_analyzer(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(turn_emulated_vad_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(turn_emulated_vad_timeout=AGGREGATION_TIMEOUT)
) )
@@ -424,7 +460,7 @@ class BaseTestUserContextAggregator:
] ]
expected_down_frames = [ expected_down_frames = [
SpeechControlParamsFrame, SpeechControlParamsFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
expected_up_frames = [EmulateUserStartedSpeakingFrame, EmulateUserStoppedSpeakingFrame] expected_up_frames = [EmulateUserStartedSpeakingFrame, EmulateUserStoppedSpeakingFrame]
@@ -438,10 +474,12 @@ class BaseTestUserContextAggregator:
async def test_it(self): async def test_it(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context context
) # No aggregation timeout; this tests VAD emulation ) # No aggregation timeout; this tests VAD emulation
frames_to_send = [ frames_to_send = [
@@ -451,7 +489,7 @@ class BaseTestUserContextAggregator:
TranscriptionFrame(text="Hello Pipecat!", user_id="cat", timestamp=""), TranscriptionFrame(text="Hello Pipecat!", user_id="cat", timestamp=""),
SleepFrame(sleep=AGGREGATION_SLEEP), SleepFrame(sleep=AGGREGATION_SLEEP),
] ]
expected_down_frames = [SpeechControlParamsFrame, *self.EXPECTED_CONTEXT_FRAMES] expected_down_frames = [SpeechControlParamsFrame, *self.USER_EXPECTED_CONTEXT_FRAMES]
expected_up_frames = [EmulateUserStartedSpeakingFrame, EmulateUserStoppedSpeakingFrame] expected_up_frames = [EmulateUserStartedSpeakingFrame, EmulateUserStoppedSpeakingFrame]
await run_test( await run_test(
aggregator, aggregator,
@@ -463,10 +501,12 @@ class BaseTestUserContextAggregator:
async def test_sie_delay_it(self): async def test_sie_delay_it(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT) context, params=LLMUserAggregatorParams(aggregation_timeout=AGGREGATION_TIMEOUT)
) )
frames_to_send = [ frames_to_send = [
@@ -482,7 +522,7 @@ class BaseTestUserContextAggregator:
expected_down_frames = [ expected_down_frames = [
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -493,7 +533,12 @@ class BaseTestUserContextAggregator:
async def test_min_words_interruption_strategy_one_word(self): 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.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" assert self.CONTEXT_FRAME_CLASS is not None, "CONTEXT_FRAME_CLASS must be set in a subclass"
assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
CONTEXT_FRAME_CLASS = self.CONTEXT_FRAME_CLASS
class ContextProcessor(FrameProcessor): class ContextProcessor(FrameProcessor):
def __init__(self): def __init__(self):
@@ -503,13 +548,13 @@ class BaseTestUserContextAggregator:
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
if isinstance(frame, OpenAILLMContextFrame): if isinstance(frame, CONTEXT_FRAME_CLASS):
self.context_received = True self.context_received = True
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.USER_AGGREGATOR_CLASS(context)
context_processor = ContextProcessor() context_processor = ContextProcessor()
pipeline = Pipeline([aggregator, context_processor]) pipeline = Pipeline([aggregator, context_processor])
@@ -537,7 +582,12 @@ class BaseTestUserContextAggregator:
async def test_min_words_interruption_strategy_two_words(self): 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.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" assert self.CONTEXT_FRAME_CLASS is not None, "CONTEXT_FRAME_CLASS must be set in a subclass"
assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
CONTEXT_FRAME_CLASS = self.CONTEXT_FRAME_CLASS
class ContextProcessor(FrameProcessor): class ContextProcessor(FrameProcessor):
def __init__(self): def __init__(self):
@@ -547,7 +597,7 @@ class BaseTestUserContextAggregator:
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
if isinstance(frame, OpenAILLMContextFrame): if isinstance(frame, CONTEXT_FRAME_CLASS):
self.context_received = True self.context_received = True
elif isinstance(frame, InterruptionFrame): elif isinstance(frame, InterruptionFrame):
self.context_received = False self.context_received = False
@@ -555,7 +605,7 @@ class BaseTestUserContextAggregator:
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.USER_AGGREGATOR_CLASS(context)
context_processor = ContextProcessor() context_processor = ContextProcessor()
pipeline = Pipeline([aggregator, context_processor]) pipeline = Pipeline([aggregator, context_processor])
@@ -572,7 +622,7 @@ class BaseTestUserContextAggregator:
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
InterruptionFrame, InterruptionFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.USER_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
pipeline, pipeline,
@@ -588,11 +638,77 @@ class BaseTestUserContextAggregator:
# interruption then we have an issue. # interruption then we have an issue.
assert context_processor.context_received assert context_processor.context_received
async def test_interruption_strategy_context_order(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass"
assert self.USER_AGGREGATOR_CLASS is not None, (
"USER_AGGREGATOR_CLASS must be set in a subclass"
)
assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
class DelayedProcessor(FrameProcessor):
"""Force a delay in interruption frames.
This might give time to the assistant aggregator to update the
context before the user aggregator (which shouldn't really happen)
and reveal any issues in context ordering.
"""
def __init__(self):
super().__init__()
async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if isinstance(frame, InterruptionFrame):
await asyncio.sleep(0.3)
await self.push_frame(frame, direction)
context = self.CONTEXT_CLASS()
user_aggregator = self.USER_AGGREGATOR_CLASS(
context, params=LLMUserAggregatorParams(aggregation_timeout=1.0)
)
assistant_aggregator = self.ASSISTANT_AGGREGATOR_CLASS(context)
pipeline = Pipeline([user_aggregator, DelayedProcessor(), assistant_aggregator])
frames_to_send = [
# Aggregate assistant content.
BotStartedSpeakingFrame(),
LLMFullResponseStartFrame(),
TTSTextFrame(text="Hello, I'm your assistant"),
SleepFrame(),
# Interrupt the bot. Assistant content should be added first to the
# context, followed by user content.
UserStartedSpeakingFrame(),
TranscriptionFrame(text="Can you tell me", user_id="cat", timestamp=""),
SleepFrame(),
UserStoppedSpeakingFrame(),
]
expected_down_frames = [
BotStartedSpeakingFrame,
UserStartedSpeakingFrame,
*self.ASSISTANT_EXPECTED_CONTEXT_FRAMES,
InterruptionFrame,
UserStoppedSpeakingFrame,
*self.USER_EXPECTED_CONTEXT_FRAMES,
]
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)]
),
)
self.check_message_content(context, -1, "Can you tell me")
self.check_message_content(context, -2, "Hello, I'm your assistant")
class BaseTestAssistantContextAggregator: class BaseTestAssistantContextAggregator:
CONTEXT_CLASS = None # To be set in subclasses CONTEXT_CLASS = None # To be set in subclasses
AGGREGATOR_CLASS = None # To be set in subclasses USER_AGGREGATOR_CLASS = None # To be set in subclasses
EXPECTED_CONTEXT_FRAMES = None # To be set in subclasses ASSISTANT_AGGREGATOR_CLASS = None # To be set in subclasses
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [] # To be set in subclasses
def create_assistant_aggregator_params( def create_assistant_aggregator_params(
self, **kwargs self, **kwargs
@@ -612,10 +728,12 @@ class BaseTestAssistantContextAggregator:
async def test_empty(self): async def test_empty(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.ASSISTANT_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(
@@ -626,16 +744,18 @@ class BaseTestAssistantContextAggregator:
async def test_single_text(self): async def test_single_text(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.ASSISTANT_AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
LLMFullResponseStartFrame(), LLMFullResponseStartFrame(),
TextFrame(text="Hello Pipecat!"), TextFrame(text="Hello Pipecat!"),
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [*self.EXPECTED_CONTEXT_FRAMES] expected_down_frames = [*self.ASSISTANT_EXPECTED_CONTEXT_FRAMES]
await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
@@ -645,10 +765,12 @@ class BaseTestAssistantContextAggregator:
async def test_multiple_text(self): async def test_multiple_text(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.ASSISTANT_AGGREGATOR_CLASS(
context, params=self.create_assistant_aggregator_params(expect_stripped_words=False) context, params=self.create_assistant_aggregator_params(expect_stripped_words=False)
) )
frames_to_send = [ frames_to_send = [
@@ -659,7 +781,7 @@ class BaseTestAssistantContextAggregator:
TextFrame(text="you?"), TextFrame(text="you?"),
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [*self.EXPECTED_CONTEXT_FRAMES] expected_down_frames = [*self.ASSISTANT_EXPECTED_CONTEXT_FRAMES]
await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
@@ -669,10 +791,12 @@ class BaseTestAssistantContextAggregator:
async def test_multiple_text_stripped(self): async def test_multiple_text_stripped(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.ASSISTANT_AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
LLMFullResponseStartFrame(), LLMFullResponseStartFrame(),
TextFrame(text="Hello"), TextFrame(text="Hello"),
@@ -681,7 +805,7 @@ class BaseTestAssistantContextAggregator:
TextFrame(text="you?"), TextFrame(text="you?"),
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [*self.EXPECTED_CONTEXT_FRAMES] expected_down_frames = [*self.ASSISTANT_EXPECTED_CONTEXT_FRAMES]
await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
@@ -691,10 +815,12 @@ class BaseTestAssistantContextAggregator:
async def test_multiple_llm_responses(self): async def test_multiple_llm_responses(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.ASSISTANT_AGGREGATOR_CLASS(
context, params=self.create_assistant_aggregator_params(expect_stripped_words=False) context, params=self.create_assistant_aggregator_params(expect_stripped_words=False)
) )
frames_to_send = [ frames_to_send = [
@@ -707,7 +833,10 @@ class BaseTestAssistantContextAggregator:
TextFrame(text="you?"), TextFrame(text="you?"),
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [*self.EXPECTED_CONTEXT_FRAMES, *self.EXPECTED_CONTEXT_FRAMES] expected_down_frames = [
*self.ASSISTANT_EXPECTED_CONTEXT_FRAMES,
*self.ASSISTANT_EXPECTED_CONTEXT_FRAMES,
]
await run_test( await run_test(
aggregator, aggregator,
frames_to_send=frames_to_send, frames_to_send=frames_to_send,
@@ -718,10 +847,12 @@ class BaseTestAssistantContextAggregator:
async def test_multiple_llm_responses_interruption(self): 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.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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS( aggregator = self.ASSISTANT_AGGREGATOR_CLASS(
context, params=self.create_assistant_aggregator_params(expect_stripped_words=False) context, params=self.create_assistant_aggregator_params(expect_stripped_words=False)
) )
frames_to_send = [ frames_to_send = [
@@ -737,9 +868,9 @@ class BaseTestAssistantContextAggregator:
LLMFullResponseEndFrame(), LLMFullResponseEndFrame(),
] ]
expected_down_frames = [ expected_down_frames = [
*self.EXPECTED_CONTEXT_FRAMES, *self.ASSISTANT_EXPECTED_CONTEXT_FRAMES,
InterruptionFrame, InterruptionFrame,
*self.EXPECTED_CONTEXT_FRAMES, *self.ASSISTANT_EXPECTED_CONTEXT_FRAMES,
] ]
await run_test( await run_test(
aggregator, aggregator,
@@ -751,10 +882,12 @@ class BaseTestAssistantContextAggregator:
async def test_function_call(self): async def test_function_call(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.ASSISTANT_AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
FunctionCallInProgressFrame( FunctionCallInProgressFrame(
function_name="get_weather", function_name="get_weather",
@@ -780,7 +913,9 @@ class BaseTestAssistantContextAggregator:
async def test_function_call_on_context_updated(self): async def test_function_call_on_context_updated(self):
assert self.CONTEXT_CLASS is not None, "CONTEXT_CLASS must be set in a subclass" 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" assert self.ASSISTANT_AGGREGATOR_CLASS is not None, (
"ASSISTANT_AGGREGATOR_CLASS must be set in a subclass"
)
context_updated = False context_updated = False
@@ -789,7 +924,7 @@ class BaseTestAssistantContextAggregator:
context_updated = True context_updated = True
context = self.CONTEXT_CLASS() context = self.CONTEXT_CLASS()
aggregator = self.AGGREGATOR_CLASS(context) aggregator = self.ASSISTANT_AGGREGATOR_CLASS(context)
frames_to_send = [ frames_to_send = [
FunctionCallInProgressFrame( FunctionCallInProgressFrame(
function_name="get_weather", function_name="get_weather",
@@ -824,7 +959,14 @@ class BaseTestAssistantContextAggregator:
class TestLLMUserContextAggregator(BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase): class TestLLMUserContextAggregator(BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase):
CONTEXT_CLASS = OpenAILLMContext CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = LLMUserContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
USER_AGGREGATOR_CLASS = LLMUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = LLMAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
# #
@@ -836,7 +978,14 @@ class TestAnthropicUserContextAggregator(
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = AnthropicLLMContext CONTEXT_CLASS = AnthropicLLMContext
AGGREGATOR_CLASS = AnthropicUserContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
USER_AGGREGATOR_CLASS = AnthropicUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = AnthropicAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
def check_message_multi_content( def check_message_multi_content(
self, context: OpenAILLMContext, content_index: int, index: int, content: str self, context: OpenAILLMContext, content_index: int, index: int, content: str
@@ -849,8 +998,14 @@ class TestAnthropicAssistantContextAggregator(
BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = AnthropicLLMContext CONTEXT_CLASS = AnthropicLLMContext
AGGREGATOR_CLASS = AnthropicAssistantContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame] USER_AGGREGATOR_CLASS = AnthropicUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = AnthropicAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
def check_message_multi_content( def check_message_multi_content(
self, context: OpenAILLMContext, content_index: int, index: int, content: str self, context: OpenAILLMContext, content_index: int, index: int, content: str
@@ -871,7 +1026,14 @@ class TestAWSBedrockUserContextAggregator(
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = AWSBedrockLLMContext CONTEXT_CLASS = AWSBedrockLLMContext
AGGREGATOR_CLASS = AWSBedrockUserContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
USER_AGGREGATOR_CLASS = AWSBedrockUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = AWSBedrockAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
def check_message_multi_content( def check_message_multi_content(
self, context: OpenAILLMContext, content_index: int, index: int, content: str self, context: OpenAILLMContext, content_index: int, index: int, content: str
@@ -884,8 +1046,14 @@ class TestAWSBedrockAssistantContextAggregator(
BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = AWSBedrockLLMContext CONTEXT_CLASS = AWSBedrockLLMContext
AGGREGATOR_CLASS = AWSBedrockAssistantContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame] USER_AGGREGATOR_CLASS = AWSBedrockUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = AWSBedrockAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
def check_message_multi_content( def check_message_multi_content(
self, context: OpenAILLMContext, content_index: int, index: int, content: str self, context: OpenAILLMContext, content_index: int, index: int, content: str
@@ -908,7 +1076,14 @@ class TestGoogleUserContextAggregator(
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = GoogleLLMContext CONTEXT_CLASS = GoogleLLMContext
AGGREGATOR_CLASS = GoogleUserContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
USER_AGGREGATOR_CLASS = GoogleUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = GoogleAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
def check_message_content(self, context: OpenAILLMContext, index: int, content: str): def check_message_content(self, context: OpenAILLMContext, index: int, content: str):
obj = context.messages[index].to_json_dict() obj = context.messages[index].to_json_dict()
@@ -925,8 +1100,14 @@ class TestGoogleAssistantContextAggregator(
BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = GoogleLLMContext CONTEXT_CLASS = GoogleLLMContext
AGGREGATOR_CLASS = GoogleAssistantContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame] USER_AGGREGATOR_CLASS = GoogleUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = GoogleAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
def check_message_content(self, context: OpenAILLMContext, index: int, content: str): def check_message_content(self, context: OpenAILLMContext, index: int, content: str):
obj = context.messages[index].to_json_dict() obj = context.messages[index].to_json_dict()
@@ -952,26 +1133,53 @@ class TestOpenAIUserContextAggregator(
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = OpenAILLMContext CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = OpenAIUserContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
USER_AGGREGATOR_CLASS = OpenAIUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = OpenAIAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
class TestOpenAIAssistantContextAggregator( class TestOpenAIAssistantContextAggregator(
BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = OpenAILLMContext CONTEXT_CLASS = OpenAILLMContext
AGGREGATOR_CLASS = OpenAIAssistantContextAggregator CONTEXT_FRAME_CLASS = OpenAILLMContextFrame
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame] USER_AGGREGATOR_CLASS = OpenAIUserContextAggregator
USER_EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = OpenAIAssistantContextAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [
OpenAILLMContextFrame,
OpenAILLMContextAssistantTimestampFrame,
]
# #
# Universal # Universal
# #
class TestLLMUserAggregator(BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase):
CONTEXT_CLASS = LLMContext
CONTEXT_FRAME_CLASS = LLMContextFrame
USER_AGGREGATOR_CLASS = LLMUserAggregator
USER_EXPECTED_CONTEXT_FRAMES = [LLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = LLMAssistantAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [LLMContextFrame, LLMContextAssistantTimestampFrame]
class TestLLMAssistantAggregator( class TestLLMAssistantAggregator(
BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase BaseTestAssistantContextAggregator, unittest.IsolatedAsyncioTestCase
): ):
CONTEXT_CLASS = OpenAILLMContext CONTEXT_CLASS = LLMContext
AGGREGATOR_CLASS = LLMAssistantAggregator CONTEXT_FRAME_CLASS = LLMContextFrame
EXPECTED_CONTEXT_FRAMES = [LLMContextFrame, LLMContextAssistantTimestampFrame] USER_AGGREGATOR_CLASS = LLMUserAggregator
USER_EXPECTED_CONTEXT_FRAMES = [LLMContextFrame]
ASSISTANT_AGGREGATOR_CLASS = LLMAssistantAggregator
ASSISTANT_EXPECTED_CONTEXT_FRAMES = [LLMContextFrame, LLMContextAssistantTimestampFrame]
# Override to remove 'expect_stripped_words' parameter, which is deprecated # Override to remove 'expect_stripped_words' parameter, which is deprecated
# for LLMAssistantAggregator # for LLMAssistantAggregator