fix(datafusion): strip trailing semicolon before subquery-wrapping in query (#2430)

This commit is contained in:
Bartok
2026-07-06 10:11:22 +08:00
committed by GitHub
parent fe92e5b2a2
commit 8929764310
2 changed files with 82 additions and 1 deletions
+21 -1
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import io
import re
from pathlib import Path
import pyarrow as pa
@@ -11,6 +12,22 @@ from wren.connector.base import ConnectorABC
from wren.model import DataFusionConnectionInfo
from wren.model.error import ErrorCode, WrenError
_TRAILING_SEMICOLONS_RE = re.compile(r"[;\s]+\Z")
def _strip_trailing_semicolon(sql: str) -> str:
"""Strip the terminating run of ``;`` characters and surrounding whitespace.
Wrapping user SQL as ``SELECT * FROM ({sql}) AS _q LIMIT N`` breaks when
``sql`` ends in a semicolon — DataFusion rejects ``SELECT 1;`` inside a
subquery (``sql parser error: Expected: an expression, found: ;``). Only
the *terminating* run of semicolons/whitespace is stripped, so semicolons
inside string literals (e.g. ``SELECT 'a;b' FROM t``) are preserved.
Mirrors the postgres/redshift/duckdb connectors, which already strip
before subquery-wrapping.
"""
return _TRAILING_SEMICOLONS_RE.sub("", sql)
class DataFusionConnector(ConnectorABC):
"""DataFusion-native connector for local file analysis.
@@ -30,7 +47,10 @@ class DataFusionConnector(ConnectorABC):
def query(self, sql: str, limit: int | None = None) -> pa.Table:
if limit is not None:
sql = f"SELECT * FROM ({sql}) AS _q LIMIT {int(limit)}"
sql = (
f"SELECT * FROM ({_strip_trailing_semicolon(sql)}) "
f"AS _q LIMIT {int(limit)}"
)
ipc_bytes = self.ctx.query(sql)
reader = ipc.open_stream(io.BytesIO(bytes(ipc_bytes)))
return reader.read_all()
@@ -0,0 +1,61 @@
"""Trailing-semicolon stripping for the DataFusion connector (mocked ctx).
The DataFusion connector wraps user SQL as ``SELECT * FROM (...) AS _q LIMIT N``
when a limit is supplied. If the user SQL ends in a semicolon, DataFusion
rejects ``SELECT 1;`` inside a subquery
(``sql parser error: Expected: an expression, found: ;``). These tests use a
mocked ``ctx`` and assert on the SQL string the connector builds, so no native
runtime file registration is required. Mirrors the postgres/redshift/duckdb
connector semicolon tests.
"""
from unittest.mock import MagicMock
import pyarrow as pa
import pyarrow.ipc as ipc
from wren.connector.datafusion import (
DataFusionConnector,
_strip_trailing_semicolon,
)
def _make_mock_connector() -> tuple[DataFusionConnector, MagicMock]:
"""Build a DataFusionConnector bypassing __init__ (no real runtime)."""
connector = DataFusionConnector.__new__(DataFusionConnector)
ctx = MagicMock()
# ctx.query must return IPC-stream bytes that read back into a table.
empty = pa.table({"x": pa.array([], type=pa.int64())})
sink = pa.BufferOutputStream()
with ipc.new_stream(sink, empty.schema) as writer:
writer.write_table(empty)
ctx.query.return_value = sink.getvalue().to_pybytes()
connector.ctx = ctx
return connector, ctx
def test_query_strips_trailing_semicolon_before_subquery_wrap() -> None:
connector, ctx = _make_mock_connector()
connector.query("SELECT 1;", limit=5)
(sent,), _ = ctx.query.call_args
assert sent == "SELECT * FROM (SELECT 1) AS _q LIMIT 5"
assert ";)" not in sent
def test_query_without_limit_is_unwrapped() -> None:
connector, ctx = _make_mock_connector()
connector.query("SELECT 1;")
(sent,), _ = ctx.query.call_args
# No limit -> no subquery wrapping; passed through verbatim.
assert sent == "SELECT 1;"
def test_helper_preserves_semicolon_inside_string_literal() -> None:
sql = "SELECT 'a;b' AS x"
assert _strip_trailing_semicolon(sql) == sql
def test_helper_no_trailing_semicolon_unchanged() -> None:
assert _strip_trailing_semicolon("SELECT 1") == "SELECT 1"