tests pass
This commit is contained in:
@@ -18,12 +18,13 @@ class Message:
|
|||||||
|
|
||||||
class MessageHandler:
|
class MessageHandler:
|
||||||
def __init__(self, intro):
|
def __init__(self, intro):
|
||||||
self.messages = [Message("system", time.time(), intro)]
|
self.messages: list[Message] = [Message("system", time.time(), intro)]
|
||||||
self.last_user_message_idx = None
|
self.last_user_message_idx:int | None = None
|
||||||
|
self.finalized_user_message_idx: int | None = None
|
||||||
|
|
||||||
def add_user_message(self, message):
|
def add_user_message(self, message) -> None:
|
||||||
if (self.last_user_message_idx is not None and self.last_user_message_idx != self.finalized_user_message_idx):
|
if self.last_user_message_idx is not None and self.last_user_message_idx != self.finalized_user_message_idx:
|
||||||
previous_message = self.messages[self.last_user_message_idx].message
|
previous_message: str = self.messages[self.last_user_message_idx].message
|
||||||
self.messages[self.last_user_message_idx] = Message(
|
self.messages[self.last_user_message_idx] = Message(
|
||||||
"user", time.time(), ' '.join([previous_message, message])
|
"user", time.time(), ' '.join([previous_message, message])
|
||||||
)
|
)
|
||||||
@@ -33,22 +34,22 @@ class MessageHandler:
|
|||||||
|
|
||||||
self.last_user_message_idx = len(self.messages) - 1
|
self.last_user_message_idx = len(self.messages) - 1
|
||||||
|
|
||||||
def add_assistant_message(self, message):
|
def add_assistant_message(self, message) -> None:
|
||||||
if self.messages[-1].type == "assistant":
|
if self.messages[-1].type == "assistant":
|
||||||
self.messages[-1].message += " " + message
|
self.messages[-1].message += " " + message
|
||||||
else:
|
else:
|
||||||
self.messages.append(Message("assistant", time.time(), message))
|
self.messages.append(Message("assistant", time.time(), message))
|
||||||
|
|
||||||
def add_assistant_messages(self, messages):
|
def add_assistant_messages(self, messages) -> None:
|
||||||
self.messages.append(Message("assistant", time.time(), " ".join(messages)))
|
self.messages.append(Message("assistant", time.time(), " ".join(messages)))
|
||||||
|
|
||||||
def get_llm_messages(self):
|
def get_llm_messages(self) -> list[dict[str, str]]:
|
||||||
return [{"role": m.type, "content": m.message} for m in self.messages]
|
return [{"role": m.type, "content": m.message} for m in self.messages]
|
||||||
|
|
||||||
def finalize_user_message(self):
|
def finalize_user_message(self) -> None:
|
||||||
pass
|
self.finalized_user_message_idx = self.last_user_message_idx
|
||||||
|
|
||||||
def shutdown(self):
|
def shutdown(self) -> None:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class IndexingMessageHandler(MessageHandler):
|
class IndexingMessageHandler(MessageHandler):
|
||||||
@@ -66,8 +67,6 @@ class IndexingMessageHandler(MessageHandler):
|
|||||||
self.index_writer_thread = Thread(target=self.indexer_writer, daemon=True)
|
self.index_writer_thread = Thread(target=self.indexer_writer, daemon=True)
|
||||||
self.index_writer_thread.start()
|
self.index_writer_thread.start()
|
||||||
|
|
||||||
self.finalized_user_message_idx = None
|
|
||||||
|
|
||||||
self.logger = logging.getLogger("bot-instance")
|
self.logger = logging.getLogger("bot-instance")
|
||||||
|
|
||||||
def shutdown(self):
|
def shutdown(self):
|
||||||
@@ -75,7 +74,7 @@ class IndexingMessageHandler(MessageHandler):
|
|||||||
self.index_message_queue.put(None)
|
self.index_message_queue.put(None)
|
||||||
self.index_writer_thread.join()
|
self.index_writer_thread.join()
|
||||||
|
|
||||||
def indexer_writer(self):
|
def indexer_writer(self) -> None:
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
message_idx = self.index_message_queue.get()
|
message_idx = self.index_message_queue.get()
|
||||||
@@ -123,7 +122,7 @@ class IndexingMessageHandler(MessageHandler):
|
|||||||
return user_message
|
return user_message
|
||||||
|
|
||||||
def finalize_user_message(self):
|
def finalize_user_message(self):
|
||||||
self.finalized_user_message_idx = self.last_user_message_idx
|
super().finalize_user_message()
|
||||||
self.write_messages_to_index()
|
self.write_messages_to_index()
|
||||||
|
|
||||||
def write_messages_to_index(self):
|
def write_messages_to_index(self):
|
||||||
|
|||||||
@@ -26,7 +26,6 @@ class MockLLMService(LLMService):
|
|||||||
time.sleep(0.1)
|
time.sleep(0.1)
|
||||||
yield i
|
yield i
|
||||||
|
|
||||||
|
|
||||||
class MockImageService(ImageGenService):
|
class MockImageService(ImageGenService):
|
||||||
def run_image_gen(self, sentence) -> None:
|
def run_image_gen(self, sentence) -> None:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import unittest
|
|||||||
from unittest.mock import MagicMock, call
|
from unittest.mock import MagicMock, call
|
||||||
|
|
||||||
from message_handler.message_handler import MessageHandler, IndexingMessageHandler
|
from message_handler.message_handler import MessageHandler, IndexingMessageHandler
|
||||||
from services.ai_services import AIService, AIServiceConfig
|
from services.ai_services import AIService, AIServiceConfig, TTSService, LLMService, ImageGenService
|
||||||
from storage.search import SearchIndexer
|
from storage.search import SearchIndexer
|
||||||
|
|
||||||
|
|
||||||
@@ -44,7 +44,7 @@ class TestMessageHandler(unittest.TestCase):
|
|||||||
message_handler = MessageHandler("System prompt")
|
message_handler = MessageHandler("System prompt")
|
||||||
message_handler.add_user_message("User message")
|
message_handler.add_user_message("User message")
|
||||||
message_handler.add_assistant_message("Assistant message")
|
message_handler.add_assistant_message("Assistant message")
|
||||||
message_handler.add_user_message("User message plus something else")
|
message_handler.add_user_message("plus something else")
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
message_handler.get_llm_messages(),
|
message_handler.get_llm_messages(),
|
||||||
[
|
[
|
||||||
@@ -57,6 +57,7 @@ class TestMessageHandler(unittest.TestCase):
|
|||||||
message_handler = MessageHandler("System prompt")
|
message_handler = MessageHandler("System prompt")
|
||||||
message_handler.add_user_message("User message")
|
message_handler.add_user_message("User message")
|
||||||
message_handler.add_assistant_message("Assistant message")
|
message_handler.add_assistant_message("Assistant message")
|
||||||
|
message_handler.finalize_user_message()
|
||||||
message_handler.add_user_message("other user message")
|
message_handler.add_user_message("other user message")
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
message_handler.get_llm_messages(),
|
message_handler.get_llm_messages(),
|
||||||
@@ -69,25 +70,36 @@ class TestMessageHandler(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class MockAIService(AIService):
|
class MockTTSService(TTSService):
|
||||||
def __init__(self, **kwargs):
|
def run_tts(self, sentence):
|
||||||
super().__init__(**kwargs)
|
for word in sentence.split(" "):
|
||||||
|
time.sleep(0.1)
|
||||||
|
yield bytes(word, "utf-8")
|
||||||
|
|
||||||
def run_llm(self, messages, latest_user_message=None, stream=True):
|
|
||||||
return {"choices": [{"message": {"content": "Parsed user message."}}]}
|
class MockLLMService(LLMService):
|
||||||
|
def run_llm(self, messages) -> str:
|
||||||
|
return "Parsed user message."
|
||||||
|
|
||||||
|
class MockImageService(ImageGenService):
|
||||||
|
def run_image_gen(self, sentence) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class TestIndexingMessageHandler(unittest.TestCase):
|
class TestIndexingMessageHandler(unittest.TestCase):
|
||||||
def test_user_message_finalized(self):
|
def test_user_message_finalized(self):
|
||||||
mock_ai_service = MockAIService()
|
mock_tts_service = MockTTSService()
|
||||||
|
mock_llm_service = MockLLMService()
|
||||||
|
mock_image_service = MockImageService()
|
||||||
|
|
||||||
service_config = AIServiceConfig(
|
service_config = AIServiceConfig(
|
||||||
mock_ai_service, mock_ai_service, mock_ai_service
|
tts=mock_tts_service, llm=mock_llm_service, image=mock_image_service
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_indexer = MagicMock(spec=SearchIndexer)
|
mock_indexer = MagicMock(spec=SearchIndexer)
|
||||||
|
|
||||||
message_handler = IndexingMessageHandler(
|
message_handler = IndexingMessageHandler(
|
||||||
"Hello world", "story_id", service_config, mock_indexer
|
"Hello world", service_config, mock_indexer
|
||||||
)
|
)
|
||||||
message_handler.add_user_message("User message")
|
message_handler.add_user_message("User message")
|
||||||
message_handler.add_assistant_message("Assistant message will be ignored")
|
message_handler.add_assistant_message("Assistant message will be ignored")
|
||||||
@@ -101,11 +113,10 @@ class TestIndexingMessageHandler(unittest.TestCase):
|
|||||||
message_handler.write_messages_to_index()
|
message_handler.write_messages_to_index()
|
||||||
|
|
||||||
time.sleep(0.5)
|
time.sleep(0.5)
|
||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
mock_indexer.mock_calls,
|
mock_indexer.mock_calls,
|
||||||
[
|
[
|
||||||
call.index_text("Parsed user message."),
|
call.index_text('"Parsed user message."'),
|
||||||
call.index_text("New assistant message will not be ignored"),
|
call.index_text("New assistant message will not be ignored"),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
@@ -119,7 +130,7 @@ class TestIndexingMessageHandler(unittest.TestCase):
|
|||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
mock_indexer.mock_calls,
|
mock_indexer.mock_calls,
|
||||||
[
|
[
|
||||||
call.index_text("Parsed user message."),
|
call.index_text('"Parsed user message."'),
|
||||||
call.index_text("Assistant message second time"),
|
call.index_text("Assistant message second time"),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user