Team Ai
Apppublic

prazy1208/text2sql

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
test_gen_sql_agent.py78 linesDownload Raw Back to tests
1"""Unit tests for backend.agents.gen_sql_agent (LLM mocked)."""2 3from unittest.mock import patch4 5from backend.agents import gen_sql_agent6 7 8def test_run_gen_sql_empty_rephrased():9    out = gen_sql_agent.run_gen_sql(10        "retail",11        "",12        [],13        [],14        ["public.t"],15        {},16    )17    assert out["generated_sql"] == ""18    assert "Missing" in out["reasoning_summary"]19 20 21def test_run_gen_sql_no_tables():22    out = gen_sql_agent.run_gen_sql(23        "retail",24        "Show revenue",25        [],26        [],27        [],28        {},29    )30    assert out["generated_sql"] == ""31    assert "No tables" in out["reasoning_summary"]32 33 34def test_run_gen_sql_parses_llm_json():35    payload = '{"generated_sql": "SELECT 1", "reasoning_summary": "smoke"}'36    with patch.object(gen_sql_agent, "chat_completion", return_value=payload) as mock_chat:37        out = gen_sql_agent.run_gen_sql(38            "retail",39            "Count rows",40            ["rule"],41            [{"id": 1, "query_type": "agg", "question_text": "q", "sql_query": "SELECT 1"}],42            ["public.orders"],43            {"public.orders": ["id"]},44        )45    assert out["generated_sql"] == "SELECT 1"46    assert out["reasoning_summary"] == "smoke"47    mock_chat.assert_called_once()48    _args, kwargs = mock_chat.call_args49    assert kwargs.get("agent_name") == gen_sql_agent.AGENT_GEN_SQL50 51 52def test_run_gen_sql_invalid_json_from_model():53    with patch.object(gen_sql_agent, "chat_completion", return_value="not json"):54        out = gen_sql_agent.run_gen_sql(55            "retail",56            "Q",57            [],58            [],59            ["public.t"],60            {"public.t": ["id"]},61        )62    assert out["generated_sql"] == ""63    assert "Could not extract SQL" in out["reasoning_summary"]64 65 66def test_run_gen_sql_llm_raises():67    with patch.object(gen_sql_agent, "chat_completion", side_effect=RuntimeError("api down")):68        out = gen_sql_agent.run_gen_sql(69            "retail",70            "Q",71            [],72            [],73            ["public.t"],74            {"public.t": ["id"]},75        )76    assert out["generated_sql"] == ""77    assert "api down" in out["reasoning_summary"]78