Spaces:
Sleeping
Sleeping
File size: 6,455 Bytes
8429e5e 61f1033 8429e5e 61f1033 8429e5e 61f1033 8429e5e | 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 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 | """Tests for gazet.sql — SQL rewriting, normalization, and helpers."""
from unittest.mock import MagicMock, patch
import pandas as pd
from gazet.sql import (
_normalize_ne_subtypes,
_rewrite_data_paths,
_strip_fences,
run_geo_sql_dspy,
run_geo_sql_gguf,
)
class TestStripFences:
def test_plain_sql(self):
assert _strip_fences("SELECT * FROM foo") == "SELECT * FROM foo"
def test_sql_backtick_fences(self):
raw = "```sql\nSELECT id FROM bar\n```"
assert _strip_fences(raw) == "SELECT id FROM bar"
def test_backtick_fences_no_lang(self):
raw = "```\nSELECT 1\n```"
assert _strip_fences(raw) == "SELECT 1"
def test_none_input(self):
assert _strip_fences(None) == ""
def test_empty_string(self):
assert _strip_fences("") == ""
def test_partial_fence_leading(self):
raw = "```sql\nSELECT id"
assert _strip_fences(raw) == "SELECT id"
def test_partial_fence_trailing(self):
raw = "SELECT id\n```"
assert _strip_fences(raw) == "SELECT id"
def test_preserves_inner_backticks(self):
raw = "```sql\nSELECT `column` FROM table\n```"
assert _strip_fences(raw) == "SELECT `column` FROM table"
class TestRewriteDataPaths:
def test_symbolic_divisions_area(self):
sql = "SELECT * FROM read_parquet('divisions_area')"
result = _rewrite_data_paths(sql)
assert "divisions_area" in result
assert "read_parquet('divisions_area')" not in result
def test_symbolic_natural_earth(self):
sql = "SELECT * FROM read_parquet('natural_earth')"
result = _rewrite_data_paths(sql)
assert "natural_earth" in result
assert "read_parquet('natural_earth')" not in result
def test_hallucinated_division_path(self):
sql = "SELECT * FROM read_parquet('/data/overture/division_area/foo.parquet')"
result = _rewrite_data_paths(sql)
# The original hallucinated path should be gone
assert "/data/overture/division_area/foo.parquet" not in result
def test_hallucinated_natural_earth_path(self):
sql = (
"SELECT * FROM read_parquet('/some/natural_earth_geoparquet/data.parquet')"
)
result = _rewrite_data_paths(sql)
assert "/some/natural_earth_geoparquet/data.parquet" not in result
def test_double_quotes(self):
sql = 'SELECT * FROM read_parquet("divisions_area")'
result = _rewrite_data_paths(sql)
assert 'read_parquet("divisions_area")' not in result
def test_no_false_positive_unrelated_table(self):
sql = "SELECT * FROM read_parquet('some_other_table')"
result = _rewrite_data_paths(sql)
# Should be unchanged
assert "some_other_table" in result
class TestNormalizeNeSubtypes:
def test_lowercase_river(self):
sql = "WHERE n.subtype = 'River'"
result = _normalize_ne_subtypes(sql)
assert "'river'" in result
def test_lowercase_lake(self):
sql = "WHERE n.subtype = 'Lake'"
result = _normalize_ne_subtypes(sql)
assert "'lake'" in result
def test_lowercase_ocean(self):
sql = "WHERE subtype = 'Ocean'"
result = _normalize_ne_subtypes(sql)
assert "'ocean'" in result
def test_lowercase_sea(self):
sql = "WHERE subtype = 'Sea'"
result = _normalize_ne_subtypes(sql)
assert "'sea'" in result
def test_lowercase_range_mtn(self):
sql = "WHERE subtype = 'Range/mtn'"
result = _normalize_ne_subtypes(sql)
assert "'range/mtn'" in result
def test_terrain_area_replacement(self):
sql = "WHERE n.subtype = 'Terrain area'"
result = _normalize_ne_subtypes(sql)
assert "range/mtn" in result
assert "peninsula" in result
assert "depression" in result
def test_terrain_area_in_clause(self):
sql = "WHERE n.subtype IN ('Terrain area')"
result = _normalize_ne_subtypes(sql)
assert "range/mtn" in result
def test_already_lowercase_unchanged(self):
sql = "WHERE n.subtype = 'river'"
result = _normalize_ne_subtypes(sql)
assert result == sql
def test_island_group(self):
sql = "WHERE subtype = 'Island group'"
result = _normalize_ne_subtypes(sql)
assert "'island group'" in result
class TestRunGeoSqlGguf:
def test_empty_candidates_returns_none_result(self, con):
empty_df = pd.DataFrame()
events = list(run_geo_sql_gguf(con, "get Paris", empty_df))
assert len(events) >= 1
assert events[-1]["type"] == "result"
assert events[-1]["df"] is None
@patch("gazet.sql.generate_sql")
def test_execution_flow(self, mock_generate, con):
mock_generate.return_value = "SELECT 1"
# A query that succeeds but returns no useful geometry rows
events = list(run_geo_sql_gguf(con, "test", pd.DataFrame({"id": ["1"]})))
# Should emit at least sql_attempt and result
types = [e["type"] for e in events]
assert "sql_attempt" in types
class TestRunGeoSqlDspy:
def test_empty_candidates_returns_none_result(self, con):
empty_df = pd.DataFrame()
events = list(run_geo_sql_dspy(con, "get Paris", empty_df))
assert events[-1]["type"] == "result"
assert events[-1]["df"] is None
@patch("gazet.sql.write_sql")
def test_successful_sql(self, mock_write, con):
# Return a result object with .sql attribute
mock_pred = MagicMock()
mock_pred.sql = "SELECT 1 as id"
mock_write.return_value = mock_pred
df = pd.DataFrame(
{"id": ["x1"], "name": ["test"], "source": ["divisions_area"]}
)
events = list(run_geo_sql_dspy(con, "test", df, max_iterations=1))
types = [e["type"] for e in events]
assert "sql_attempt" in types
@patch("gazet.sql.write_sql")
def test_exhausts_iterations(self, mock_write, con):
mock_pred = MagicMock()
mock_pred.sql = "INVALID SQL" # will cause execution error
mock_write.return_value = mock_pred
df = pd.DataFrame(
{"id": ["x1"], "name": ["test"], "source": ["divisions_area"]}
)
events = list(run_geo_sql_dspy(con, "test", df, max_iterations=2))
# Should exhaust iterations and yield final result
assert events[-1]["type"] == "result"
|