Team Ai
Apppublic

prazy1208/text2sql

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
test_query_route_gen_sql.py447 linesDownload Raw Back to tests
1"""Integration-style tests for POST /query Gen-SQL path (DB and LLM mocked)."""2 3import uuid4from contextlib import ExitStack5from unittest.mock import patch6 7import pytest8from fastapi.testclient import TestClient9 10from backend.api.main import app11 12 13@pytest.fixture14def client():15    return TestClient(app)16 17 18def test_post_query_includes_generated_sql_and_persists_gen_sql_row(client):19    sid = str(uuid.uuid4())20    validation_ok = {21        "validation_passed": True,22        "validation_error_codes": "",23        "validation_error_message": "",24        "blocked_keywords": "",25        "is_single_statement": True,26        "is_select_only": True,27    }28    with ExitStack() as stack:29        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))30        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))31        stack.enter_context(patch("backend.api.routes.query.get_recent_chat_messages", return_value=[]))32        stack.enter_context(patch("backend.api.routes.query.get_session_memory", return_value={}))33        stack.enter_context(patch("backend.api.routes.query.insert_intent_review", return_value=1))34        stack.enter_context(35            patch("backend.api.routes.query.list_relationships_from_metadata", return_value=[])36        )37        stack.enter_context(patch("backend.api.routes.query.insert_intent_output", return_value=1))38        stack.enter_context(patch("backend.api.routes.query.insert_table_agent_output", return_value=10))39        stack.enter_context(patch("backend.api.routes.query.insert_column_agent_output", return_value=20))40        stack.enter_context(patch("backend.api.routes.query.insert_few_shot_agent_output", return_value=30))41        insert_gen = stack.enter_context(42            patch("backend.api.routes.query.insert_gen_sql_agent_output", return_value=100)43        )44        stack.enter_context(45            patch(46                "backend.api.routes.query.run_intent",47                return_value={48                    "rephrased_question": "Total sales last month",49                    "resolved_question": "Total sales last month",50                    "confidence_score": 96,51                    "clarification_question": "",52                    "keywords": ["sales"],53                    "business_insights": ["Use net amounts"],54                },55            )56        )57        stack.enter_context(58            patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["public.orders"]})59        )60        stack.enter_context(61            patch(62                "backend.api.routes.query.run_column_agent",63                return_value={"selected_columns": {"public.orders": ["id", "amount"]}},64            )65        )66        stack.enter_context(patch("backend.api.routes.query.run_few_shot_agent", return_value={"few_shot_examples": []}))67        stack.enter_context(68            patch(69                "backend.api.routes.query.run_gen_sql",70                return_value={71                    "generated_sql": "SELECT sum(amount) FROM public.orders",72                    "reasoning_summary": "agg",73                },74            )75        )76        stack.enter_context(patch("backend.api.routes.query.validate_generated_sql", return_value=validation_ok))77 78        res = client.post(79            "/query",80            json={"message": "sales?", "use_case": "retail", "session_id": sid},81        )82 83    assert res.status_code == 20084    body = res.json()85    assert body["generated_sql"] == "SELECT sum(amount) FROM public.orders"86    assert body["rephrased_question"] == "Total sales last month"87    assert body["clarification_question"] == ""88    insert_gen.assert_called_once()89    args, kwargs = insert_gen.call_args90    assert args[0] == 1  # intent_output_id91    assert args[1] == "SELECT sum(amount) FROM public.orders"92    assert args[2] == "agg"  # reasoning_summary93    assert args[3] is True  # validation_passed94 95 96def test_post_query_merges_validation_failure_into_error(client):97    sid = str(uuid.uuid4())98    validation_bad = {99        "validation_passed": False,100        "validation_error_codes": "FORBIDDEN_KEYWORD",101        "validation_error_message": "blocked",102        "blocked_keywords": "DELETE",103        "is_single_statement": True,104        "is_select_only": False,105    }106    with ExitStack() as stack:107        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))108        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))109        stack.enter_context(patch("backend.api.routes.query.get_recent_chat_messages", return_value=[]))110        stack.enter_context(patch("backend.api.routes.query.get_session_memory", return_value={}))111        stack.enter_context(patch("backend.api.routes.query.insert_intent_review", return_value=1))112        stack.enter_context(113            patch("backend.api.routes.query.list_relationships_from_metadata", return_value=[])114        )115        stack.enter_context(patch("backend.api.routes.query.insert_intent_output", return_value=1))116        stack.enter_context(patch("backend.api.routes.query.insert_table_agent_output", return_value=10))117        stack.enter_context(patch("backend.api.routes.query.insert_column_agent_output", return_value=20))118        stack.enter_context(patch("backend.api.routes.query.insert_few_shot_agent_output", return_value=30))119        stack.enter_context(patch("backend.api.routes.query.insert_gen_sql_agent_output", return_value=100))120        stack.enter_context(121            patch(122                "backend.api.routes.query.run_intent",123                return_value={124                    "rephrased_question": "Show totals across categories for the retail domain",125                    "resolved_question": "Show totals across categories for the retail domain",126                    "confidence_score": 96,127                    "clarification_question": "",128                    "keywords": ["totals", "categories"],129                    "business_insights": [],130                },131            )132        )133        stack.enter_context(patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["a.b"]}))134        stack.enter_context(patch("backend.api.routes.query.run_column_agent", return_value={"selected_columns": {"a.b": ["x"]}}))135        stack.enter_context(patch("backend.api.routes.query.run_few_shot_agent", return_value={"few_shot_examples": []}))136        stack.enter_context(137            patch(138                "backend.api.routes.query.run_gen_sql",139                return_value={"generated_sql": "SELECT 1; DELETE FROM a.b", "reasoning_summary": ""},140            )141        )142        stack.enter_context(patch("backend.api.routes.query.validate_generated_sql", return_value=validation_bad))143 144        res = client.post(145            "/query",146            json={"message": "show totals by category", "use_case": "retail", "session_id": sid},147        )148 149    assert res.status_code == 200150    body = res.json()151    assert "SQL validation" in (body.get("error") or "")152    assert body["generated_sql"]  # still returned for inspection153    assert body["clarification_question"] == ""154 155 156def test_post_query_low_confidence_returns_confirmation_prompt(client):157    sid = str(uuid.uuid4())158    with ExitStack() as stack:159        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))160        insert_chat = stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))161        stack.enter_context(patch("backend.api.routes.query.get_recent_chat_messages", return_value=[]))162        stack.enter_context(patch("backend.api.routes.query.get_session_memory", return_value={}))163        stack.enter_context(patch("backend.api.routes.query.list_relationships_from_metadata", return_value=[]))164        stack.enter_context(patch("backend.api.routes.query.insert_intent_output", return_value=11))165        insert_review = stack.enter_context(patch("backend.api.routes.query.insert_intent_review", return_value=22))166        run_table = stack.enter_context(patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["a.b"]}))167        stack.enter_context(168            patch(169                "backend.api.routes.query.run_intent",170                return_value={171                    "rephrased_question": "Find branch-wise high spenders for last month",172                    "resolved_question": "Find branch-wise high spenders for last month",173                    "confidence_score": 62,174                    "clarification_question": "Did you mean: find branch-wise high spenders for last month?",175                    "keywords": ["ranking", "spend", "branch"],176                    "business_insights": ["High activity entities"],177                },178            )179        )180 181        res = client.post(182            "/query",183            json={"message": "high spenders by branch last month", "use_case": "finance", "session_id": sid},184        )185 186    assert res.status_code == 200187    body = res.json()188    assert body["needs_confirmation"] is True189    assert body["conversation_state"] == "waiting_intent_confirmation"190    assert body["intent_confidence"] == 62191    assert body["generated_sql"] == ""192    assert body["pending_intent_id"] == 11193    run_table.assert_not_called()194    insert_review.assert_called_once()195    # user turn + assistant confirmation prompt196    assert insert_chat.call_count >= 2197 198 199def test_post_query_intent_confirmation_yes_continues_pipeline(client):200    sid = str(uuid.uuid4())201    validation_ok = {202        "validation_passed": True,203        "validation_error_codes": "",204        "validation_error_message": "",205        "blocked_keywords": "",206        "is_single_statement": True,207        "is_select_only": True,208    }209    with ExitStack() as stack:210        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))211        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))212        stack.enter_context(patch("backend.api.routes.query.get_session_memory", return_value={}))213        stack.enter_context(214            patch(215                "backend.api.routes.query.get_latest_pending_intent",216                return_value={217                    "intent_output_id": 44,218                    "rephrased_question": "Top customers by spend last month",219                    "user_input": "top spenders",220                    "keywords": ["ranking", "spend"],221                    "business_insights": ["High activity entities"],222                    "confidence_score": 65,223                },224            )225        )226        update_review = stack.enter_context(227            patch("backend.api.routes.query.update_intent_review_status", return_value=None)228        )229        stack.enter_context(patch("backend.api.routes.query.list_relationships_from_metadata", return_value=[]))230        stack.enter_context(patch("backend.api.routes.query.insert_table_agent_output", return_value=9))231        stack.enter_context(patch("backend.api.routes.query.insert_column_agent_output", return_value=19))232        stack.enter_context(patch("backend.api.routes.query.insert_few_shot_agent_output", return_value=29))233        stack.enter_context(patch("backend.api.routes.query.insert_gen_sql_agent_output", return_value=39))234        run_table = stack.enter_context(235            patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["finance_schema.transactions"]})236        )237        stack.enter_context(238            patch(239                "backend.api.routes.query.run_column_agent",240                return_value={"selected_columns": {"finance_schema.transactions": ["transaction_id", "amount"]}},241            )242        )243        stack.enter_context(patch("backend.api.routes.query.run_few_shot_agent", return_value={"few_shot_examples": []}))244        stack.enter_context(245            patch(246                "backend.api.routes.query.run_gen_sql",247                return_value={248                    "generated_sql": "SELECT customer_id, sum(amount) FROM finance_schema.transactions GROUP BY customer_id",249                    "reasoning_summary": "aggregation",250                },251            )252        )253        stack.enter_context(patch("backend.api.routes.query.validate_generated_sql", return_value=validation_ok))254 255        res = client.post(256            "/query",257            json={258                "message": "yes",259                "use_case": "finance",260                "session_id": sid,261                "message_type": "intent_confirmation",262                "confirmation": "yes",263            },264        )265 266    assert res.status_code == 200267    body = res.json()268    assert body["conversation_state"] == "completed"269    assert body["needs_confirmation"] is False270    assert body["generated_sql"]271    update_review.assert_called_once_with(44, "confirmed")272    run_table.assert_called_once()273 274 275def test_post_query_intent_confirmation_no_waits_for_rephrase(client):276    sid = str(uuid.uuid4())277    with ExitStack() as stack:278        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))279        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))280        stack.enter_context(patch("backend.api.routes.query.get_session_memory", return_value={}))281        stack.enter_context(282            patch(283                "backend.api.routes.query.get_latest_pending_intent",284                return_value={285                    "intent_output_id": 51,286                    "rephrased_question": "Branch-wise top spenders for last month",287                    "user_input": "top spenders",288                    "keywords": ["ranking", "branch"],289                    "business_insights": ["High activity entities"],290                    "confidence_score": 61,291                },292            )293        )294        update_review = stack.enter_context(295            patch("backend.api.routes.query.update_intent_review_status", return_value=None)296        )297        run_table = stack.enter_context(298            patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["x.y"]})299        )300 301        res = client.post(302            "/query",303            json={304                "message": "no",305                "use_case": "finance",306                "session_id": sid,307                "message_type": "intent_confirmation",308                "confirmation": "no",309            },310        )311 312    assert res.status_code == 200313    body = res.json()314    assert body["conversation_state"] == "waiting_user_rephrase"315    assert body["needs_confirmation"] is False316    assert body["generated_sql"] == ""317    assert body["pending_intent_id"] == 51318    update_review.assert_called_once_with(51, "rejected")319    run_table.assert_not_called()320 321 322def test_post_query_trivial_greeting_skips_confirmation_and_waits_for_query(client):323    """Pure greetings get a short reply and waiting_analytical_query — no Yes/No on meta-intent."""324    sid = str(uuid.uuid4())325    with ExitStack() as stack:326        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))327        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))328        stack.enter_context(patch("backend.api.routes.query.get_recent_chat_messages", return_value=[]))329        stack.enter_context(patch("backend.api.routes.query.get_session_memory", return_value={}))330        stack.enter_context(patch("backend.api.routes.query.list_relationships_from_metadata", return_value=[]))331        stack.enter_context(patch("backend.api.routes.query.insert_intent_output", return_value=3))332        stack.enter_context(patch("backend.api.routes.query.insert_intent_review", return_value=4))333        run_intent = stack.enter_context(patch("backend.api.routes.query.run_intent"))334        run_table = stack.enter_context(patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["a.b"]}))335 336        res = client.post(337            "/query",338            json={"message": "hi", "use_case": "retail", "session_id": sid},339        )340 341    assert res.status_code == 200342    body = res.json()343    assert body["conversation_state"] == "waiting_analytical_query"344    assert body["needs_confirmation"] is False345    assert body["pending_intent_id"] is None346    assert "analyze" in (body.get("clarification_question") or "").lower()347    run_intent.assert_not_called()348    run_table.assert_not_called()349 350 351def test_post_query_open_invite_yes_returns_waiting_analytical_query(client):352    sid = str(uuid.uuid4())353    with ExitStack() as stack:354        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))355        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))356        stack.enter_context(357            patch(358                "backend.api.routes.query.get_latest_pending_intent",359                return_value={360                    "intent_output_id": 70,361                    "rephrased_question": "User may want to explore analytics",362                    "user_input": "hello",363                    "keywords": ["explore"],364                    "business_insights": [],365                    "confidence_score": 55,366                },367            )368        )369        stack.enter_context(370            patch(371                "backend.api.routes.query.get_session_memory",372                return_value={"pending_confirm_kind": "open_invite", "pending_intent_output_id": 70},373            )374        )375        update_review = stack.enter_context(376            patch("backend.api.routes.query.update_intent_review_status", return_value=None)377        )378        run_table = stack.enter_context(patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["x.y"]}))379 380        res = client.post(381            "/query",382            json={383                "message": "yes",384                "use_case": "finance",385                "session_id": sid,386                "message_type": "intent_confirmation",387                "confirmation": "yes",388            },389        )390 391    assert res.status_code == 200392    body = res.json()393    assert body["conversation_state"] == "waiting_analytical_query"394    assert body["pending_intent_id"] is None395    assert body["generated_sql"] == ""396    update_review.assert_called_once_with(70, "rejected")397    run_table.assert_not_called()398 399 400def test_post_query_open_invite_no_returns_conversation_ended(client):401    sid = str(uuid.uuid4())402    with ExitStack() as stack:403        stack.enter_context(patch("backend.api.routes.query.session_exists", return_value=True))404        stack.enter_context(patch("backend.api.routes.query.insert_chat_message", return_value=1))405        stack.enter_context(406            patch(407                "backend.api.routes.query.get_latest_pending_intent",408                return_value={409                    "intent_output_id": 71,410                    "rephrased_question": "User may want to explore analytics",411                    "user_input": "hello",412                    "keywords": ["explore"],413                    "business_insights": [],414                    "confidence_score": 55,415                },416            )417        )418        stack.enter_context(419            patch(420                "backend.api.routes.query.get_session_memory",421                return_value={"pending_confirm_kind": "open_invite", "pending_intent_output_id": 71},422            )423        )424        update_review = stack.enter_context(425            patch("backend.api.routes.query.update_intent_review_status", return_value=None)426        )427        run_table = stack.enter_context(patch("backend.api.routes.query.run_table_agent", return_value={"selected_tables": ["x.y"]}))428 429        res = client.post(430            "/query",431            json={432                "message": "no",433                "use_case": "finance",434                "session_id": sid,435                "message_type": "intent_confirmation",436                "confirmation": "no",437            },438        )439 440    assert res.status_code == 200441    body = res.json()442    assert body["conversation_state"] == "conversation_ended"443    assert body["clarification_question"] == "Ok, thank you!"444    assert body["pending_intent_id"] is None445    update_review.assert_called_once_with(71, "rejected")446    run_table.assert_not_called()447