From 683d3c6446fec787249f01ab50eaae81ed3f933d Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Fri, 24 Jul 2026 02:10:26 -0400 Subject: [PATCH 1/8] fix(spark): push LIMIT into SQL instead of client slice Avoid full result materialization before Arrow slice; wrap dry_run in LIMIT 0 subquery after stripping trailing semicolons. --- core/wren/src/wren/connector/spark.py | 18 ++++++++++------ core/wren/tests/unit/test_spark_semicolon.py | 22 +++++++++++++------- 2 files changed, 26 insertions(+), 14 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index 6612d807ec..08f5d61131 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -22,20 +22,26 @@ def _create_session(self): ) def query(self, sql: str, limit: int | None = None) -> pa.Table: - df = self.connection.sql(strip_trailing_semicolon(sql)).toPandas() + # Push LIMIT into Spark SQL so the engine does not materialize the full + # result only for a client-side slice. Strip trailing ``;`` first so the + # outer subscript form stays valid. + cleaned = strip_trailing_semicolon(sql) + if limit is not None: + cleaned = f"SELECT * FROM ({cleaned}) AS _q LIMIT {int(limit)}" + df = self.connection.sql(cleaned).toPandas() if hasattr(df, "attrs") and df.attrs: df.attrs = { k: v for k, v in df.attrs.items() if k not in ("metrics", "observed_metrics") } - arrow_table = pa.Table.from_pandas(df) - if limit is not None: - arrow_table = arrow_table.slice(0, limit) - return arrow_table + return pa.Table.from_pandas(df) def dry_run(self, sql: str) -> None: - self.connection.sql(strip_trailing_semicolon(sql)).limit(0).count() + # Prefer a LIMIT 0 subquery wrapper (like other connectors) so EXPLAIN + # is unnecessary and a trailing semicolon cannot break Spark SQL. + cleaned = strip_trailing_semicolon(sql) + self.connection.sql(f"SELECT * FROM ({cleaned}) AS _q LIMIT 0").count() def close(self) -> None: if self._closed: diff --git a/core/wren/tests/unit/test_spark_semicolon.py b/core/wren/tests/unit/test_spark_semicolon.py index 4d1f5e9af0..008abb8e3c 100644 --- a/core/wren/tests/unit/test_spark_semicolon.py +++ b/core/wren/tests/unit/test_spark_semicolon.py @@ -1,7 +1,9 @@ -"""Trailing-semicolon stripping for the Spark connector (mocked session).""" +"""Trailing-semicolon stripping + LIMIT pushdown for the Spark connector.""" from unittest.mock import MagicMock +import pandas as pd + from wren.connector.base import strip_trailing_semicolon from wren.connector.spark import SparkConnector @@ -16,19 +18,23 @@ def _make_mock_connector() -> tuple[SparkConnector, MagicMock]: def test_query_strips_trailing_semicolon_before_sql() -> None: connector, session = _make_mock_connector() - # pandas DF mock for pa.Table.from_pandas - import pandas as pd - session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": [1, 2, 3]}) - connector.query("SELECT 1;", limit=2) + connector.query("SELECT 1;") session.sql.assert_called_once_with("SELECT 1") -def test_dry_run_strips_trailing_semicolon() -> None: +def test_query_pushes_limit_into_sql_after_strip() -> None: + connector, session = _make_mock_connector() + session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": [1, 2]}) + connector.query("SELECT 1 AS x;", limit=2) + session.sql.assert_called_once_with("SELECT * FROM (SELECT 1 AS x) AS _q LIMIT 2") + + +def test_dry_run_wraps_limit_zero_after_strip() -> None: connector, session = _make_mock_connector() connector.dry_run("SELECT 1; \n") - session.sql.assert_called_once_with("SELECT 1") - session.sql.return_value.limit.assert_called_once_with(0) + session.sql.assert_called_once_with("SELECT * FROM (SELECT 1) AS _q LIMIT 0") + session.sql.return_value.count.assert_called_once_with() def test_helper_preserves_semicolon_inside_string_literal() -> None: From 10b46d2bc74fe405bb43209c4399fdd8c5ac0bdb Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Fri, 24 Jul 2026 02:15:11 -0400 Subject: [PATCH 2/8] fix(spark): reject negative query limits before SQL construction --- core/wren/src/wren/connector/spark.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index 08f5d61131..1119e6c3d8 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -27,7 +27,10 @@ def query(self, sql: str, limit: int | None = None) -> pa.Table: # outer subscript form stays valid. cleaned = strip_trailing_semicolon(sql) if limit is not None: - cleaned = f"SELECT * FROM ({cleaned}) AS _q LIMIT {int(limit)}" + coerced = int(limit) + if coerced < 0: + raise ValueError(f"limit must be non-negative, got {coerced}") + cleaned = f"SELECT * FROM ({cleaned}) AS _q LIMIT {coerced}" df = self.connection.sql(cleaned).toPandas() if hasattr(df, "attrs") and df.attrs: df.attrs = { From 58b4a7bf415897a079c546715f6acd686c4db7b7 Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Thu, 30 Jul 2026 23:10:00 -0400 Subject: [PATCH 3/8] fix(spark): keep dry_run on DataFrame API to preserve non-subquery validation --- core/wren/src/wren/connector/spark.py | 7 ++++--- core/wren/tests/unit/test_spark_semicolon.py | 9 ++++++--- 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index 1119e6c3d8..47fb1bbb24 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -41,10 +41,11 @@ def query(self, sql: str, limit: int | None = None) -> pa.Table: return pa.Table.from_pandas(df) def dry_run(self, sql: str) -> None: - # Prefer a LIMIT 0 subquery wrapper (like other connectors) so EXPLAIN - # is unnecessary and a trailing semicolon cannot break Spark SQL. + # Validate via the DataFrame API so statements that are not legal as a + # subquery (SHOW TABLES, DESCRIBE, ...) are still accepted, exactly as + # before. Only strip a trailing ``;`` so it cannot break Spark SQL. cleaned = strip_trailing_semicolon(sql) - self.connection.sql(f"SELECT * FROM ({cleaned}) AS _q LIMIT 0").count() + self.connection.sql(cleaned).limit(0).count() def close(self) -> None: if self._closed: diff --git a/core/wren/tests/unit/test_spark_semicolon.py b/core/wren/tests/unit/test_spark_semicolon.py index 008abb8e3c..b4bf54caad 100644 --- a/core/wren/tests/unit/test_spark_semicolon.py +++ b/core/wren/tests/unit/test_spark_semicolon.py @@ -30,11 +30,14 @@ def test_query_pushes_limit_into_sql_after_strip() -> None: session.sql.assert_called_once_with("SELECT * FROM (SELECT 1 AS x) AS _q LIMIT 2") -def test_dry_run_wraps_limit_zero_after_strip() -> None: +def test_dry_run_validates_via_dataframe_after_strip() -> None: connector, session = _make_mock_connector() connector.dry_run("SELECT 1; \n") - session.sql.assert_called_once_with("SELECT * FROM (SELECT 1) AS _q LIMIT 0") - session.sql.return_value.count.assert_called_once_with() + # Validation stays on the DataFrame API so non-subquery statements + # (SHOW TABLES, DESCRIBE, ...) remain valid; only the trailing ; is stripped. + session.sql.assert_called_once_with("SELECT 1") + session.sql.return_value.limit.assert_called_once_with(0) + session.sql.return_value.limit.return_value.count.assert_called_once_with() def test_helper_preserves_semicolon_inside_string_literal() -> None: From 79c1465b499a5a0962f88c9fd290b1f47b2c9799 Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Sun, 2 Aug 2026 22:36:16 -0400 Subject: [PATCH 4/8] fix(spark): multiline LIMIT wrap + non-subqueryable fallback Address goldmedal review on #2574: - Mirror snowflake multiline wrap so trailing `--` comments cannot swallow the closing paren/alias/LIMIT - When limit is set for SHOW/DESCRIBE/etc., keep DataFrame client slice (MCP always passes DEFAULT_ROW_LIMIT) - Cover comment wrap, limit=0, negative limit, and SHOW TABLES path --- core/wren/src/wren/connector/spark.py | 48 +++++++++++++++++--- core/wren/tests/unit/test_spark_semicolon.py | 39 +++++++++++++++- 2 files changed, 79 insertions(+), 8 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index 47fb1bbb24..7c73657a66 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -1,8 +1,30 @@ +import re + import pyarrow as pa from wren.connector.base import ConnectorABC, strip_trailing_semicolon from wren.model import SparkConnectionInfo +# Statements that are not legal as a subquery on Spark. When a limit is +# supplied for these, fall back to the DataFrame client-side slice so MCP +# `run_sql` (which always passes a default limit) keeps working as on main. +_NON_SUBQUERYABLE = re.compile( + r"^\s*(SHOW|DESCRIBE|DESC|EXPLAIN|USE|SET|RESET|CACHE|UNCACHE|CLEAR|REFRESH|" + r"MSCK|ANALYZE|LIST|ADD|REMOVE|GET|PUT|DFS|CREATE|DROP|ALTER|TRUNCATE|INSERT|" + r"UPDATE|DELETE|MERGE|LOAD|WITH\s+.*\bINSERT\b)\b", + re.IGNORECASE | re.DOTALL, +) + + +def _coerce_limit(limit: int | None) -> int | None: + """Validate and coerce a user-supplied ``limit`` to a non-negative ``int``.""" + if limit is None: + return None + coerced = int(limit) + if coerced < 0: + raise ValueError(f"limit must be non-negative, got {coerced}") + return coerced + class SparkConnector(ConnectorABC): def __init__(self, connection_info: SparkConnectionInfo): @@ -24,14 +46,26 @@ def _create_session(self): def query(self, sql: str, limit: int | None = None) -> pa.Table: # Push LIMIT into Spark SQL so the engine does not materialize the full # result only for a client-side slice. Strip trailing ``;`` first so the - # outer subscript form stays valid. + # outer form stays valid. cleaned = strip_trailing_semicolon(sql) - if limit is not None: - coerced = int(limit) - if coerced < 0: - raise ValueError(f"limit must be non-negative, got {coerced}") - cleaned = f"SELECT * FROM ({cleaned}) AS _q LIMIT {coerced}" - df = self.connection.sql(cleaned).toPandas() + coerced = _coerce_limit(limit) + if coerced is not None and not _NON_SUBQUERYABLE.match(cleaned): + # Place the user SQL on its own line so a trailing line comment + # (`-- ...`) cannot swallow the closing paren, alias, or LIMIT. + # Mirrors snowflake.py. + executed = ( + "SELECT * FROM (\n" + f"{cleaned}\n" + f") AS _q LIMIT {coerced}" + ) + df = self.connection.sql(executed).toPandas() + else: + # Unlimited path, or non-subqueryable statements (SHOW/DESCRIBE/…). + # Preserve main behaviour: DataFrame API + optional client slice. + frame = self.connection.sql(cleaned) + if coerced is not None: + frame = frame.limit(coerced) + df = frame.toPandas() if hasattr(df, "attrs") and df.attrs: df.attrs = { k: v diff --git a/core/wren/tests/unit/test_spark_semicolon.py b/core/wren/tests/unit/test_spark_semicolon.py index b4bf54caad..dd8f307d3c 100644 --- a/core/wren/tests/unit/test_spark_semicolon.py +++ b/core/wren/tests/unit/test_spark_semicolon.py @@ -3,6 +3,7 @@ from unittest.mock import MagicMock import pandas as pd +import pytest from wren.connector.base import strip_trailing_semicolon from wren.connector.spark import SparkConnector @@ -27,7 +28,43 @@ def test_query_pushes_limit_into_sql_after_strip() -> None: connector, session = _make_mock_connector() session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": [1, 2]}) connector.query("SELECT 1 AS x;", limit=2) - session.sql.assert_called_once_with("SELECT * FROM (SELECT 1 AS x) AS _q LIMIT 2") + session.sql.assert_called_once_with( + "SELECT * FROM (\nSELECT 1 AS x\n) AS _q LIMIT 2" + ) + + +def test_query_trailing_line_comment_does_not_swallow_wrap() -> None: + connector, session = _make_mock_connector() + session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": [1]}) + connector.query("SELECT 1 -- note", limit=10) + session.sql.assert_called_once_with( + "SELECT * FROM (\nSELECT 1 -- note\n) AS _q LIMIT 10" + ) + + +def test_query_limit_zero_pushes_limit_zero() -> None: + connector, session = _make_mock_connector() + session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": []}) + connector.query("SELECT 1 AS x", limit=0) + session.sql.assert_called_once_with( + "SELECT * FROM (\nSELECT 1 AS x\n) AS _q LIMIT 0" + ) + + +def test_query_negative_limit_raises() -> None: + connector, _session = _make_mock_connector() + with pytest.raises(ValueError, match="non-negative"): + connector.query("SELECT 1", limit=-1) + + +def test_query_show_tables_with_limit_uses_dataframe_slice() -> None: + """Non-subqueryable statements keep DataFrame path (MCP always passes limit).""" + connector, session = _make_mock_connector() + frame = session.sql.return_value + frame.limit.return_value.toPandas.return_value = pd.DataFrame({"tableName": ["t"]}) + connector.query("SHOW TABLES", limit=500) + session.sql.assert_called_once_with("SHOW TABLES") + frame.limit.assert_called_once_with(500) def test_dry_run_validates_via_dataframe_after_strip() -> None: From d9e812d449b719ad493ce9ee28bed8404ccc6fc7 Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Sun, 2 Aug 2026 22:39:07 -0400 Subject: [PATCH 5/8] style(spark): ruff-format multiline LIMIT wrap f-string CI ruff format --check wants the subquery wrap as a single f-string. --- core/wren/src/wren/connector/spark.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index 7c73657a66..5d6de972f4 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -53,11 +53,7 @@ def query(self, sql: str, limit: int | None = None) -> pa.Table: # Place the user SQL on its own line so a trailing line comment # (`-- ...`) cannot swallow the closing paren, alias, or LIMIT. # Mirrors snowflake.py. - executed = ( - "SELECT * FROM (\n" - f"{cleaned}\n" - f") AS _q LIMIT {coerced}" - ) + executed = f"SELECT * FROM (\n{cleaned}\n) AS _q LIMIT {coerced}" df = self.connection.sql(executed).toPandas() else: # Unlimited path, or non-subqueryable statements (SHOW/DESCRIBE/…). From cfc51d047ca3051c7b3a6a5b09e3c276af596c23 Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Sun, 2 Aug 2026 22:41:55 -0400 Subject: [PATCH 6/8] fix(spark): skip leading SQL comments in non-subqueryable check Classify SHOW/DESCRIBE/... after stripping leading -- and /* */ comments so MCP default limits still use the DataFrame path for commented meta SQL. --- core/wren/src/wren/connector/spark.py | 22 ++++++++++++++++++-- core/wren/tests/unit/test_spark_semicolon.py | 11 ++++++++++ 2 files changed, 31 insertions(+), 2 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index 5d6de972f4..e81c783ed5 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -9,12 +9,30 @@ # supplied for these, fall back to the DataFrame client-side slice so MCP # `run_sql` (which always passes a default limit) keeps working as on main. _NON_SUBQUERYABLE = re.compile( - r"^\s*(SHOW|DESCRIBE|DESC|EXPLAIN|USE|SET|RESET|CACHE|UNCACHE|CLEAR|REFRESH|" + r"^(SHOW|DESCRIBE|DESC|EXPLAIN|USE|SET|RESET|CACHE|UNCACHE|CLEAR|REFRESH|" r"MSCK|ANALYZE|LIST|ADD|REMOVE|GET|PUT|DFS|CREATE|DROP|ALTER|TRUNCATE|INSERT|" r"UPDATE|DELETE|MERGE|LOAD|WITH\s+.*\bINSERT\b)\b", re.IGNORECASE | re.DOTALL, ) +# Leading line/block comments (and whitespace) before the first keyword. +_LEADING_SQL_NOISE = re.compile( + r"(?s)^(?:\s|--[^\n]*(?:\n|$)|/\*.*?\*/)+" +) + + +def _strip_leading_sql_comments(sql: str) -> str: + """Remove leading whitespace and SQL comments so keyword classification works.""" + prev = None + while prev != sql: + prev = sql + sql = _LEADING_SQL_NOISE.sub("", sql, count=1) + return sql.lstrip() + + +def _is_non_subqueryable(sql: str) -> bool: + return bool(_NON_SUBQUERYABLE.match(_strip_leading_sql_comments(sql))) + def _coerce_limit(limit: int | None) -> int | None: """Validate and coerce a user-supplied ``limit`` to a non-negative ``int``.""" @@ -49,7 +67,7 @@ def query(self, sql: str, limit: int | None = None) -> pa.Table: # outer form stays valid. cleaned = strip_trailing_semicolon(sql) coerced = _coerce_limit(limit) - if coerced is not None and not _NON_SUBQUERYABLE.match(cleaned): + if coerced is not None and not _is_non_subqueryable(cleaned): # Place the user SQL on its own line so a trailing line comment # (`-- ...`) cannot swallow the closing paren, alias, or LIMIT. # Mirrors snowflake.py. diff --git a/core/wren/tests/unit/test_spark_semicolon.py b/core/wren/tests/unit/test_spark_semicolon.py index dd8f307d3c..ab23a7eca2 100644 --- a/core/wren/tests/unit/test_spark_semicolon.py +++ b/core/wren/tests/unit/test_spark_semicolon.py @@ -67,6 +67,17 @@ def test_query_show_tables_with_limit_uses_dataframe_slice() -> None: frame.limit.assert_called_once_with(500) +def test_query_show_tables_with_leading_comment_uses_dataframe_slice() -> None: + """Leading -- / /* */ comments must not force subquery wrap on SHOW.""" + connector, session = _make_mock_connector() + frame = session.sql.return_value + frame.limit.return_value.toPandas.return_value = pd.DataFrame({"tableName": ["t"]}) + sql = "-- metadata\nSHOW TABLES" + connector.query(sql, limit=500) + session.sql.assert_called_once_with(sql) + frame.limit.assert_called_once_with(500) + + def test_dry_run_validates_via_dataframe_after_strip() -> None: connector, session = _make_mock_connector() connector.dry_run("SELECT 1; \n") From 57a102f92ea79e3a98b9c4701aaf16d116b93a11 Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Sun, 2 Aug 2026 22:44:17 -0400 Subject: [PATCH 7/8] style(spark): ruff-format _LEADING_SQL_NOISE on one line --- core/wren/src/wren/connector/spark.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index e81c783ed5..e786e9b6d1 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -16,9 +16,7 @@ ) # Leading line/block comments (and whitespace) before the first keyword. -_LEADING_SQL_NOISE = re.compile( - r"(?s)^(?:\s|--[^\n]*(?:\n|$)|/\*.*?\*/)+" -) +_LEADING_SQL_NOISE = re.compile(r"(?s)^(?:\s|--[^\n]*(?:\n|$)|/\*.*?\*/)+") def _strip_leading_sql_comments(sql: str) -> str: From 3c63da0632099a9028d2b0a9a4dd604b034b0785 Mon Sep 17 00:00:00 2001 From: Bartok9 <259807879+Bartok9@users.noreply.github.com> Date: Mon, 3 Aug 2026 00:08:47 -0400 Subject: [PATCH 8/8] fix(spark): limit via DataFrame.limit only (drop SQL wrap) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit goldmedal measured DataFrame.limit() as server-side CollectLimit — identical plan to subquery LIMIT wrap. Drop wrap + classifier; fix is limit before toPandas, not string surgery. Keep non-negative coerce. --- core/wren/src/wren/connector/spark.py | 59 +++----------------- core/wren/tests/unit/test_spark_semicolon.py | 53 ++++++------------ 2 files changed, 26 insertions(+), 86 deletions(-) diff --git a/core/wren/src/wren/connector/spark.py b/core/wren/src/wren/connector/spark.py index e786e9b6d1..9ef294c8f7 100644 --- a/core/wren/src/wren/connector/spark.py +++ b/core/wren/src/wren/connector/spark.py @@ -1,36 +1,8 @@ -import re - import pyarrow as pa from wren.connector.base import ConnectorABC, strip_trailing_semicolon from wren.model import SparkConnectionInfo -# Statements that are not legal as a subquery on Spark. When a limit is -# supplied for these, fall back to the DataFrame client-side slice so MCP -# `run_sql` (which always passes a default limit) keeps working as on main. -_NON_SUBQUERYABLE = re.compile( - r"^(SHOW|DESCRIBE|DESC|EXPLAIN|USE|SET|RESET|CACHE|UNCACHE|CLEAR|REFRESH|" - r"MSCK|ANALYZE|LIST|ADD|REMOVE|GET|PUT|DFS|CREATE|DROP|ALTER|TRUNCATE|INSERT|" - r"UPDATE|DELETE|MERGE|LOAD|WITH\s+.*\bINSERT\b)\b", - re.IGNORECASE | re.DOTALL, -) - -# Leading line/block comments (and whitespace) before the first keyword. -_LEADING_SQL_NOISE = re.compile(r"(?s)^(?:\s|--[^\n]*(?:\n|$)|/\*.*?\*/)+") - - -def _strip_leading_sql_comments(sql: str) -> str: - """Remove leading whitespace and SQL comments so keyword classification works.""" - prev = None - while prev != sql: - prev = sql - sql = _LEADING_SQL_NOISE.sub("", sql, count=1) - return sql.lstrip() - - -def _is_non_subqueryable(sql: str) -> bool: - return bool(_NON_SUBQUERYABLE.match(_strip_leading_sql_comments(sql))) - def _coerce_limit(limit: int | None) -> int | None: """Validate and coerce a user-supplied ``limit`` to a non-negative ``int``.""" @@ -60,24 +32,15 @@ def _create_session(self): ) def query(self, sql: str, limit: int | None = None) -> pa.Table: - # Push LIMIT into Spark SQL so the engine does not materialize the full - # result only for a client-side slice. Strip trailing ``;`` first so the - # outer form stays valid. - cleaned = strip_trailing_semicolon(sql) + # Apply limit via DataFrame.limit before toPandas so Spark pushes a + # CollectLimit into the plan (server-side). Avoid post-Arrow slice and + # SQL subquery wraps — both unnecessary on the DataFrame API and the + # latter breaks SHOW/DESCRIBE-style statements. coerced = _coerce_limit(limit) - if coerced is not None and not _is_non_subqueryable(cleaned): - # Place the user SQL on its own line so a trailing line comment - # (`-- ...`) cannot swallow the closing paren, alias, or LIMIT. - # Mirrors snowflake.py. - executed = f"SELECT * FROM (\n{cleaned}\n) AS _q LIMIT {coerced}" - df = self.connection.sql(executed).toPandas() - else: - # Unlimited path, or non-subqueryable statements (SHOW/DESCRIBE/…). - # Preserve main behaviour: DataFrame API + optional client slice. - frame = self.connection.sql(cleaned) - if coerced is not None: - frame = frame.limit(coerced) - df = frame.toPandas() + frame = self.connection.sql(strip_trailing_semicolon(sql)) + if coerced is not None: + frame = frame.limit(coerced) + df = frame.toPandas() if hasattr(df, "attrs") and df.attrs: df.attrs = { k: v @@ -87,11 +50,7 @@ def query(self, sql: str, limit: int | None = None) -> pa.Table: return pa.Table.from_pandas(df) def dry_run(self, sql: str) -> None: - # Validate via the DataFrame API so statements that are not legal as a - # subquery (SHOW TABLES, DESCRIBE, ...) are still accepted, exactly as - # before. Only strip a trailing ``;`` so it cannot break Spark SQL. - cleaned = strip_trailing_semicolon(sql) - self.connection.sql(cleaned).limit(0).count() + self.connection.sql(strip_trailing_semicolon(sql)).limit(0).count() def close(self) -> None: if self._closed: diff --git a/core/wren/tests/unit/test_spark_semicolon.py b/core/wren/tests/unit/test_spark_semicolon.py index ab23a7eca2..bb2149fee5 100644 --- a/core/wren/tests/unit/test_spark_semicolon.py +++ b/core/wren/tests/unit/test_spark_semicolon.py @@ -1,4 +1,4 @@ -"""Trailing-semicolon stripping + LIMIT pushdown for the Spark connector.""" +"""Trailing-semicolon stripping + DataFrame limit for the Spark connector.""" from unittest.mock import MagicMock @@ -24,41 +24,35 @@ def test_query_strips_trailing_semicolon_before_sql() -> None: session.sql.assert_called_once_with("SELECT 1") -def test_query_pushes_limit_into_sql_after_strip() -> None: +def test_query_limit_uses_dataframe_limit_before_to_pandas() -> None: connector, session = _make_mock_connector() - session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": [1, 2]}) + frame = session.sql.return_value + frame.limit.return_value.toPandas.return_value = pd.DataFrame({"x": [1, 2]}) connector.query("SELECT 1 AS x;", limit=2) - session.sql.assert_called_once_with( - "SELECT * FROM (\nSELECT 1 AS x\n) AS _q LIMIT 2" - ) - - -def test_query_trailing_line_comment_does_not_swallow_wrap() -> None: - connector, session = _make_mock_connector() - session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": [1]}) - connector.query("SELECT 1 -- note", limit=10) - session.sql.assert_called_once_with( - "SELECT * FROM (\nSELECT 1 -- note\n) AS _q LIMIT 10" - ) + session.sql.assert_called_once_with("SELECT 1 AS x") + frame.limit.assert_called_once_with(2) + frame.limit.return_value.toPandas.assert_called_once_with() + # Must not call toPandas on the unlimited frame. + frame.toPandas.assert_not_called() -def test_query_limit_zero_pushes_limit_zero() -> None: +def test_query_limit_zero_uses_dataframe_limit_zero() -> None: connector, session = _make_mock_connector() - session.sql.return_value.toPandas.return_value = pd.DataFrame({"x": []}) + frame = session.sql.return_value + frame.limit.return_value.toPandas.return_value = pd.DataFrame({"x": []}) connector.query("SELECT 1 AS x", limit=0) - session.sql.assert_called_once_with( - "SELECT * FROM (\nSELECT 1 AS x\n) AS _q LIMIT 0" - ) + session.sql.assert_called_once_with("SELECT 1 AS x") + frame.limit.assert_called_once_with(0) def test_query_negative_limit_raises() -> None: - connector, _session = _make_mock_connector() + connector, session = _make_mock_connector() with pytest.raises(ValueError, match="non-negative"): connector.query("SELECT 1", limit=-1) + session.sql.assert_not_called() -def test_query_show_tables_with_limit_uses_dataframe_slice() -> None: - """Non-subqueryable statements keep DataFrame path (MCP always passes limit).""" +def test_query_show_tables_with_limit_uses_dataframe_limit() -> None: connector, session = _make_mock_connector() frame = session.sql.return_value frame.limit.return_value.toPandas.return_value = pd.DataFrame({"tableName": ["t"]}) @@ -67,22 +61,9 @@ def test_query_show_tables_with_limit_uses_dataframe_slice() -> None: frame.limit.assert_called_once_with(500) -def test_query_show_tables_with_leading_comment_uses_dataframe_slice() -> None: - """Leading -- / /* */ comments must not force subquery wrap on SHOW.""" - connector, session = _make_mock_connector() - frame = session.sql.return_value - frame.limit.return_value.toPandas.return_value = pd.DataFrame({"tableName": ["t"]}) - sql = "-- metadata\nSHOW TABLES" - connector.query(sql, limit=500) - session.sql.assert_called_once_with(sql) - frame.limit.assert_called_once_with(500) - - def test_dry_run_validates_via_dataframe_after_strip() -> None: connector, session = _make_mock_connector() connector.dry_run("SELECT 1; \n") - # Validation stays on the DataFrame API so non-subquery statements - # (SHOW TABLES, DESCRIBE, ...) remain valid; only the trailing ; is stripped. session.sql.assert_called_once_with("SELECT 1") session.sql.return_value.limit.assert_called_once_with(0) session.sql.return_value.limit.return_value.count.assert_called_once_with()