Akshanshsensei/PDF-Constrained-Conversational-Agent
1
1import unittest2from unittest.mock import patch, MagicMock3from agent.generator import LLMGenerator4from core.metadata_store import DocumentMetadata5from core.tools import ToolExecutor6 7class MockChunk:8 def __init__(self, text):9 self.text = text10 11class MockStream:12 def __init__(self, chunks):13 self.chunks = chunks14 15 def __iter__(self):16 for c in self.chunks:17 yield MockChunk(c)18 19class TestGenerateWithTools(unittest.TestCase):20 def setUp(self):21 self.generator = LLMGenerator()22 self.metadata = DocumentMetadata(1, 10, 100, {}, [0], "fake text")23 self.executor = ToolExecutor(self.metadata)24 25 @patch('google.genai.Client')26 def test_no_tool_call(self, mock_client):27 # Mock LLM returning a direct answer28 mock_instance = mock_client.return_value29 mock_instance.models.generate_content_stream.return_value = MockStream(["Hello ", "world!"])30 31 stream = self.generator.generate_with_tools("query", [], [], self.executor)32 33 results = list(stream)34 self.assertEqual("".join(results), "Hello world!")35 36 @patch('google.genai.Client')37 def test_with_tool_call(self, mock_client):38 # First iteration returns a tool call39 # Second iteration returns the final answer40 mock_instance = mock_client.return_value41 42 # We need side_effect to return different streams on each call43 mock_instance.models.generate_content_stream.side_effect = [44 MockStream(['<tool_call>{"name": "count_word", "args": {"word": "test"}}</tool_call>']),45 MockStream(["The word test appears ", "3 times."])46 ]47 48 # Override executor to return a fake result without actually running re.findall on fake text49 self.executor.execute = MagicMock()50 self.executor.execute.return_value = MagicMock(error=None, output=3)51 52 stream = self.generator.generate_with_tools("query", [], [], self.executor)53 54 results = list(stream)55 56 # The first item should be a tool status dict57 self.assertTrue(isinstance(results[0], dict))58 self.assertEqual(results[0]["type"], "tool_status")59 self.assertEqual(results[0]["name"], "count_word")60 61 # The remaining items should be strings62 final_answer = "".join(results[1:])63 self.assertEqual(final_answer, "The word test appears 3 times.")64 65 # Ensure LLM was called twice66 self.assertEqual(mock_instance.models.generate_content_stream.call_count, 2)67 