diff --git a/sqlit/domains/connections/providers/oracle/adapter.py b/sqlit/domains/connections/providers/oracle/adapter.py index 76edc895..2f0503c5 100644 --- a/sqlit/domains/connections/providers/oracle/adapter.py +++ b/sqlit/domains/connections/providers/oracle/adapter.py @@ -2,6 +2,7 @@ from __future__ import annotations +import re from typing import TYPE_CHECKING, Any from sqlit.domains.connections.providers.adapters.base import ( @@ -18,6 +19,31 @@ from sqlit.domains.connections.domain.config import ConnectionConfig +_LEADING_SQL_COMMENTS = re.compile( + r"^\s*(?:(?:--[^\n]*(?:\n|$))|(?:/\*.*?\*/\s*))*", + re.DOTALL, +) +_PLSQL_START = re.compile( + r"^(?:BEGIN|DECLARE)\b|^CREATE\s+(?:OR\s+REPLACE\s+)?" + r"(?:(?:NON)?EDITIONABLE\s+)?" + r"(?:FUNCTION|PACKAGE|PROCEDURE|TRIGGER|TYPE\s+BODY)\b", + re.IGNORECASE, +) + + +def _prepare_statement(query: str) -> str: + """Remove SQL*Plus terminators that python-oracledb does not accept.""" + statement = query.rstrip() + if not statement.endswith(";"): + return query + + without_leading_comments = _LEADING_SQL_COMMENTS.sub("", statement) + if _PLSQL_START.match(without_leading_comments): + return query + + return statement[:-1].rstrip() + + class OracleAdapter(DatabaseAdapter): """Adapter for Oracle Database using oracledb. @@ -339,7 +365,7 @@ def execute_query(self, conn: Any, query: str, max_rows: int | None = None) -> t """Execute a query on Oracle with optional row limit.""" cursor = conn.cursor() try: - cursor.execute(query) + cursor.execute(_prepare_statement(query)) if cursor.description: columns = [col[0] for col in cursor.description] if max_rows is not None: @@ -359,7 +385,7 @@ def execute_non_query(self, conn: Any, query: str) -> int: """Execute a non-query on Oracle.""" cursor = conn.cursor() try: - cursor.execute(query) + cursor.execute(_prepare_statement(query)) rowcount = int(cursor.rowcount) conn.commit() return rowcount diff --git a/tests/connections/providers/oracle/test_statement_execution.py b/tests/connections/providers/oracle/test_statement_execution.py new file mode 100644 index 00000000..73eda739 --- /dev/null +++ b/tests/connections/providers/oracle/test_statement_execution.py @@ -0,0 +1,77 @@ +"""Oracle-specific statement execution behavior.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from sqlit.domains.connections.providers.oracle.adapter import OracleAdapter + + +@pytest.fixture +def adapter() -> OracleAdapter: + return OracleAdapter() + + +@pytest.fixture +def mock_conn() -> MagicMock: + conn = MagicMock() + cursor = conn.cursor.return_value + cursor.description = None + cursor.rowcount = 0 + return conn + + +@pytest.mark.parametrize( + ("query", "expected"), + [ + ("SELECT ';' AS value FROM DUAL;", "SELECT ';' AS value FROM DUAL"), + ("SELECT 1 FROM DUAL; \n", "SELECT 1 FROM DUAL"), + ("SELECT 1 FROM DUAL", "SELECT 1 FROM DUAL"), + ], +) +def test_execute_query_removes_sql_statement_terminator( + adapter: OracleAdapter, + mock_conn: MagicMock, + query: str, + expected: str, +) -> None: + """python-oracledb rejects SQL statements ending in a semicolon.""" + adapter.execute_query(mock_conn, query) + + mock_conn.cursor.return_value.execute.assert_called_once_with(expected) + + +def test_execute_non_query_removes_sql_statement_terminator( + adapter: OracleAdapter, + mock_conn: MagicMock, +) -> None: + """Issue #260: ALTER SESSION must reach python-oracledb without ``;``.""" + adapter.execute_non_query(mock_conn, "ALTER SESSION SET EDITION = V0;") + + mock_conn.cursor.return_value.execute.assert_called_once_with("ALTER SESSION SET EDITION = V0") + + +@pytest.mark.parametrize( + "statement", + [ + "BEGIN NULL; END;", + "DECLARE value NUMBER := 1; BEGIN NULL; END;", + "CREATE OR REPLACE PROCEDURE p AS BEGIN NULL; END;", + "-- setup\nBEGIN NULL; END;", + "/* setup */ CREATE OR REPLACE EDITIONABLE FUNCTION f RETURN NUMBER " + "AS BEGIN RETURN 1; END;", + "CREATE OR REPLACE TYPE BODY t AS MEMBER PROCEDURE p IS " + "BEGIN NULL; END; END;", + ], +) +def test_execute_non_query_preserves_plsql_terminator( + adapter: OracleAdapter, + mock_conn: MagicMock, + statement: str, +) -> None: + """The final semicolon is part of PL/SQL syntax, not a client terminator.""" + adapter.execute_non_query(mock_conn, statement) + + mock_conn.cursor.return_value.execute.assert_called_once_with(statement)