prazy1208/text2sql
0
1"""Unit tests for backend.services.sql_validator."""2 3import pytest4 5from backend.services.sql_validator import validate_generated_sql6 7 8def test_empty_sql_fails():9 r = validate_generated_sql("")10 assert r["validation_passed"] is False11 assert "EMPTY_SQL" in r["validation_error_codes"]12 13 14def test_select_single_statement_passes():15 sql = "SELECT a FROM public.orders WHERE id = 1"16 r = validate_generated_sql(17 sql,18 selected_tables=["public.orders"],19 )20 assert r["validation_passed"] is True21 assert r["is_single_statement"] is True22 assert r["is_select_only"] is True23 24 25def test_with_select_passes():26 sql = "WITH x AS (SELECT 1 AS n) SELECT n FROM x"27 r = validate_generated_sql(sql, selected_tables=None)28 assert r["validation_passed"] is True29 30 31def test_explain_analyze_select_passes():32 sql = "EXPLAIN ANALYZE SELECT 1"33 r = validate_generated_sql(sql)34 assert r["validation_passed"] is True35 36 37def test_multiple_statements_fails():38 sql = "SELECT 1; SELECT 2"39 r = validate_generated_sql(sql)40 assert r["validation_passed"] is False41 assert "MULTIPLE_STATEMENTS" in r["validation_error_codes"]42 43 44@pytest.mark.parametrize(45 "snippet",46 [47 "INSERT INTO t VALUES (1)",48 "UPDATE t SET a = 1",49 "DELETE FROM t",50 "DROP TABLE t",51 "CREATE TABLE t (id int)",52 "TRUNCATE t",53 ],54)55def test_forbidden_keyword_fails(snippet):56 r = validate_generated_sql(snippet)57 assert r["validation_passed"] is False58 assert "FORBIDDEN_KEYWORD" in r["validation_error_codes"]59 assert r["blocked_keywords"]60 61 62def test_keyword_in_string_literal_ignored():63 sql = "SELECT 'DELETE' AS x FROM public.t"64 r = validate_generated_sql(sql, selected_tables=["public.t"])65 assert r["validation_passed"] is True66 67 68def test_select_into_temp_fails():69 sql = "SELECT 1 INTO TEMP foo"70 r = validate_generated_sql(sql)71 assert r["validation_passed"] is False72 assert "SELECT_INTO_DDL" in r["validation_error_codes"]73 74 75def test_table_not_in_selection_fails():76 sql = "SELECT 1 FROM other.orders"77 r = validate_generated_sql(78 sql,79 selected_tables=["public.orders"],80 )81 assert r["validation_passed"] is False82 assert "TABLE_NOT_IN_SELECTION" in r["validation_error_codes"]83 84 85def test_from_join_table_must_match_selected():86 sql = "SELECT 1 FROM retail.customers c JOIN retail.orders o ON c.id = o.customer_id"87 r = validate_generated_sql(88 sql,89 selected_tables=["retail.customers", "retail.orders"],90 )91 assert r["validation_passed"] is True92 