CoolFace
Apppublic

Akshanshsensei/PDF-Constrained-Conversational-Agent

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes
test_generate_with_tools.py67 linesDownload Raw Back to tests
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