Skip to content

Commit 26d8077

Browse files
chore: improve SQL parsing (#26767)
1 parent a75bb76 commit 26d8077

27 files changed

Lines changed: 394 additions & 196 deletions

File tree

superset-frontend/cypress-base/cypress/e2e/explore/AdhocMetrics.test.ts

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ describe('AdhocMetrics', () => {
2525
});
2626

2727
it('Clear metric and set simple adhoc metric', () => {
28-
const metric = 'sum(num_girls)';
28+
const metric = 'SUM(num_girls)';
2929
const metricName = 'Sum Girls';
3030
cy.get('[data-test=metrics]')
3131
.find('[data-test="remove-control-button"]')

superset-frontend/cypress-base/cypress/e2e/explore/visualizations/table.test.ts

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -100,7 +100,7 @@ describe('Visualization > Table', () => {
100100
});
101101
cy.verifySliceSuccess({
102102
waitAlias: '@chartData',
103-
querySubstring: /group by.*name/i,
103+
querySubstring: /group by\n.*name/i,
104104
chartSelector: 'table',
105105
});
106106
});
@@ -246,7 +246,7 @@ describe('Visualization > Table', () => {
246246
cy.visitChartByParams(formData);
247247
cy.verifySliceSuccess({
248248
waitAlias: '@chartData',
249-
querySubstring: /group by.*state/i,
249+
querySubstring: /group by\n.*state/i,
250250
chartSelector: 'table',
251251
});
252252
cy.get('td').contains(/\d*%/);

superset-frontend/src/SqlLab/actions/sqlLab.js

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -921,6 +921,7 @@ export function formatQuery(queryEditor) {
921921
const { sql } = getUpToDateQuery(getState(), queryEditor);
922922
return SupersetClient.post({
923923
endpoint: `/api/v1/sqllab/format_sql/`,
924+
// TODO (betodealmeida): pass engine as a parameter for better formatting
924925
body: JSON.stringify({ sql }),
925926
headers: { 'Content-Type': 'application/json' },
926927
}).then(({ json }) => {

superset/connectors/sqla/models.py

Lines changed: 3 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,6 @@
3333
import numpy as np
3434
import pandas as pd
3535
import sqlalchemy as sa
36-
import sqlparse
3736
from flask import escape, Markup
3837
from flask_appbuilder import Model
3938
from flask_appbuilder.security.sqla.models import User
@@ -104,7 +103,6 @@
104103
ExploreMixin,
105104
ImportExportMixin,
106105
QueryResult,
107-
QueryStringExtended,
108106
validate_adhoc_subquery,
109107
)
110108
from superset.models.slice import Slice
@@ -1099,7 +1097,9 @@ def _process_sql_expression(
10991097

11001098

11011099
class SqlaTable(
1102-
Model, BaseDatasource, ExploreMixin
1100+
Model,
1101+
BaseDatasource,
1102+
ExploreMixin,
11031103
): # pylint: disable=too-many-public-methods
11041104
"""An ORM object for SqlAlchemy table references"""
11051105

@@ -1413,26 +1413,6 @@ def mutate_query_from_config(self, sql: str) -> str:
14131413
def get_template_processor(self, **kwargs: Any) -> BaseTemplateProcessor:
14141414
return get_template_processor(table=self, database=self.database, **kwargs)
14151415

1416-
def get_query_str_extended(
1417-
self,
1418-
query_obj: QueryObjectDict,
1419-
mutate: bool = True,
1420-
) -> QueryStringExtended:
1421-
sqlaq = self.get_sqla_query(**query_obj)
1422-
sql = self.database.compile_sqla_query(sqlaq.sqla_query)
1423-
sql = self._apply_cte(sql, sqlaq.cte)
1424-
sql = sqlparse.format(sql, reindent=True)
1425-
if mutate:
1426-
sql = self.mutate_query_from_config(sql)
1427-
return QueryStringExtended(
1428-
applied_template_filters=sqlaq.applied_template_filters,
1429-
applied_filter_columns=sqlaq.applied_filter_columns,
1430-
rejected_filter_columns=sqlaq.rejected_filter_columns,
1431-
labels_expected=sqlaq.labels_expected,
1432-
prequeries=sqlaq.prequeries,
1433-
sql=sql,
1434-
)
1435-
14361416
def get_query_str(self, query_obj: QueryObjectDict) -> str:
14371417
query_str_ext = self.get_query_str_extended(query_obj)
14381418
all_queries = query_str_ext.prequeries + [query_str_ext.sql]
@@ -1474,33 +1454,6 @@ def get_from_clause(
14741454

14751455
return from_clause, cte
14761456

1477-
def get_rendered_sql(
1478-
self, template_processor: BaseTemplateProcessor | None = None
1479-
) -> str:
1480-
"""
1481-
Render sql with template engine (Jinja).
1482-
"""
1483-
1484-
sql = self.sql
1485-
if template_processor:
1486-
try:
1487-
sql = template_processor.process_template(sql)
1488-
except TemplateError as ex:
1489-
raise QueryObjectValidationError(
1490-
_(
1491-
"Error while rendering virtual dataset query: %(msg)s",
1492-
msg=ex.message,
1493-
)
1494-
) from ex
1495-
sql = sqlparse.format(sql.strip("\t\r\n; "), strip_comments=True)
1496-
if not sql:
1497-
raise QueryObjectValidationError(_("Virtual dataset query cannot be empty"))
1498-
if len(sqlparse.split(sql)) > 1:
1499-
raise QueryObjectValidationError(
1500-
_("Virtual dataset query cannot consist of multiple statements")
1501-
)
1502-
return sql
1503-
15041457
def adhoc_metric_to_sqla(
15051458
self,
15061459
metric: AdhocMetric,

superset/db_engine_specs/base.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,7 @@
5959
from superset.constants import TimeGrain as TimeGrainConstants
6060
from superset.databases.utils import make_url_safe
6161
from superset.errors import ErrorLevel, SupersetError, SupersetErrorType
62-
from superset.sql_parse import ParsedQuery, Table
62+
from superset.sql_parse import ParsedQuery, SQLScript, Table
6363
from superset.superset_typing import ResultSetColumnType, SQLAColumnType
6464
from superset.utils import core as utils
6565
from superset.utils.core import ColumnSpec, GenericDataType
@@ -1448,7 +1448,7 @@ def select_star( # pylint: disable=too-many-arguments,too-many-locals
14481448
qry = partition_query
14491449
sql = database.compile_sqla_query(qry)
14501450
if indent:
1451-
sql = sqlparse.format(sql, reindent=True)
1451+
sql = SQLScript(sql, engine=cls.engine).format()
14521452
return sql
14531453

14541454
@classmethod

superset/db_engine_specs/postgres.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,6 @@
2424
from re import Pattern
2525
from typing import Any, TYPE_CHECKING
2626

27-
import sqlparse
2827
from flask_babel import gettext as __
2928
from sqlalchemy.dialects.postgresql import DOUBLE_PRECISION, ENUM, JSON
3029
from sqlalchemy.dialects.postgresql.base import PGInspector
@@ -37,6 +36,7 @@
3736
from superset.errors import ErrorLevel, SupersetError, SupersetErrorType
3837
from superset.exceptions import SupersetException, SupersetSecurityException
3938
from superset.models.sql_lab import Query
39+
from superset.sql_parse import SQLScript
4040
from superset.utils import core as utils
4141
from superset.utils.core import GenericDataType
4242

@@ -281,8 +281,9 @@ def get_default_schema_for_query(
281281
This method simply uses the parent method after checking that there are no
282282
malicious path setting in the query.
283283
"""
284-
sql = sqlparse.format(query.sql, strip_comments=True)
285-
if re.search(r"set\s+search_path\s*=", sql, re.IGNORECASE):
284+
script = SQLScript(query.sql, engine=cls.engine)
285+
settings = script.get_settings()
286+
if "search_path" in settings:
286287
raise SupersetSecurityException(
287288
SupersetError(
288289
error_type=SupersetErrorType.QUERY_SECURITY_ACCESS_ERROR,

superset/errors.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,7 @@ class SupersetErrorType(StrEnum):
8383
RESULTS_BACKEND_ERROR = "RESULTS_BACKEND_ERROR"
8484
ASYNC_WORKERS_ERROR = "ASYNC_WORKERS_ERROR"
8585
ADHOC_SUBQUERY_NOT_ALLOWED_ERROR = "ADHOC_SUBQUERY_NOT_ALLOWED_ERROR"
86+
INVALID_SQL_ERROR = "INVALID_SQL_ERROR"
8687

8788
# Generic errors
8889
GENERIC_COMMAND_ERROR = "GENERIC_COMMAND_ERROR"
@@ -176,6 +177,7 @@ class SupersetErrorType(StrEnum):
176177
SupersetErrorType.INVALID_PAYLOAD_SCHEMA_ERROR: [1020],
177178
SupersetErrorType.INVALID_CTAS_QUERY_ERROR: [1023],
178179
SupersetErrorType.INVALID_CVAS_QUERY_ERROR: [1024, 1025],
180+
SupersetErrorType.INVALID_SQL_ERROR: [1003],
179181
SupersetErrorType.SQLLAB_TIMEOUT_ERROR: [1026, 1027],
180182
SupersetErrorType.OBJECT_DOES_NOT_EXIST_ERROR: [1029],
181183
SupersetErrorType.SYNTAX_ERROR: [1030],

superset/exceptions.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -295,3 +295,20 @@ def __init__(self, exc: ValidationError, payload: dict[str, Any]):
295295
extra={"messages": exc.messages, "payload": payload},
296296
)
297297
super().__init__(error)
298+
299+
300+
class SupersetParseError(SupersetErrorException):
301+
"""
302+
Exception to be raised when we fail to parse SQL.
303+
"""
304+
305+
status = 422
306+
307+
def __init__(self, sql: str, engine: Optional[str] = None):
308+
error = SupersetError(
309+
message=_("The SQL is invalid and cannot be parsed."),
310+
error_type=SupersetErrorType.INVALID_SQL_ERROR,
311+
level=ErrorLevel.ERROR,
312+
extra={"sql": sql, "engine": engine},
313+
)
314+
super().__init__(error)

superset/models/helpers.py

Lines changed: 20 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,7 @@
6464
ColumnNotFoundException,
6565
QueryClauseValidationException,
6666
QueryObjectValidationError,
67+
SupersetParseError,
6768
SupersetSecurityException,
6869
)
6970
from superset.extensions import feature_flag_manager
@@ -73,6 +74,8 @@
7374
insert_rls_in_predicate,
7475
ParsedQuery,
7576
sanitize_clause,
77+
SQLScript,
78+
SQLStatement,
7679
)
7780
from superset.superset_typing import (
7881
AdhocMetric,
@@ -901,12 +904,18 @@ def _apply_cte(sql: str, cte: Optional[str]) -> str:
901904
return sql
902905

903906
def get_query_str_extended(
904-
self, query_obj: QueryObjectDict, mutate: bool = True
907+
self,
908+
query_obj: QueryObjectDict,
909+
mutate: bool = True,
905910
) -> QueryStringExtended:
906911
sqlaq = self.get_sqla_query(**query_obj)
907912
sql = self.database.compile_sqla_query(sqlaq.sqla_query)
908913
sql = self._apply_cte(sql, sqlaq.cte)
909-
sql = sqlparse.format(sql, reindent=True)
914+
try:
915+
sql = SQLStatement(sql, engine=self.db_engine_spec.engine).format()
916+
except SupersetParseError:
917+
logger.warning("Unable to parse SQL to format it, passing it as-is")
918+
910919
if mutate:
911920
sql = self.mutate_query_from_config(sql)
912921
return QueryStringExtended(
@@ -1054,7 +1063,8 @@ def assign_column_label(df: pd.DataFrame) -> Optional[pd.DataFrame]:
10541063
)
10551064

10561065
def get_rendered_sql(
1057-
self, template_processor: Optional[BaseTemplateProcessor] = None
1066+
self,
1067+
template_processor: Optional[BaseTemplateProcessor] = None,
10581068
) -> str:
10591069
"""
10601070
Render sql with template engine (Jinja).
@@ -1071,13 +1081,16 @@ def get_rendered_sql(
10711081
msg=ex.message,
10721082
)
10731083
) from ex
1074-
sql = sqlparse.format(sql.strip("\t\r\n; "), strip_comments=True)
1075-
if not sql:
1076-
raise QueryObjectValidationError(_("Virtual dataset query cannot be empty"))
1077-
if len(sqlparse.split(sql)) > 1:
1084+
1085+
script = SQLScript(sql.strip("\t\r\n; "), engine=self.db_engine_spec.engine)
1086+
if len(script.statements) > 1:
10781087
raise QueryObjectValidationError(
10791088
_("Virtual dataset query cannot consist of multiple statements")
10801089
)
1090+
1091+
sql = script.statements[0].format(comments=False)
1092+
if not sql:
1093+
raise QueryObjectValidationError(_("Virtual dataset query cannot be empty"))
10811094
return sql
10821095

10831096
def text(self, clause: str) -> TextClause:

0 commit comments

Comments
 (0)