mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
fix(datafusion): strip trailing semicolon before subquery-wrapping in query (#2430)
This commit is contained in:
@@ -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"
|
||||
Reference in New Issue
Block a user