tests: fix langchanin tests
This commit is contained in:
@@ -7,9 +7,9 @@
|
|||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
|
EndFrame,
|
||||||
LLMFullResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
LLMFullResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
StopTaskFrame,
|
|
||||||
TextFrame,
|
TextFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
UserStartedSpeakingFrame,
|
UserStartedSpeakingFrame,
|
||||||
@@ -32,6 +32,7 @@ from langchain_core.language_models import FakeStreamingListLLM
|
|||||||
class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
||||||
class MockProcessor(FrameProcessor):
|
class MockProcessor(FrameProcessor):
|
||||||
def __init__(self, name):
|
def __init__(self, name):
|
||||||
|
super().__init__()
|
||||||
self.name = name
|
self.name = name
|
||||||
self.token: list[str] = []
|
self.token: list[str] = []
|
||||||
# Start collecting tokens when we see the start frame
|
# Start collecting tokens when we see the start frame
|
||||||
@@ -55,13 +56,13 @@ class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
|||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.expected_response = "Hello dear human"
|
self.expected_response = "Hello dear human"
|
||||||
self.fake_llm = FakeStreamingListLLM(responses=[self.expected_response])
|
self.fake_llm = FakeStreamingListLLM(responses=[self.expected_response])
|
||||||
self.mock_proc = self.MockProcessor("token_collector")
|
|
||||||
|
|
||||||
async def test_langchain(self):
|
async def test_langchain(self):
|
||||||
messages = [("system", "Say hello to {name}"), ("human", "{input}")]
|
messages = [("system", "Say hello to {name}"), ("human", "{input}")]
|
||||||
prompt = ChatPromptTemplate.from_messages(messages).partial(name="Thomas")
|
prompt = ChatPromptTemplate.from_messages(messages).partial(name="Thomas")
|
||||||
chain = prompt | self.fake_llm
|
chain = prompt | self.fake_llm
|
||||||
proc = LangchainProcessor(chain=chain)
|
proc = LangchainProcessor(chain=chain)
|
||||||
|
self.mock_proc = self.MockProcessor("token_collector")
|
||||||
|
|
||||||
tma_in = LLMUserResponseAggregator(messages)
|
tma_in = LLMUserResponseAggregator(messages)
|
||||||
tma_out = LLMAssistantResponseAggregator(messages)
|
tma_out = LLMAssistantResponseAggregator(messages)
|
||||||
@@ -81,7 +82,7 @@ class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
|||||||
UserStartedSpeakingFrame(),
|
UserStartedSpeakingFrame(),
|
||||||
TranscriptionFrame(text="Hi World", user_id="user", timestamp="now"),
|
TranscriptionFrame(text="Hi World", user_id="user", timestamp="now"),
|
||||||
UserStoppedSpeakingFrame(),
|
UserStoppedSpeakingFrame(),
|
||||||
StopTaskFrame(),
|
EndFrame(),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user