File size: 1,755 Bytes
0cac9cf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
import unittest
from unittest.mock import patch

from rich.console import Console

from queryquest.sql.handoff import extract_sql_statements, expose_sql_statements


class SqlHandoffTests(unittest.TestCase):
    def test_extract_sql_statements_plain_json(self) -> None:
        output = '{"sql_statements": ["SELECT * FROM table1"], "explanation": "ok"}'
        self.assertEqual(extract_sql_statements(output), ["SELECT * FROM table1"])

    def test_extract_sql_statements_fenced_with_explanation_text(self) -> None:
        output = (
            "I can do that.\n"
            "```json\n"
            "{\n"
            '  "sql_statements": ["SELECT * FROM listings"],\n'
            '  "explanation": "preview"\n'
            "}\n"
            "```"
        )
        self.assertEqual(extract_sql_statements(output), ["SELECT * FROM listings"])

    def test_extract_sql_statements_embedded_json_object(self) -> None:
        output = "Use this payload: {\"sql_statements\": [\"SELECT 1\"], \"explanation\": \"ok\"} thanks"
        self.assertEqual(extract_sql_statements(output), ["SELECT 1"])

    def test_extract_sql_statements_invalid_json(self) -> None:
        self.assertEqual(extract_sql_statements("not json"), [])

    @patch("queryquest.sql.handoff.append_log")
    @patch("queryquest.sql.handoff.execute_sql_statements")
    def test_expose_sql_statements_executes_and_logs(self, mock_execute, mock_log) -> None:
        console = Console(record=True)
        sql = ["SELECT 1"]

        expose_sql_statements(sql, provider="groq", model="llama", console=console)

        mock_execute.assert_called_once_with(sql, console=console, excel_dir=None)
        self.assertTrue(mock_log.called)


if __name__ == "__main__":
    unittest.main()