Skip to content

Commit ae4bcb2

Browse files
committed
pr feedback
1 parent 3ec5e35 commit ae4bcb2

3 files changed

Lines changed: 157 additions & 195 deletions

File tree

sqlglot/optimizer/qualify_columns.py

Lines changed: 76 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -641,6 +641,70 @@ def _convert_columns_to_dots(scope: Scope, resolver: Resolver) -> None:
641641
scope.clear_cache()
642642

643643

644+
def _qualify_positional_column(
645+
scope: Scope,
646+
resolver: Resolver,
647+
column: exp.Column,
648+
column_table: str,
649+
column_source: exp.Expr | Scope,
650+
source_columns: t.Sequence[str],
651+
pivots: t.Sequence[exp.Pivot],
652+
allow_partial_qualification: bool,
653+
) -> bool:
654+
"""Resolve a positional column, returning whether to skip further qualification."""
655+
if (
656+
not resolver.dialect.SUPPORTS_POSITIONAL_COLUMN_REFS
657+
or not isinstance(column.this, exp.Parameter)
658+
or not isinstance(position := column.this.this, exp.Literal)
659+
or not position.is_int
660+
):
661+
return False
662+
663+
# Pivots from unrelated sources may share this scope, so prefer an exact
664+
# output-alias match. For an aliasless chain, fall back to the last operator
665+
# on the referenced source.
666+
scope_pivot = next((pivot for pivot in scope.pivots if pivot.alias == column_table), None)
667+
if not scope_pivot:
668+
scope_pivot = next(
669+
(
670+
pivot
671+
for pivot in reversed(scope.pivots)
672+
if pivot.parent and pivot.parent.alias_or_name == column_table
673+
),
674+
None,
675+
)
676+
if (
677+
pivots
678+
or scope_pivot
679+
or (isinstance(column_source, exp.Table) and column_source.alias_column_names)
680+
or (isinstance(column_source, Scope) and column_source.outer_columns)
681+
or not source_columns
682+
or "*" in source_columns
683+
):
684+
if scope_pivot:
685+
column.set("table", exp.to_identifier(scope_pivot.alias))
686+
return True
687+
688+
positional_columns = resolver.get_source_columns(column_table, only_visible=True)
689+
position_value = int(position.to_py())
690+
if not 1 <= position_value <= len(positional_columns):
691+
if allow_partial_qualification:
692+
return True
693+
raise OptimizeError(
694+
f"Positional reference ${position_value} is out of range for source '{column_table}'"
695+
)
696+
697+
positional_name = positional_columns[position_value - 1]
698+
if positional_columns.count(positional_name) > 1:
699+
return True
700+
701+
positional_identifier = exp.to_identifier(positional_name)
702+
resolver.dialect.quote_identifier(positional_identifier, identify=False)
703+
column.set("this", positional_identifier)
704+
705+
return False
706+
707+
644708
def _qualify_columns(
645709
scope: Scope,
646710
resolver: Resolver,
@@ -658,62 +722,23 @@ def _qualify_columns(
658722
pivots = (
659723
column_source.args.get("pivots", []) if isinstance(column_source, exp.Table) else []
660724
)
661-
if source_columns and pivots:
725+
if pivots:
662726
# Each operator's input is the previous one's output
663727
for pivot in pivots:
664728
source_columns = pivot.output_columns(source_columns)
665729

666-
if (
667-
resolver.dialect.SUPPORTS_POSITIONAL_COLUMN_REFS
668-
and isinstance(column.this, exp.Parameter)
669-
and isinstance(position := column.this.this, exp.Literal)
670-
and position.is_int
730+
if _qualify_positional_column(
731+
scope,
732+
resolver,
733+
column,
734+
column_table,
735+
column_source,
736+
source_columns,
737+
pivots,
738+
allow_partial_qualification,
671739
):
672-
# Pivots from unrelated sources may share this scope, so prefer an exact
673-
# output-alias match. For an aliasless chain, fall back to the last operator
674-
# on the referenced source.
675-
scope_pivot = next(
676-
(pivot for pivot in scope.pivots if pivot.alias == column_table), None
677-
)
678-
if not scope_pivot:
679-
scope_pivot = next(
680-
(
681-
pivot
682-
for pivot in reversed(scope.pivots)
683-
if pivot.parent and pivot.parent.alias_or_name == column_table
684-
),
685-
None,
686-
)
687-
if (
688-
pivots
689-
or scope_pivot
690-
or (isinstance(column_source, exp.Table) and column_source.alias_column_names)
691-
or (isinstance(column_source, Scope) and column_source.outer_columns)
692-
or not source_columns
693-
or "*" in source_columns
694-
):
695-
if scope_pivot:
696-
column.set("table", exp.to_identifier(scope_pivot.alias))
697-
continue
698-
699-
positional_columns = resolver.get_source_columns(column_table, only_visible=True)
700-
position_value = int(position.to_py())
701-
if not 1 <= position_value <= len(positional_columns):
702-
if allow_partial_qualification:
703-
continue
704-
raise OptimizeError(
705-
f"Positional reference ${position_value} is out of range for source '{column_table}'"
706-
)
707-
708-
positional_name = positional_columns[position_value - 1]
709-
if positional_columns.count(positional_name) > 1:
710-
continue
711-
712-
positional_identifier = exp.to_identifier(positional_name)
713-
resolver.dialect.quote_identifier(positional_identifier, identify=False)
714-
715-
column.set("this", positional_identifier)
716-
column_name = column.name
740+
continue
741+
column_name = column.name
717742
if (
718743
not allow_partial_qualification
719744
and source_columns

tests/fixtures/optimizer/qualify_columns.sql

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1246,3 +1246,75 @@ SELECT piv.id AS id, piv.month AS month, piv.revenue AS revenue FROM (SELECT unp
12461246
# dialect: duckdb
12471247
SELECT u.* FROM unpivotable UNPIVOT(revenue FOR month IN (jan, feb)) UNPIVOT(headcount FOR region IN (north, south)) AS u;
12481248
SELECT u.id AS id, u.month AS month, u.revenue AS revenue, u.region AS region, u.headcount AS headcount FROM unpivotable AS unpivotable UNPIVOT(revenue FOR month IN (jan, feb)) UNPIVOT(headcount FOR region IN (north, south)) AS u;
1249+
1250+
# title: resolve Snowflake positional references inside composed expressions
1251+
# execute: false
1252+
# dialect: snowflake
1253+
WITH t AS (SELECT PARSE_JSON('{"id": 1}') AS data, 'file' AS filepath) SELECT t.$1:"id"::INT AS id, t.$2 AS filename FROM t;
1254+
WITH T AS (SELECT PARSE_JSON('{"id": 1}') AS DATA, 'file' AS FILEPATH) SELECT CAST(GET_PATH(T.DATA, 'id') AS INT) AS ID, T.FILEPATH AS FILENAME FROM T AS T;
1255+
1256+
# title: resolve Snowflake positional reference to a quoted column
1257+
# execute: false
1258+
# dialect: snowflake
1259+
WITH t AS (SELECT 1 AS "lower") SELECT t.$1 FROM t;
1260+
WITH T AS (SELECT 1 AS "lower") SELECT T."lower" AS "lower" FROM T AS T;
1261+
1262+
# title: preserve ambiguous Snowflake positional reference and resolve unique one
1263+
# execute: false
1264+
# dialect: snowflake
1265+
WITH t AS (SELECT 1 AS x, 2 AS x, 3 AS y) SELECT t.$2, t.$3 FROM t;
1266+
WITH T AS (SELECT 1 AS X, 2 AS X, 3 AS Y) SELECT T.$2 AS _COL_0, T.Y AS Y FROM T AS T;
1267+
1268+
# title: preserve Snowflake positional reference with CTE column aliases
1269+
# execute: false
1270+
# dialect: snowflake
1271+
WITH t("lower") AS (SELECT 1) SELECT t.$1 FROM t;
1272+
WITH T AS (SELECT 1 AS "lower") SELECT T.$1 AS _COL_0 FROM T AS T;
1273+
1274+
# title: preserve Snowflake positional reference through UNPIVOT
1275+
# execute: false
1276+
# dialect: snowflake
1277+
WITH t AS (SELECT 1 AS "SELECT", 2 AS a, 3 AS b) SELECT t.$1 FROM t UNPIVOT(v FOR k IN (a, b));
1278+
WITH T AS (SELECT 1 AS "SELECT", 2 AS A, 3 AS B) SELECT T.$1 AS _COL_0 FROM T AS T UNPIVOT(V FOR K IN (A, B)) AS T;
1279+
1280+
# title: preserve Snowflake positional reference through aliased PIVOT
1281+
# execute: false
1282+
# dialect: snowflake
1283+
WITH t AS (SELECT 1 AS id, 2 AS v, 'a' AS k) SELECT p.$2 FROM t PIVOT(SUM(v) FOR k IN ('a' AS "SELECT")) AS p;
1284+
WITH T AS (SELECT 1 AS ID, 2 AS V, 'a' AS K) SELECT P.$2 AS _COL_0 FROM T AS T PIVOT(SUM(T.V) FOR T.K IN ('a' AS "SELECT")) AS P;
1285+
1286+
# title: retarget Snowflake positional reference to generated PIVOT alias
1287+
# execute: false
1288+
# dialect: snowflake
1289+
WITH t AS (SELECT 1 AS id, 2 AS v, 'a' AS k) SELECT t.$2 FROM t PIVOT(SUM(v) FOR k IN ('a' AS a));
1290+
WITH T AS (SELECT 1 AS ID, 2 AS V, 'a' AS K) SELECT _0.$2 AS _COL_0 FROM T AS T PIVOT(SUM(T.V) FOR T.K IN ('a' AS A)) AS _0;
1291+
1292+
# title: retarget Snowflake positional reference to final chained operator
1293+
# execute: false
1294+
# dialect: snowflake
1295+
WITH t AS (SELECT 1 AS id, 2 AS v, 'a' AS k) SELECT t.$1 FROM t PIVOT(SUM(v) FOR k IN ('a' AS a)) UNPIVOT(v2 FOR k2 IN (a));
1296+
WITH T AS (SELECT 1 AS ID, 2 AS V, 'a' AS K) SELECT T.$1 AS _COL_0 FROM T AS T PIVOT(SUM(T.V) FOR T.K IN ('a' AS A)) UNPIVOT(V2 FOR K2 IN (A)) AS T;
1297+
1298+
# title: resolve Snowflake positional reference beside unrelated PIVOT
1299+
# execute: false
1300+
# dialect: snowflake
1301+
WITH t AS (SELECT 1 AS x), s AS (SELECT 1 AS id, 2 AS v, 'a' AS k) SELECT t.$1, p.$2 FROM t CROSS JOIN s PIVOT(SUM(v) FOR k IN ('a' AS a)) AS p;
1302+
WITH T AS (SELECT 1 AS X), S AS (SELECT 1 AS ID, 2 AS V, 'a' AS K) SELECT T.X AS X, P.$2 AS _COL_1 FROM T AS T CROSS JOIN S AS S PIVOT(SUM(S.V) FOR S.K IN ('a' AS A)) AS P;
1303+
1304+
# title: preserve Snowflake positional reference from unknown source
1305+
# execute: false
1306+
# dialect: snowflake
1307+
SELECT t.$1 FROM source AS t;
1308+
SELECT T.$1 AS _COL_0 FROM SOURCE AS T;
1309+
1310+
# title: preserve Snowflake positional reference from star source
1311+
# execute: false
1312+
# dialect: snowflake
1313+
WITH t AS (SELECT * FROM source) SELECT t.$1 FROM t;
1314+
WITH T AS (SELECT * FROM SOURCE AS SOURCE) SELECT T.$1 AS _COL_0 FROM T AS T;
1315+
1316+
# title: preserve Snowflake positional reference with table alias column list
1317+
# execute: false
1318+
# dialect: snowflake
1319+
SELECT x.$1 FROM x AS x(alias_name);
1320+
SELECT X.$1 AS _COL_0 FROM X AS X(ALIAS_NAME);

tests/test_optimizer.py

Lines changed: 9 additions & 144 deletions
Original file line numberDiff line numberDiff line change
@@ -1106,157 +1106,22 @@ def test_qualify_columns_struct_star_expansion_types(self):
11061106
)
11071107
self.assertEqual(qualified.selects[0].type.sql("bigquery"), "INT64")
11081108

1109-
def test_qualify_columns_pivot_with_unknown_source(self):
1110-
self.assertEqual(
1111-
qualify(
1112-
parse_one("SELECT u.z FROM u PIVOT(SUM(f) FOR h IN ('x', 'y')) AS u"),
1113-
quote_identifiers=False,
1114-
).sql(),
1115-
"SELECT u.z AS z FROM u AS u PIVOT(SUM(u.f) FOR u.h IN ('x', 'y')) AS u",
1116-
)
1117-
1118-
def test_qualify_columns_pivot_with_known_source(self):
1119-
with self.assertRaisesRegex(OptimizeError, "Unknown column: q"):
1120-
qualify(
1121-
parse_one("SELECT u.q FROM u PIVOT(SUM(f) FOR h IN ('x', 'y')) AS u"),
1122-
schema={"u": {"f": "INT", "h": "TEXT", "z": "INT"}},
1123-
quote_identifiers=False,
1124-
)
1125-
1126-
def test_qualify_snowflake_positional_columns(self):
1127-
expression = parse_one(
1128-
"""
1129-
WITH t AS (
1130-
SELECT PARSE_JSON('{"id": 1}') AS data, 'file' AS filepath
1131-
)
1132-
SELECT t.$1:"id"::INT AS id, t.$2 AS filename
1133-
FROM t
1134-
""",
1135-
dialect="snowflake",
1136-
)
1137-
1138-
self.assertEqual(
1139-
qualify(expression, dialect="snowflake").sql("snowflake"),
1140-
"""WITH "T" AS (SELECT PARSE_JSON('{"id": 1}') AS "DATA", 'file' AS "FILEPATH") """
1141-
"""SELECT CAST(GET_PATH("T"."DATA", 'id') AS INT) AS "ID", """
1142-
""""T"."FILEPATH" AS "FILENAME" FROM "T" AS "T\"""",
1143-
)
1144-
1109+
def test_qualify_snowflake_positional_column_with_visible_schema(self):
11451110
visible_schema = MappingSchema(
11461111
{"t": {"hidden": "INT", "HAS SPACE": "INT"}},
11471112
visible={"T": {"HAS SPACE"}},
11481113
dialect="snowflake",
11491114
)
1150-
cases = (
1151-
(
1152-
'WITH t AS (SELECT 1 AS "lower") SELECT t.$1 FROM t',
1153-
'WITH T AS (SELECT 1 AS "lower") SELECT T."lower" AS "lower" FROM T AS T',
1154-
{},
1155-
),
1156-
(
1157-
"SELECT t.$1 FROM t",
1158-
'SELECT T."HAS SPACE" AS "HAS SPACE" FROM T AS T',
1159-
{"schema": visible_schema},
1160-
),
1161-
)
1162-
1163-
for sql, expected, options in cases:
1164-
with self.subTest(sql=sql):
1165-
self.assertEqual(
1166-
qualify(
1167-
parse_one(sql, dialect="snowflake"),
1168-
dialect="snowflake",
1169-
quote_identifiers=False,
1170-
**options,
1171-
).sql("snowflake"),
1172-
expected,
1173-
)
1174-
1175-
def test_qualify_snowflake_positional_columns_preserve_unresolved(self):
1176-
cases = (
1177-
(
1178-
"duplicate output name",
1179-
"WITH t AS (SELECT 1 AS x, 2 AS x, 3 AS y) SELECT t.$2, t.$3 FROM t",
1180-
("T.$2 AS _COL_0", "T.Y AS Y"),
1181-
{},
1182-
),
1183-
(
1184-
"CTE column aliases",
1185-
'WITH t("lower") AS (SELECT 1) SELECT t.$1 FROM t',
1186-
("T.$1 AS _COL_0",),
1187-
{},
1188-
),
1189-
(
1190-
"UNPIVOT",
1191-
'WITH t AS (SELECT 1 AS "SELECT", 2 AS a, 3 AS b) '
1192-
"SELECT t.$1 FROM t UNPIVOT(v FOR k IN (a, b))",
1193-
("T.$1 AS _COL_0",),
1194-
{},
1195-
),
1196-
(
1197-
"PIVOT",
1198-
"WITH t AS (SELECT 1 AS id, 2 AS v, 'a' AS k) "
1199-
"SELECT p.$2 FROM t PIVOT(SUM(v) FOR k IN ('a' AS \"SELECT\")) AS p",
1200-
("P.$2 AS _COL_0",),
1201-
{},
1202-
),
1203-
(
1204-
"aliasless PIVOT",
1205-
"WITH t AS (SELECT 1 AS id, 2 AS v, 'a' AS k) "
1206-
"SELECT t.$2 FROM t PIVOT(SUM(v) FOR k IN ('a' AS a))",
1207-
("_0.$2 AS _COL_0",),
1208-
{},
1209-
),
1210-
(
1211-
"chained aliasless PIVOT",
1212-
"WITH t AS (SELECT 1 AS id, 2 AS v, 'a' AS k) "
1213-
"SELECT t.$1 FROM t PIVOT(SUM(v) FOR k IN ('a' AS a)) "
1214-
"UNPIVOT(v2 FOR k2 IN (a))",
1215-
("T.$1 AS _COL_0",),
1216-
{},
1217-
),
1218-
(
1219-
"unrelated PIVOT",
1220-
"WITH t AS (SELECT 1 AS x), "
1221-
"s AS (SELECT 1 AS id, 2 AS v, 'a' AS k) "
1222-
"SELECT t.$1, p.$2 FROM t CROSS JOIN "
1223-
"s PIVOT(SUM(v) FOR k IN ('a' AS a)) AS p",
1224-
("T.X AS X", "P.$2 AS _COL_1"),
1225-
{},
1226-
),
1227-
(
1228-
"unknown source",
1229-
"SELECT t.$1 FROM source AS t",
1230-
("T.$1 AS _COL_0",),
1231-
{},
1232-
),
1233-
(
1234-
"star source",
1235-
"WITH t AS (SELECT * FROM source) SELECT t.$1 FROM t",
1236-
("T.$1 AS _COL_0",),
1237-
{},
1238-
),
1239-
(
1240-
"table alias column list",
1241-
"SELECT t.$1 FROM source AS t(alias_name)",
1242-
("T.$1 AS _COL_0",),
1243-
{"schema": {"source": {"x": "INT"}}},
1244-
),
1115+
self.assertEqual(
1116+
qualify(
1117+
parse_one("SELECT t.$1 FROM t", dialect="snowflake"),
1118+
dialect="snowflake",
1119+
quote_identifiers=False,
1120+
schema=visible_schema,
1121+
).sql("snowflake"),
1122+
'SELECT T."HAS SPACE" AS "HAS SPACE" FROM T AS T',
12451123
)
12461124

1247-
for name, sql, expected, options in cases:
1248-
with self.subTest(name=name):
1249-
expression = qualify(
1250-
parse_one(sql, dialect="snowflake"),
1251-
dialect="snowflake",
1252-
quote_identifiers=False,
1253-
**options,
1254-
)
1255-
self.assertEqual(
1256-
tuple(selection.sql("snowflake") for selection in expression.selects),
1257-
expected,
1258-
)
1259-
12601125
def test_qualify_positional_columns_is_snowflake_only(self):
12611126
expression = parse_one(
12621127
"WITH t AS (SELECT 1 AS a) SELECT t.$1 FROM t",

0 commit comments

Comments
 (0)