Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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()

# Align unlimited execute with other connectors (mysql/mssql/etc.):
# strip a terminating `;` before send. Limited composition still needs
# a clean inner SQL so `;` cannot break the subquery 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