tests: add bedrock context aggregator tests
This commit is contained in:
@@ -1 +1 @@
|
|||||||
-e ".[anthropic,google,langchain]"
|
-e ".[anthropic,aws,google,langchain]"
|
||||||
|
|||||||
@@ -40,6 +40,11 @@ from pipecat.services.anthropic.llm import (
|
|||||||
AnthropicLLMContext,
|
AnthropicLLMContext,
|
||||||
AnthropicUserContextAggregator,
|
AnthropicUserContextAggregator,
|
||||||
)
|
)
|
||||||
|
from pipecat.services.aws.llm import (
|
||||||
|
BedrockAssistantContextAggregator,
|
||||||
|
BedrockLLMContext,
|
||||||
|
BedrockUserContextAggregator,
|
||||||
|
)
|
||||||
from pipecat.services.google.llm import (
|
from pipecat.services.google.llm import (
|
||||||
GoogleAssistantContextAggregator,
|
GoogleAssistantContextAggregator,
|
||||||
GoogleLLMContext,
|
GoogleLLMContext,
|
||||||
@@ -669,26 +674,6 @@ class TestLLMUserContextAggregator(BaseTestUserContextAggregator, unittest.Isola
|
|||||||
AGGREGATOR_CLASS = LLMUserContextAggregator
|
AGGREGATOR_CLASS = LLMUserContextAggregator
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# OpenAI
|
|
||||||
#
|
|
||||||
|
|
||||||
|
|
||||||
class TestOpenAIUserContextAggregator(
|
|
||||||
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
|
|
||||||
):
|
|
||||||
CONTEXT_CLASS = OpenAILLMContext
|
|
||||||
AGGREGATOR_CLASS = OpenAIUserContextAggregator
|
|
||||||
|
|
||||||
|
|
||||||
class TestOpenAIAssistantContextAggregator(
|
|
||||||
BaseTestAssistantContextAggreagator, unittest.IsolatedAsyncioTestCase
|
|
||||||
):
|
|
||||||
CONTEXT_CLASS = OpenAILLMContext
|
|
||||||
AGGREGATOR_CLASS = OpenAIAssistantContextAggregator
|
|
||||||
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame]
|
|
||||||
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Anthropic
|
# Anthropic
|
||||||
#
|
#
|
||||||
@@ -724,6 +709,43 @@ class TestAnthropicAssistantContextAggregator(
|
|||||||
assert context.messages[index]["content"][0]["content"] == json.dumps(content)
|
assert context.messages[index]["content"][0]["content"] == json.dumps(content)
|
||||||
|
|
||||||
|
|
||||||
|
#
|
||||||
|
# AWS (Bedrock)
|
||||||
|
#
|
||||||
|
|
||||||
|
|
||||||
|
class TestBedrockUserContextAggregator(
|
||||||
|
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
|
||||||
|
):
|
||||||
|
CONTEXT_CLASS = BedrockLLMContext
|
||||||
|
AGGREGATOR_CLASS = BedrockUserContextAggregator
|
||||||
|
|
||||||
|
def check_message_multi_content(
|
||||||
|
self, context: OpenAILLMContext, content_index: int, index: int, content: str
|
||||||
|
):
|
||||||
|
messages = context.messages[content_index]
|
||||||
|
assert messages["content"][index]["text"] == content
|
||||||
|
|
||||||
|
|
||||||
|
class TestBedrockAssistantContextAggregator(
|
||||||
|
BaseTestAssistantContextAggreagator, unittest.IsolatedAsyncioTestCase
|
||||||
|
):
|
||||||
|
CONTEXT_CLASS = BedrockLLMContext
|
||||||
|
AGGREGATOR_CLASS = BedrockAssistantContextAggregator
|
||||||
|
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame]
|
||||||
|
|
||||||
|
def check_message_multi_content(
|
||||||
|
self, context: OpenAILLMContext, content_index: int, index: int, content: str
|
||||||
|
):
|
||||||
|
messages = context.messages[content_index]
|
||||||
|
assert messages["content"][index]["text"] == content
|
||||||
|
|
||||||
|
def check_function_call_result(self, context: OpenAILLMContext, index: int, content: Any):
|
||||||
|
assert context.messages[index]["content"][0]["toolResult"]["content"][0][
|
||||||
|
"text"
|
||||||
|
] == json.dumps(content)
|
||||||
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Google
|
# Google
|
||||||
#
|
#
|
||||||
@@ -766,3 +788,23 @@ class TestGoogleAssistantContextAggregator(
|
|||||||
def check_function_call_result(self, context: OpenAILLMContext, index: int, content: Any):
|
def check_function_call_result(self, context: OpenAILLMContext, index: int, content: Any):
|
||||||
obj = glm.Content.to_dict(context.messages[index])
|
obj = glm.Content.to_dict(context.messages[index])
|
||||||
assert obj["parts"][0]["function_response"]["response"]["value"] == json.dumps(content)
|
assert obj["parts"][0]["function_response"]["response"]["value"] == json.dumps(content)
|
||||||
|
|
||||||
|
|
||||||
|
#
|
||||||
|
# OpenAI
|
||||||
|
#
|
||||||
|
|
||||||
|
|
||||||
|
class TestOpenAIUserContextAggregator(
|
||||||
|
BaseTestUserContextAggregator, unittest.IsolatedAsyncioTestCase
|
||||||
|
):
|
||||||
|
CONTEXT_CLASS = OpenAILLMContext
|
||||||
|
AGGREGATOR_CLASS = OpenAIUserContextAggregator
|
||||||
|
|
||||||
|
|
||||||
|
class TestOpenAIAssistantContextAggregator(
|
||||||
|
BaseTestAssistantContextAggreagator, unittest.IsolatedAsyncioTestCase
|
||||||
|
):
|
||||||
|
CONTEXT_CLASS = OpenAILLMContext
|
||||||
|
AGGREGATOR_CLASS = OpenAIAssistantContextAggregator
|
||||||
|
EXPECTED_CONTEXT_FRAMES = [OpenAILLMContextFrame, OpenAILLMContextAssistantTimestampFrame]
|
||||||
|
|||||||
Reference in New Issue
Block a user