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
12 changes: 6 additions & 6 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -6,15 +6,15 @@ repos:
- id: check-merge-conflict

- repo: https://github.com/asottile/setup-cfg-fmt
rev: v3.1.0
rev: v3.2.0
hooks:
- id: setup-cfg-fmt
args:
- --include-version-classifiers
- --min-py-version=3.10

- repo: https://github.com/myint/autoflake
rev: v2.3.1
rev: v2.4.0
hooks:
- id: autoflake
args:
Expand All @@ -26,12 +26,12 @@ repos:
- --remove-unused-variables

- repo: https://github.com/PyCQA/isort
rev: "6.1.0"
rev: "9.0.2"
hooks:
- id: isort

- repo: https://github.com/psf/black
rev: 25.11.0
rev: 26.10.0
hooks:
- id: black
language_version: python3.10
Expand All @@ -40,10 +40,10 @@ repos:
rev: v0.4.0
hooks:
- id: black_nbconvert
additional_dependencies: ["setuptools>=69,<82", "black==25.11.0", "tomli"] # 82+ removes pkg_resources
additional_dependencies: ["setuptools>=69,<82", "black==26.10.0", "tomli"] # 82+ removes pkg_resources

- repo: https://github.com/PyCQA/flake8
rev: "7.3.0"
rev: "7.4.1"
hooks:
- id: flake8
additional_dependencies:
Expand Down
12 changes: 4 additions & 8 deletions examples/dbapi.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -179,12 +179,10 @@
"metadata": {},
"outputs": [],
"source": [
"cursor.execute(\n",
" \"\"\"\n",
"cursor.execute(\"\"\"\n",
" select * from test_table where id < 4;\n",
" select * from test_table where id > 2;\n",
"\"\"\"\n",
")\n",
"\"\"\")\n",
"print(\"First query: \", cursor.fetchall())\n",
"assert cursor.nextset()\n",
"print(\"Second query: \", cursor.fetchall())\n",
Expand All @@ -209,13 +207,11 @@
"outputs": [],
"source": [
"try:\n",
" cursor.execute(\n",
" \"\"\"\n",
" cursor.execute(\"\"\"\n",
" select * from test_table where id < 4;\n",
" select * from test_table where wrong_field > 2;\n",
" select * from test_table\n",
" \"\"\"\n",
" )\n",
" \"\"\")\n",
"except OperationalError:\n",
" pass\n",
"cursor.fetchall()"
Expand Down
20 changes: 10 additions & 10 deletions setup.cfg
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ project_urls =
packages = find:
install_requires =
aiorwlock==1.5.1
anyio>=4.12.1
anyio>=4.15.1
appdirs>=1.4.4
appdirs-stubs>=0.2.0
async-generator>=1.10
Expand All @@ -36,8 +36,8 @@ install_requires =
pydantic>=2.13.5,<3.0.0
python-dateutil>=2.9.0.post0
readerwriterlock>=1.0.9
sqlparse==0.5.5
trio>=0.31.0
sqlparse==0.6.0
trio>=0.34.0
truststore>=0.10.4
python_requires = >=3.10
include_package_data = True
Expand All @@ -52,21 +52,21 @@ ciso8601 =
ciso8601==2.3.3
dev =
devtools==0.12.2
mypy>=1.19.1,<2
pre-commit==4.3.0
mypy>=2.4.0,<3
pre-commit==4.6.2
psutil==7.2.2
pyfakefs>=5.10.2,<6
pytest==8.4.2
pyfakefs>=6.2.0,<7
pytest==9.1.1
pytest-cov==7.1.0
pytest-httpx>=0.35.0
pytest-mock==3.15.1
pytest-httpx>=0.36.2
pytest-mock==3.16.0
pytest-timeout==2.4.0
pytest-trio==0.8.0
pytest-xdist==3.8.0
trio-typing[mypy]>=0.10,<0.11
types-cryptography==3.3.23.2
docs =
sphinx>=7.4.7,<10
sphinx>=8.1.3,<10
sphinx-rtd-theme>=3.1.0,<4

[options.package_data]
Expand Down
2 changes: 1 addition & 1 deletion src/firebolt/common/cursor/base_cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -293,7 +293,7 @@ def setoutputsize(self, size: int, column: Optional[int] = None) -> None:
def close(self) -> None:
"""Terminate an ongoing query (if any) and mark connection as closed."""
self._state = CursorState.CLOSED
self.connection._remove_cursor(self) # type:ignore
self.connection._remove_cursor(self) # type: ignore

def __del__(self) -> None:
self.close()
Expand Down
71 changes: 9 additions & 62 deletions src/firebolt/common/statement_formatter.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,73 +27,20 @@
NotSupportedError,
)

_original_change_splitlevel = _StatementSplitter._change_splitlevel

def _patched_change_splitlevel(self, ttype, value): # type: ignore[no-untyped-def]
"""Patched version of StatementSplitter._change_splitlevel.

Fixes CASE...END level tracking outside of CREATE blocks.
See: https://github.com/andialbrecht/sqlparse/pull/839
"""
if ttype is _T.Punctuation and value == "(":
return 1
elif ttype is _T.Punctuation and value == ")":
return -1
elif ttype not in _T.Keyword:
return 0

unified = value.upper()

if ttype is _T.Keyword.DDL and unified.startswith("CREATE"):
self._is_create = True
return 0

if unified == "DECLARE" and self._is_create and self._begin_depth == 0:
self._in_declare = True
return 1

if unified == "BEGIN":
self._begin_depth += 1
self._seen_begin = True
if self._is_create:
return 1
return 0

def _patched_change_splitlevel(self, ttype, value): # type: ignore[no-untyped-def]
"""Track CASE outside BEGIN blocks until sqlparse fixes PR 839."""
# https://github.com/andialbrecht/sqlparse/pull/839
if (
self._seen_begin
and (ttype is _T.Keyword or ttype is _T.Name)
and unified
in (
"TRANSACTION",
"WORK",
"TRAN",
"DISTRIBUTED",
"DEFERRED",
"IMMEDIATE",
"EXCLUSIVE",
)
ttype in _T.Keyword
and value.upper() == "CASE"
and "BEGIN" not in self._block_stack
):
self._begin_depth = max(0, self._begin_depth - 1)
self._seen_begin = False
return 0

if unified == "END":
if not self._in_case:
self._begin_depth = max(0, self._begin_depth - 1)
else:
self._in_case = False
return -1

if unified == "CASE":
self._in_case = True
self._block_stack.append("CASE")
return 1

if unified in ("IF", "FOR", "WHILE") and self._is_create and self._begin_depth > 0:
return 1

if unified in ("END IF", "END FOR", "END WHILE"):
return -1

return 0
return _original_change_splitlevel(self, ttype, value)


setattr(_StatementSplitter, "_change_splitlevel", _patched_change_splitlevel)
Expand Down
2 changes: 0 additions & 2 deletions tests/integration/utils/test_usage_tracker.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,6 @@
]


