prazy1208/text2sql
0
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 