Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions core/wren/src/wren/connector/trino.py
Original file line number Diff line number Diff line change
Expand Up @@ -485,10 +485,12 @@ def query(self, sql: str, limit: int | None = None) -> pa.Table:
limit = coerce_limit(limit)
trino = _import_trino()

# Strip terminating `;` for unlimited execute too — Trino's statement
# path is stricter than most engines about bare terminators (CLI also
# strips client-side). Limited path already strips inside the wrap.
sql = strip_trailing_semicolon(sql)
if limit is not None:
sql = (
f"SELECT * FROM ({strip_trailing_semicolon(sql)}) AS _sub LIMIT {limit}"
)
sql = f"SELECT * FROM ({sql}) AS _sub LIMIT {limit}"
try:
with contextlib.closing(self.connection.cursor()) as cursor:
cursor.execute(sql)
Expand Down
42 changes: 42 additions & 0 deletions core/wren/tests/unit/test_trino_semicolon_unlimited.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
"""Trino unlimited query strips trailing semicolons before execute."""

from __future__ import annotations

from unittest.mock import MagicMock, patch

import pyarrow as pa
import pytest


@pytest.fixture
def connector():
with patch("wren.connector.trino._import_trino") as imp:
mod = MagicMock()
imp.return_value = mod
from wren.connector.trino import TrinoConnector

c = TrinoConnector.__new__(TrinoConnector)
c.connection = MagicMock()
c._closed = False
yield c, mod


def test_query_without_limit_strips_trailing_semicolon(connector) -> None:
c, _mod = connector
cursor = MagicMock()
c.connection.cursor.return_value = cursor
# _build_trino_arrow_table path — mock fetch
with patch("wren.connector.trino._build_trino_arrow_table", return_value=pa.table({"x": [1]})):
c.query("SELECT 1;")
cursor.execute.assert_called_once_with("SELECT 1")


def test_query_with_limit_strips_inside_wrap(connector) -> None:
c, _mod = connector
cursor = MagicMock()
c.connection.cursor.return_value = cursor
with patch("wren.connector.trino._build_trino_arrow_table", return_value=pa.table({"x": [1]})):
c.query("SELECT 1;", limit=5)
sent = cursor.execute.call_args[0][0]
assert "SELECT 1;" not in sent
assert "LIMIT 5" in sent
Loading