@mark.xdist_group(name="usage_tracker")
@fixture(scope="module")
def create_cli_mock():
# Cleanup before starting
Expand All @@ -33,7 +32,6 @@ def create_cli_mock():
rmtree(TEST_FOLDER)


@mark.xdist_group(name="usage_tracker")
@fixture(scope="module")
def test_model():
with open(TEST_SCRIPT_MODEL) as f:
Expand Down
17 changes: 17 additions & 0 deletions tests/unit/common/test_typing_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -307,6 +307,23 @@ def test_create_statement_formatter_invalid_version() -> None:
assert "Unsupported version: 3" in str(excinfo.value)


@mark.parametrize(
"expression",
[
"CASE WHEN 1 THEN 'a' ELSE 'b' END",
"CASE WHEN 1 THEN CASE WHEN 2 THEN 'a' END ELSE 'b' END",
],
)
def test_case_followed_by_quoted_semicolon(
formatter: StatementFormatter, expression: str
) -> None:
sql = f"SELECT {expression}, test IN ('foo \\', 'foo;') FROM t; SELECT 2;"
assert formatter.split_format_sql(sql, None) == [
f"SELECT {expression}, test IN ('foo \\', 'foo;') FROM t",
"SELECT 2",
]


def test_patched_change_splitlevel(formatter: StatementFormatter) -> None:
# Testing CREATE, DECLARE, BEGIN, END, CASE, IF, FOR, WHILE
# These exercise _patched_change_splitlevel via split_format_sql which calls parse_sql
Expand Down
Loading