tests pass

This commit is contained in:
Moishe Lettvin
2023-12-26 14:13:10 -05:00
parent e724720e76
commit df536b0ad0
3 changed files with 38 additions and 29 deletions

View File

@@ -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):

View File

@@ -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

View File

@@ -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"),
], ],
) )