Skip to content

Commit 4b69d18

Browse files
authored
Fix(snowflake)!: use TO_GEOGRAPHY, TO_GEOMETRY instead of casts (tobymao#4017)
1 parent 2d4483c commit 4b69d18

2 files changed

Lines changed: 31 additions & 19 deletions

File tree

sqlglot/dialects/snowflake.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -910,6 +910,14 @@ def timestampfromparts_sql(self, expression: exp.TimestampFromParts) -> str:
910910

911911
return rename_func("TIMESTAMP_FROM_PARTS")(self, expression)
912912

913+
def cast_sql(self, expression: exp.Cast, safe_prefix: t.Optional[str] = None) -> str:
914+
if expression.is_type(exp.DataType.Type.GEOGRAPHY):
915+
return self.func("TO_GEOGRAPHY", expression.this)
916+
if expression.is_type(exp.DataType.Type.GEOMETRY):
917+
return self.func("TO_GEOMETRY", expression.this)
918+
919+
return super().cast_sql(expression, safe_prefix=safe_prefix)
920+
913921
def trycast_sql(self, expression: exp.TryCast) -> str:
914922
value = expression.this
915923

tests/dialects/test_snowflake.py

Lines changed: 23 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -11,30 +11,12 @@ class TestSnowflake(Validator):
1111
dialect = "snowflake"
1212

1313
def test_snowflake(self):
14-
self.validate_identity("1 /* /* */")
15-
self.validate_identity(
16-
"SELECT * FROM table AT (TIMESTAMP => '2024-07-24') UNPIVOT(a FOR b IN (c)) AS pivot_table"
17-
)
18-
1914
self.assertEqual(
2015
# Ensures we don't fail when generating ParseJSON with the `safe` arg set to `True`
2116
self.validate_identity("""SELECT TRY_PARSE_JSON('{"x: 1}')""").sql(),
2217
"""SELECT PARSE_JSON('{"x: 1}')""",
2318
)
2419

25-
self.validate_identity(
26-
"transform(x, a int -> a + a + 1)",
27-
"TRANSFORM(x, a -> CAST(a AS INT) + CAST(a AS INT) + 1)",
28-
)
29-
30-
self.validate_all(
31-
"ARRAY_CONSTRUCT_COMPACT(1, null, 2)",
32-
write={
33-
"spark": "ARRAY_COMPACT(ARRAY(1, NULL, 2))",
34-
"snowflake": "ARRAY_CONSTRUCT_COMPACT(1, NULL, 2)",
35-
},
36-
)
37-
3820
expr = parse_one("SELECT APPROX_TOP_K(C4, 3, 5) FROM t")
3921
expr.selects[0].assert_is(exp.AggFunc)
4022
self.assertEqual(expr.sql(dialect="snowflake"), "SELECT APPROX_TOP_K(C4, 3, 5) FROM t")
@@ -98,7 +80,6 @@ def test_snowflake(self):
9880
self.validate_identity("WITH x AS (SELECT 1 AS foo) SELECT foo FROM IDENTIFIER('x')")
9981
self.validate_identity("WITH x AS (SELECT 1 AS foo) SELECT IDENTIFIER('foo') FROM x")
10082
self.validate_identity("INITCAP('iqamqinterestedqinqthisqtopic', 'q')")
101-
self.validate_identity("CAST(x AS GEOMETRY)")
10283
self.validate_identity("OBJECT_CONSTRUCT(*)")
10384
self.validate_identity("SELECT CAST('2021-01-01' AS DATE) + INTERVAL '1 DAY'")
10485
self.validate_identity("SELECT HLL(*)")
@@ -115,6 +96,10 @@ def test_snowflake(self):
11596
self.validate_identity("ALTER TABLE a SWAP WITH b")
11697
self.validate_identity("SELECT MATCH_CONDITION")
11798
self.validate_identity("SELECT * REPLACE (CAST(col AS TEXT) AS scol) FROM t")
99+
self.validate_identity("1 /* /* */")
100+
self.validate_identity(
101+
"SELECT * FROM table AT (TIMESTAMP => '2024-07-24') UNPIVOT(a FOR b IN (c)) AS pivot_table"
102+
)
118103
self.validate_identity(
119104
"SELECT * FROM quarterly_sales PIVOT(SUM(amount) FOR quarter IN ('2023_Q1', '2023_Q2', '2023_Q3', '2023_Q4', '2024_Q1') DEFAULT ON NULL (0)) ORDER BY empid"
120105
)
@@ -139,6 +124,18 @@ def test_snowflake(self):
139124
self.validate_identity(
140125
"SELECT * FROM DATA AS DATA_L ASOF JOIN DATA AS DATA_R MATCH_CONDITION (DATA_L.VAL > DATA_R.VAL) ON DATA_L.ID = DATA_R.ID"
141126
)
127+
self.validate_identity(
128+
"CAST(x AS GEOGRAPHY)",
129+
"TO_GEOGRAPHY(x)",
130+
)
131+
self.validate_identity(
132+
"CAST(x AS GEOMETRY)",
133+
"TO_GEOMETRY(x)",
134+
)
135+
self.validate_identity(
136+
"transform(x, a int -> a + a + 1)",
137+
"TRANSFORM(x, a -> CAST(a AS INT) + CAST(a AS INT) + 1)",
138+
)
142139
self.validate_identity(
143140
"SELECT * FROM s WHERE c NOT IN (1, 2, 3)",
144141
"SELECT * FROM s WHERE NOT c IN (1, 2, 3)",
@@ -308,6 +305,13 @@ def test_snowflake(self):
308305
"SELECT * RENAME (a AS b), c AS d FROM xxx",
309306
)
310307

308+
self.validate_all(
309+
"ARRAY_CONSTRUCT_COMPACT(1, null, 2)",
310+
write={
311+
"spark": "ARRAY_COMPACT(ARRAY(1, NULL, 2))",
312+
"snowflake": "ARRAY_CONSTRUCT_COMPACT(1, NULL, 2)",
313+
},
314+
)
311315
self.validate_all(
312316
"OBJECT_CONSTRUCT_KEEP_NULL('key_1', 'one', 'key_2', NULL)",
313317
read={

0 commit comments

Comments
 (0)