From 2679d94af2234895743085cfd7416f89c414d232 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:12:56 +0000 Subject: [PATCH 01/10] fix: address adapter, fixture, and migration review regressions --- docs/changelog.rst | 31 ++++++- docs/usage/testing.rst | 29 +++++-- sqlspec/adapters/aiomysql/_typing.py | 42 ++++++---- sqlspec/adapters/aiomysql/config.py | 9 +- sqlspec/adapters/asyncmy/_typing.py | 82 ++++++++++--------- sqlspec/adapters/asyncmy/config.py | 12 +-- sqlspec/adapters/db2/_typing.py | 46 +++++------ sqlspec/adapters/db2/core.py | 11 +-- sqlspec/adapters/db2/pool.py | 22 ++++- sqlspec/adapters/mysqlconnector/_typing.py | 42 ++++++---- sqlspec/adapters/mysqlconnector/config.py | 78 +++++------------- sqlspec/adapters/oracledb/_json_handlers.py | 39 +++------ sqlspec/adapters/oracledb/_typing.py | 27 ++++-- sqlspec/adapters/oracledb/config.py | 55 +------------ sqlspec/adapters/oracledb/core.py | 29 ++++--- sqlspec/adapters/oracledb/migrations.py | 6 +- sqlspec/adapters/pymysql/_typing.py | 18 ++-- sqlspec/adapters/pymysql/pool.py | 8 +- sqlspec/adapters/spanner/data_dictionary.py | 42 +++++----- sqlspec/builder/_base.py | 9 +- sqlspec/builder/_ddl.py | 17 +--- sqlspec/builder/_dml.py | 12 ++- sqlspec/builder/_parsing_utils.py | 39 ++++++++- .../spanner/sql/googlesql/columns.sql | 1 - .../spanner/sql/googlesql/indexes.sql | 1 - .../dialects/spanner/sql/googlesql/tables.sql | 1 - .../spanner/sql/postgresql/columns.sql | 1 - .../spanner/sql/postgresql/indexes.sql | 1 - .../spanner/sql/postgresql/tables.sql | 1 - sqlspec/dialects/db2/_generators.py | 2 +- sqlspec/dialects/db2/_transforms.py | 8 +- sqlspec/dialects/spanner/_generators.py | 25 ++++-- sqlspec/dialects/spanner/_parsers.py | 23 ++++-- sqlspec/dialects/spanner/_spangres.py | 10 ++- sqlspec/dialects/spanner/_spanner.py | 6 +- sqlspec/migrations/base.py | 10 +-- sqlspec/migrations/tracker.py | 4 +- sqlspec/utils/fixtures.py | 56 +++++++------ .../unit/adapters/test_db2/test_async_pool.py | 47 +++++++++++ .../unit/adapters/test_db2/test_exceptions.py | 2 +- tests/unit/adapters/test_db2/test_pool.py | 15 ++++ .../adapters/test_db2/test_transactions.py | 18 ++++ .../adapters/test_mysql_pool_lifecycle.py | 40 +++++++++ .../test_mysqlconnector/test_config.py | 19 +++++ .../adapters/test_oracledb/test_config.py | 38 +++++++++ .../test_core_row_materialization.py | 16 ++++ .../test_oracledb/test_json_handlers.py | 27 ++++++ .../test_oracledb/test_pipeline_helpers.py | 10 +-- .../unit/adapters/test_pymysql/test_config.py | 16 ++++ .../test_data_dictionary_routing.py | 14 ++++ tests/unit/builder/test_ddl_builder.py | 72 +++++++++++++++- tests/unit/dialects/test_db2.py | 15 ++++ tests/unit/dialects/test_spanner_hints.py | 12 +++ tests/unit/dialects/test_spanner_sequences.py | 12 +++ .../migrations/test_tracker_idempotency.py | 14 ++++ tests/unit/migrations/test_version.py | 24 +++--- tests/unit/utils/test_fixtures.py | 56 +++++++++++-- 57 files changed, 893 insertions(+), 429 deletions(-) create mode 100644 tests/unit/adapters/test_mysql_pool_lifecycle.py diff --git a/docs/changelog.rst b/docs/changelog.rst index d8555cc36..27bdaf2ce 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -45,13 +45,32 @@ Unreleased **Fixed:** -* Psycopg reads COPY files in chunks, not all at once. ADK stores use - RETURNING to cut round trips. Psqlpy closes a connection if setup fails. - * Asyncpg stack telemetry reports sequential prepared execution rather than native pipelining. Each statement still returns its own result. -* Builder upserts emit ``MERGE`` for the ``db2`` dialect. +* Fixture files keep JSON strings such as ``"true"`` and ``"[1]"`` as strings. + This also works for SQLite JSON columns. Write objects and arrays directly + instead of encoding them as strings. Column filtering respects case. + See :doc:`usage/testing`. + +* DDL builders and migration trackers keep quoted table names intact. + Names with spaces and mixed-case Oracle names retain their quotes. + +* MySQL pools release connections when setup fails. They discard connections + that fail to roll back. MySQL Connector's async config keeps its direct + connection path. + +* Oracle keeps the user's Thin/Thick mode choice. Pool shutdown waits for + borrowed connections. Custom handlers still convert LOBs, and JSON handlers + preserve the user's callbacks. + +* Db2 pools clean up after failed or cancelled setup. Batch results keep an + unknown row count when the driver cannot report one. + String searches keep their start position and requested occurrence. + +* Spanner schema queries no longer require a table name. SQL output keeps JOIN + hints and plain comments. Sequence statements keep qualified names and + ``IF NOT EXISTS`` guards. * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. @@ -60,6 +79,10 @@ Unreleased parsing SQL again. ADBC keeps bound values in its ADK store queries. DuckDB Arrow loads keep sparse dictionary fields and quote table names. +* Psycopg reads COPY files in chunks, not all at once. ADK stores use + RETURNING to cut round trips. Psqlpy closes a connection if setup fails. + +* Builder upserts emit ``MERGE`` for the ``db2`` dialect. * The arrow-odbc adapter detects the SQL dialect from the ODBC driver name only, so database, host, or user names no longer select the wrong dialect. diff --git a/docs/usage/testing.rst b/docs/usage/testing.rst index 0d2af31d7..f65af6f3d 100644 --- a/docs/usage/testing.rst +++ b/docs/usage/testing.rst @@ -299,8 +299,8 @@ functionality for asynchronous drivers: - **Exact names:** table and column names are quoted in every statement, so they must match the database spelling exactly, including case. Reserved words such as ``order`` work as table or column names. On PostgreSQL, unqualified tables - resolve through the session's search path. The query builder renders Oracle - names unquoted. + resolve through the session's search path. On Oracle, explicitly quoted table + names keep their spelling; ordinary column names follow uppercase folding. - **Column types:** before inserting, the loader reads the table's columns from the driver's data dictionary on PostgreSQL-family, MySQL, DuckDB, and SQLite drivers and converts these JSON values: @@ -317,9 +317,15 @@ functionality for asynchronous drivers: - strings and numbers in ``numeric``/``decimal`` columns to ``Decimal``; - strings in ``uuid`` columns to ``UUID``; - base64 strings in ``bytea``, ``blob``, ``tinyblob``, ``mediumblob``, - ``longblob``, ``binary``, and ``varbinary`` columns to bytes (the only - conversion on SQLite); - - values of PostgreSQL and MySQL ``json``/``jsonb`` columns to JSON text. + ``longblob``, ``binary``, and ``varbinary`` columns to bytes; + - values of PostgreSQL, MySQL, DuckDB, and SQLite ``json``/``jsonb`` columns + to JSON text. SQLite converts only binary and JSON columns. + + Write JSON column values directly in the fixture. Use an object for an object, + an array for an array, and a string for a string. Strings such as ``"true"`` + and ``"[1]"`` remain strings. The loader does not decode their contents again. + Convert pre-encoded object or array strings in older fixtures to JSON objects + or arrays before loading them. Every other value is passed to the driver as decoded from JSON. That includes ``BIT`` columns (including MySQL ``BIT``), @@ -333,8 +339,13 @@ functionality for asynchronous drivers: contain them still load. Other databases are not checked for them, so SQL Server computed columns and Oracle virtual columns are exported and loaded like any other column. -- **Upserts:** ``conflict_keys`` maps a table to the columns of a unique - constraint; every entry must name a table being loaded, spelled exactly. Rows +- **Sparse rows:** rows can have different keys. Missing keys become ``None`` + within the table's combined set of columns. With ``ignore_unknown_columns=True``, + keys absent from the table's column metadata are omitted; names must match + exactly, including case. +- **Upserts:** ``conflict_keys`` maps a table to one column name or a sequence + of column names forming a unique constraint. Entries for tables outside the + current load are ignored, so one mapping can serve several subset loads. Rows for that table update the non-key columns of existing rows instead of failing; PostgreSQL ``GENERATED ALWAYS`` identity columns are never updated, and a table with nothing left to update skips the conflicting row. PostgreSQL-family, @@ -343,7 +354,9 @@ functionality for asynchronous drivers: the table. ``VALUES()`` is deprecated since MySQL 8.0.20 but kept because MariaDB does not support the row-alias form. Other dialects raise ``ValueError`` before any statement runs. Without conflict keys, a duplicate - row raises the database's integrity error. + row raises the database's integrity error. Pass ``exclude_update_columns`` + as a sequence of column names or a per-table mapping to keep those columns + unchanged on conflict; their values are still used for new rows. - **Identity columns:** on PostgreSQL, values for ``GENERATED ALWAYS`` identity columns are inserted with ``OVERRIDING SYSTEM VALUE``. CockroachDB does not accept explicit values for ``GENERATED ALWAYS`` columns; use diff --git a/sqlspec/adapters/aiomysql/_typing.py b/sqlspec/adapters/aiomysql/_typing.py index 8d40d18a5..26fee69b3 100644 --- a/sqlspec/adapters/aiomysql/_typing.py +++ b/sqlspec/adapters/aiomysql/_typing.py @@ -7,27 +7,31 @@ import contextlib from typing import TYPE_CHECKING, Any -import aiomysql -import pymysql.constants -from aiomysql import Pool as AiomysqlPool -from aiomysql import ProgrammingError as AiomysqlProgrammingError +import aiomysql as _aiomysql # pyright: ignore +from aiomysql import Connection # pyright: ignore +from aiomysql import Error as _AiomysqlError # pyright: ignore +from aiomysql import MySQLError as _AiomysqlMySQLError # pyright: ignore +from aiomysql import Pool as _AiomysqlPool # pyright: ignore +from aiomysql import ProgrammingError as AiomysqlProgrammingError # pyright: ignore from aiomysql import SSCursor as AiomysqlSSCursor from aiomysql.cursors import RE_INSERT_VALUES as AIOMYSQL_INSERT_VALUES_PATTERN -from aiomysql.cursors import Cursor as AiomysqlRawCursor -from aiomysql.cursors import DictCursor as AiomysqlDictCursor -from pymysql.err import Error as AiomysqlPymysqlError -from pymysql.err import MySQLError as AiomysqlPymysqlMySQLError +from aiomysql.cursors import Cursor as _AiomysqlCursor # pyright: ignore +from aiomysql.cursors import DictCursor as _AiomysqlDictCursor # pyright: ignore +from pymysql.constants import FIELD_TYPE as _PYMYSQL_FIELD_TYPE # pyright: ignore if TYPE_CHECKING: from collections.abc import Awaitable, Callable from types import TracebackType from typing import Protocol, TypeAlias + from pymysql.err import Error as _PymysqlError + from pymysql.err import MySQLError as _PymysqlMySQLError + from sqlspec.adapters.aiomysql.driver import AiomysqlDriver from sqlspec.core import StatementConfig class AiomysqlConnectionProtocol(Protocol): - async def cursor(self, cursor: "type[AiomysqlRawCursor] | None" = None) -> AiomysqlRawCursor: ... + async def cursor(self, cursor: "type[AiomysqlRawCursor] | None" = None) -> "AiomysqlRawCursor": ... async def commit(self) -> object: ... @@ -38,7 +42,7 @@ def close(self) -> object: ... def get_transaction_status(self) -> bool: ... class AiomysqlModuleProtocol(Protocol): - async def create_pool(self, **kwargs: Any) -> AiomysqlPool: ... + async def create_pool(self, **kwargs: Any) -> "AiomysqlPool": ... async def connect(self, **kwargs: Any) -> "AiomysqlConnection": ... @@ -46,13 +50,23 @@ class AiomysqlFieldTypeProtocol(Protocol): JSON: int AiomysqlConnection: TypeAlias = AiomysqlConnectionProtocol - AiomysqlFieldType: TypeAlias = AiomysqlFieldTypeProtocol AiomysqlModule: TypeAlias = AiomysqlModuleProtocol + AiomysqlRawCursor: TypeAlias = _AiomysqlCursor + AiomysqlDictCursor: TypeAlias = _AiomysqlDictCursor + AiomysqlFieldType: TypeAlias = AiomysqlFieldTypeProtocol + AiomysqlPool: TypeAlias = _AiomysqlPool + AiomysqlPymysqlError: TypeAlias = _PymysqlError + AiomysqlPymysqlMySQLError: TypeAlias = _PymysqlMySQLError if not TYPE_CHECKING: - AiomysqlConnection = aiomysql.Connection - AiomysqlFieldType = pymysql.constants.FIELD_TYPE - AiomysqlModule = aiomysql + AiomysqlConnection = Connection + AiomysqlModule = _aiomysql + AiomysqlRawCursor = _AiomysqlCursor + AiomysqlDictCursor = _AiomysqlDictCursor + AiomysqlFieldType = _PYMYSQL_FIELD_TYPE + AiomysqlPool = _AiomysqlPool + AiomysqlPymysqlError = _AiomysqlError + AiomysqlPymysqlMySQLError = _AiomysqlMySQLError __all__ = ( "AIOMYSQL_INSERT_VALUES_PATTERN", diff --git a/sqlspec/adapters/aiomysql/config.py b/sqlspec/adapters/aiomysql/config.py index 489b00437..d4cee9b5b 100644 --- a/sqlspec/adapters/aiomysql/config.py +++ b/sqlspec/adapters/aiomysql/config.py @@ -1,6 +1,5 @@ """aiomysql database configuration.""" -import asyncio import contextlib from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast from weakref import WeakSet @@ -182,7 +181,6 @@ def build_connection_config( config.setdefault("host", "localhost") config.setdefault("port", 3306) config.setdefault("charset", "utf8mb4") - config.setdefault("pool_recycle", 300) return _normalize_local_infile(config) @@ -204,7 +202,7 @@ async def acquire_connection(self) -> "AiomysqlConnection": try: ensure_conn = self._config._ensure_connection await ensure_conn(connection) - except Exception: + except BaseException: self._contexts.pop(id(connection), None) with contextlib.suppress(Exception): await ctx.__aexit__(None, None, None) @@ -242,7 +240,7 @@ async def __aenter__(self) -> AiomysqlConnection: try: ensure_conn = self._config._ensure_connection await ensure_conn(connection) - except Exception: + except BaseException: self._connection = None self._ctx = None with contextlib.suppress(Exception): @@ -374,8 +372,7 @@ async def _close_pool(self) -> None: """Close the actual async connection pool.""" if self.connection_instance: self.connection_instance.close() - with contextlib.suppress(Exception): - await asyncio.wait_for(self.connection_instance.wait_closed(), timeout=5.0) + await self.connection_instance.wait_closed() self.connection_instance = None async def create_connection(self) -> AiomysqlConnection: diff --git a/sqlspec/adapters/asyncmy/_typing.py b/sqlspec/adapters/asyncmy/_typing.py index 78a836aa9..48912c24d 100644 --- a/sqlspec/adapters/asyncmy/_typing.py +++ b/sqlspec/adapters/asyncmy/_typing.py @@ -6,20 +6,21 @@ import contextlib import os -from typing import TYPE_CHECKING, Any, cast - -import asyncmy -import asyncmy.constants -from asyncmy import Pool as AsyncmyPool -from asyncmy.connection import LoadLocalFile, MySQLResult +from typing import TYPE_CHECKING, Any + +import asyncmy as _asyncmy # pyright: ignore +from asyncmy import Connection # pyright: ignore +from asyncmy import errors as _asyncmy_errors # pyright: ignore +from asyncmy.connection import LoadLocalFile as _LoadLocalFile # pyright: ignore +from asyncmy.connection import MySQLResult as _AsyncmyResult # pyright: ignore +from asyncmy.constants import FIELD_TYPE as _ASYNCMY_FIELD_TYPE # pyright: ignore from asyncmy.cursors import RE_INSERT_VALUES as ASYNCMY_INSERT_VALUES_PATTERN -from asyncmy.cursors import Cursor as AsyncmyRawCursor -from asyncmy.cursors import DictCursor as AsyncmyDictCursor +from asyncmy.cursors import Cursor as _AsyncmyCursor # pyright: ignore +from asyncmy.cursors import DictCursor as _AsyncmyDictCursor # pyright: ignore from asyncmy.cursors import SSCursor as AsyncmySSCursor -from asyncmy.errors import Error as AsyncmyError -from asyncmy.errors import MySQLError as AsyncmyMySQLError from asyncmy.errors import ProgrammingError as AsyncmyProgrammingError -from asyncmy.protocol import LoadLocalPacketWrapper +from asyncmy.pool import Pool as _AsyncmyPool # pyright: ignore +from asyncmy.protocol import LoadLocalPacketWrapper as _LoadLocalPacketWrapper # pyright: ignore from sqlspec.exceptions import SQLSpecError @@ -32,7 +33,7 @@ from sqlspec.core import StatementConfig class AsyncmyConnectionProtocol(Protocol): - def cursor(self) -> AsyncmyRawCursor: ... + def cursor(self) -> "AsyncmyRawCursor": ... async def commit(self) -> object: ... @@ -43,19 +44,29 @@ def close(self) -> None: ... class AsyncmyModuleProtocol(Protocol): async def connect(self, *args: Any, **kwargs: Any) -> "AsyncmyConnection": ... - async def create_pool(self, **kwargs: Any) -> AsyncmyPool: ... + async def create_pool(self, **kwargs: Any) -> "AsyncmyPool": ... class AsyncmyFieldTypeProtocol(Protocol): JSON: int AsyncmyConnection: TypeAlias = AsyncmyConnectionProtocol + AsyncmyDictCursor: TypeAlias = _AsyncmyDictCursor + AsyncmyError: TypeAlias = _asyncmy_errors.Error AsyncmyFieldType: TypeAlias = AsyncmyFieldTypeProtocol + AsyncmyMySQLError: TypeAlias = _asyncmy_errors.MySQLError AsyncmyModule: TypeAlias = AsyncmyModuleProtocol + AsyncmyPool: TypeAlias = _AsyncmyPool + AsyncmyRawCursor: TypeAlias = _AsyncmyCursor if not TYPE_CHECKING: - AsyncmyConnection = asyncmy.Connection - AsyncmyFieldType = asyncmy.constants.FIELD_TYPE - AsyncmyModule = asyncmy + AsyncmyConnection = Connection + AsyncmyDictCursor = _AsyncmyDictCursor + AsyncmyError = _asyncmy_errors.Error + AsyncmyFieldType = _ASYNCMY_FIELD_TYPE + AsyncmyMySQLError = _asyncmy_errors.MySQLError + AsyncmyModule = _asyncmy + AsyncmyPool = _AsyncmyPool + AsyncmyRawCursor = _AsyncmyCursor __all__ = ( "ASYNCMY_INSERT_VALUES_PATTERN", @@ -168,41 +179,38 @@ def asyncmy_local_infile(connection: "AsyncmyConnection", filename: str) -> "Ite previous = raw.__dict__.get("_read_query_result", missing) async def read_result(unbuffered: bool = False) -> None: - setattr(raw, "_result", None) + raw._result = None result = _AsyncmyLocalInfileResult(raw, filename) if unbuffered: try: - init_fn = cast("Callable[[], Awaitable[None]]", result.init_unbuffered_query) - await init_fn() + await result.init_unbuffered_query() # type: ignore[no-untyped-call] except BaseException: - setattr(result, "unbuffered_active", False) - setattr(result, "connection", None) + result.unbuffered_active = False + result.connection = None raise else: - read_fn = cast("Callable[[], Awaitable[None]]", result.read) - await read_fn() - setattr(raw, "_result", result) - setattr(raw, "_affected_rows", result.affected_rows) + await result.read() # type: ignore[no-untyped-call] + raw._result = result + raw._affected_rows = result.affected_rows if result.server_status: - setattr(raw, "server_status", result.server_status) + raw.server_status = result.server_status - setattr(raw, "_read_query_result", read_result) + raw._read_query_result = read_result try: yield except BaseException: with contextlib.suppress(Exception): raw.close() - setattr(raw, "_connected", False) + raw._connected = False raise finally: if previous is missing: - if "_read_query_result" in raw.__dict__: - del raw._read_query_result + del raw._read_query_result else: - setattr(raw, "_read_query_result", previous) + raw._read_query_result = previous -class _AsyncmyLocalInfileResult(MySQLResult): +class _AsyncmyLocalInfileResult(_AsyncmyResult): """Normalize the upstream filename handoff while retaining its native sender.""" __slots__ = ("_filename",) @@ -212,17 +220,13 @@ def __init__(self, connection: Any, filename: str) -> None: self._filename = filename async def _read_load_local_packet(self, first_packet: Any) -> None: - request = LoadLocalPacketWrapper(first_packet).filename + request = _LoadLocalPacketWrapper(first_packet).filename if not self.connection._local_infile or os.fsdecode(request) != self._filename: msg = "MySQL requested an unexpected LOCAL INFILE payload." raise SQLSpecError(msg) - sender = LoadLocalFile(self._filename, self.connection) - send_data = cast("Callable[[], Awaitable[None]]", sender.send_data) - await send_data() + await _LoadLocalFile(self._filename, self.connection).send_data() # type: ignore[no-untyped-call] packet = await self.connection.read_packet() if not packet.is_ok_packet(): msg = "MySQL did not acknowledge the LOCAL INFILE payload." raise SQLSpecError(msg) - read_ok_fn: Callable[[Any], None] | None = getattr(self, "_read_ok_packet", None) - if read_ok_fn is not None: - read_ok_fn(packet) + self._read_ok_packet(packet) # type: ignore[attr-defined] diff --git a/sqlspec/adapters/asyncmy/config.py b/sqlspec/adapters/asyncmy/config.py index 34da35239..e3af7c935 100644 --- a/sqlspec/adapters/asyncmy/config.py +++ b/sqlspec/adapters/asyncmy/config.py @@ -1,6 +1,5 @@ """Asyncmy database configuration.""" -import asyncio import contextlib import inspect from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast @@ -239,7 +238,7 @@ async def acquire_connection(self) -> "AsyncmyConnection": try: ensure_conn = self._config._ensure_connection await ensure_conn(connection) - except Exception: + except BaseException: self._contexts.pop(id(connection), None) with contextlib.suppress(Exception): await ctx.__aexit__(None, None, None) @@ -273,7 +272,7 @@ async def __aenter__(self) -> AsyncmyConnection: try: ensure_conn = self._config._ensure_connection await ensure_conn(connection) - except Exception: + except BaseException: self._connection = None self._ctx = None with contextlib.suppress(Exception): @@ -391,12 +390,7 @@ async def _close_pool(self) -> None: """Close the actual async connection pool.""" if self.connection_instance: self.connection_instance.close() - with contextlib.suppress(Exception): - try: - await asyncio.wait_for(self.connection_instance.wait_closed(), timeout=5.0) - except (TimeoutError, asyncio.TimeoutError): - if hasattr(self.connection_instance, "terminate"): - self.connection_instance.terminate() + await self.connection_instance.wait_closed() self.connection_instance = None async def create_connection(self) -> AsyncmyConnection: diff --git a/sqlspec/adapters/db2/_typing.py b/sqlspec/adapters/db2/_typing.py index c2fbdca6d..b7edaf0c2 100644 --- a/sqlspec/adapters/db2/_typing.py +++ b/sqlspec/adapters/db2/_typing.py @@ -162,18 +162,18 @@ def __enter__(self) -> Any: from sqlspec.adapters.db2.driver import Db2SyncDriver self._connection = self._acquire_connection() - self._driver = Db2SyncDriver( - connection=self._connection, statement_config=self._statement_config, driver_features=self._driver_features - ) - if self._begin_transaction: - try: + try: + self._driver = Db2SyncDriver( + connection=self._connection, + statement_config=self._statement_config, + driver_features=self._driver_features, + ) + if self._begin_transaction: self._driver.begin() - except BaseException as exc: - self._release_connection(self._connection, exc_type=type(exc), exc_val=exc, exc_tb=exc.__traceback__) - self._connection = None - self._driver = None - raise - return self._prepare_driver(self._driver) + return self._prepare_driver(self._driver) + except BaseException as exc: + self.__exit__(type(exc), exc, exc.__traceback__) + raise def __exit__( self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" @@ -282,20 +282,18 @@ async def __aenter__(self) -> Any: from sqlspec.adapters.db2.driver import Db2AsyncDriver self._connection = await self._acquire_connection() - self._driver = Db2AsyncDriver( - connection=self._connection, statement_config=self._statement_config, driver_features=self._driver_features - ) - if self._begin_transaction: - try: + try: + self._driver = Db2AsyncDriver( + connection=self._connection, + statement_config=self._statement_config, + driver_features=self._driver_features, + ) + if self._begin_transaction: await self._driver.begin() - except BaseException as exc: - await self._release_connection( - self._connection, exc_type=type(exc), exc_val=exc, exc_tb=exc.__traceback__ - ) - self._connection = None - self._driver = None - raise - return self._prepare_driver(self._driver) + return self._prepare_driver(self._driver) + except BaseException as exc: + await self.__aexit__(type(exc), exc, exc.__traceback__) + raise async def __aexit__( self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" diff --git a/sqlspec/adapters/db2/core.py b/sqlspec/adapters/db2/core.py index f4286c904..aa34e7a5f 100644 --- a/sqlspec/adapters/db2/core.py +++ b/sqlspec/adapters/db2/core.py @@ -402,21 +402,22 @@ def resolve_rowcount(cursor: Any) -> int: def resolve_many_rowcount(cursor: Any, parameters: Any) -> int: - """Resolve the affected rowcount of a batch, falling back to the number of parameter sets. + """Resolve a batch rowcount without assuming one affected row per parameter set. Args: cursor: Cursor that executed the batch. parameters: Parameter sets passed to ``executemany``. Returns: - The driver-reported rowcount, else the number of parameter sets, else 0. + The driver-reported rowcount, or -1 when the count is unknown. + An empty batch has no affected rows. """ count = resolve_rowcount(cursor) if count >= 0: return count - if parameters is not None and hasattr(parameters, "__len__"): - return len(parameters) - return 0 + if parameters is None or (isinstance(parameters, (list, tuple)) and not parameters): + return 0 + return -1 def collect_rows( diff --git a/sqlspec/adapters/db2/pool.py b/sqlspec/adapters/db2/pool.py index c615599d6..37a4e9178 100644 --- a/sqlspec/adapters/db2/pool.py +++ b/sqlspec/adapters/db2/pool.py @@ -122,7 +122,12 @@ def new_connection(self) -> Any: connection = ibm_db_dbi.connect(self._dsn, "", "", "", "", {ibm_db_dbi.SQL_ATTR_AUTOCOMMIT: autocommit_mode}) if self._on_connection_create is not None: - self._on_connection_create(connection) + try: + self._on_connection_create(connection) + except BaseException: + with contextlib.suppress(Exception): + connection.close() + raise return connection @@ -433,6 +438,13 @@ async def acquire(self) -> Any: except BaseException: semaphore.release() raise + if self._closed: + try: + await self._close_connection(record.connection) + finally: + semaphore.release() + msg = "Db2 async connection pool is closed" + raise DatabaseConnectionError(msg) self._checked_out[id(record.connection)] = record return record.connection @@ -496,8 +508,12 @@ async def _checkout(self) -> _Db2PooledConnection: """Pop a reusable idle connection, or open a new one when none is left.""" while self._idle: record = self._idle.pop() - if await self._is_reusable(record): - return record + try: + if await self._is_reusable(record): + return record + except BaseException: + await self._close_connection(record.connection) + raise await self._close_connection(record.connection) connection = await self.new_connection() now = time.monotonic() diff --git a/sqlspec/adapters/mysqlconnector/_typing.py b/sqlspec/adapters/mysqlconnector/_typing.py index 066648399..3d0dc3568 100644 --- a/sqlspec/adapters/mysqlconnector/_typing.py +++ b/sqlspec/adapters/mysqlconnector/_typing.py @@ -7,16 +7,15 @@ import contextlib from typing import TYPE_CHECKING, Any -import mysql -import mysql.connector.aio -import mysql.connector.constants -from mysql.connector import Error as MysqlConnectorError -from mysql.connector import MySQLConnection as MysqlConnectorSyncConnection -from mysql.connector.abstracts import MySQLConnectionAbstract, MySQLCursorAbstract -from mysql.connector.aio.cursor import MySQLCursor as MysqlConnectorAsyncRawCursor -from mysql.connector.aio.pooling import MySQLConnectionPool as MysqlConnectorAsyncPool -from mysql.connector.cursor import MySQLCursor as MysqlConnectorSyncRawCursor -from mysql.connector.pooling import MySQLConnectionPool as MysqlConnectorConnectionPool +import mysql as _mysql +import mysql.connector as _mysql_connector +from mysql.connector import MySQLConnection as _MysqlConnectorSyncConnection +from mysql.connector import aio as _mysql_connector_aio +from mysql.connector import pooling as _mysql_connector_pooling +from mysql.connector.aio import MySQLConnection as _MysqlConnectorAsyncConnection # pyright: ignore[reportMissingImports] +from mysql.connector.aio.cursor import MySQLCursor as _MysqlConnectorAsyncRawCursor # pyright: ignore[reportMissingImports] +from mysql.connector.constants import FieldType as _MysqlConnectorFieldType +from mysql.connector.cursor import MySQLCursor as _MysqlConnectorSyncRawCursor if TYPE_CHECKING: from collections.abc import Awaitable, Callable @@ -44,29 +43,36 @@ class MysqlConnectorFieldTypeProtocol(Protocol): JSON: int class MysqlConnectorConnectorModuleProtocol(Protocol): - def connect(self, *args: Any, **kwargs: Any) -> MysqlConnectorSyncConnection: ... + def connect(self, *args: Any, **kwargs: Any) -> "MysqlConnectorSyncConnection": ... class MysqlConnectorMysqlModuleProtocol(Protocol): connector: "ClassVar[MysqlConnectorConnectorModuleProtocol]" + MysqlConnectorSyncConnection: TypeAlias = _MysqlConnectorSyncConnection MysqlConnectorAio: TypeAlias = MysqlConnectorAioModuleProtocol MysqlConnectorAsyncConnection: TypeAlias = MysqlConnectorAsyncConnectionProtocol + MysqlConnectorSyncRawCursor: TypeAlias = _MysqlConnectorSyncRawCursor + MysqlConnectorConnectionPool: TypeAlias = _mysql_connector_pooling.MySQLConnectionPool + MysqlConnectorError: TypeAlias = _mysql_connector.Error MysqlConnectorFieldType: TypeAlias = MysqlConnectorFieldTypeProtocol MysqlConnectorMysqlModule: TypeAlias = MysqlConnectorMysqlModuleProtocol + MysqlConnectorAsyncRawCursor: TypeAlias = _MysqlConnectorAsyncRawCursor if not TYPE_CHECKING: - MysqlConnectorAio = mysql.connector.aio - MysqlConnectorAsyncConnection = mysql.connector.aio.MySQLConnection - MysqlConnectorFieldType = mysql.connector.constants.FieldType - MysqlConnectorMysqlModule = mysql + MysqlConnectorAio = _mysql_connector_aio + MysqlConnectorSyncConnection = _MysqlConnectorSyncConnection + MysqlConnectorAsyncConnection = _MysqlConnectorAsyncConnection + MysqlConnectorConnectionPool = _mysql_connector_pooling.MySQLConnectionPool + MysqlConnectorError = _mysql_connector.Error + MysqlConnectorFieldType = _MysqlConnectorFieldType + MysqlConnectorMysqlModule = _mysql + MysqlConnectorSyncRawCursor = _MysqlConnectorSyncRawCursor + MysqlConnectorAsyncRawCursor = _MysqlConnectorAsyncRawCursor __all__ = ( - "MySQLConnectionAbstract", - "MySQLCursorAbstract", "MysqlConnectorAio", "MysqlConnectorAsyncConnection", "MysqlConnectorAsyncCursor", - "MysqlConnectorAsyncPool", "MysqlConnectorAsyncRawCursor", "MysqlConnectorAsyncSessionContext", "MysqlConnectorConnectionPool", diff --git a/sqlspec/adapters/mysqlconnector/config.py b/sqlspec/adapters/mysqlconnector/config.py index 40f594df0..8a64d48bf 100644 --- a/sqlspec/adapters/mysqlconnector/config.py +++ b/sqlspec/adapters/mysqlconnector/config.py @@ -11,7 +11,6 @@ MysqlConnectorAio, MysqlConnectorAsyncConnection, MysqlConnectorAsyncCursor, - MysqlConnectorAsyncPool, MysqlConnectorAsyncSessionContext, MysqlConnectorConnectionPool, MysqlConnectorMysqlModule, @@ -26,7 +25,7 @@ MysqlConnectorSyncDriver, MysqlConnectorSyncExceptionHandler, ) -from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs, SyncDatabaseConfig +from sqlspec.config import ExtensionConfigs, NoPoolAsyncConfig, SyncDatabaseConfig from sqlspec.core import TypeCoercionCapabilities from sqlspec.driver import ( AsyncPoolConnectionContext, @@ -173,10 +172,6 @@ class MysqlConnectorPoolParams(MysqlConnectorSyncConnectionParams): class MysqlConnectorAsyncConnectionParams(_MysqlConnectorBaseConnectionParams): """MysqlConnector async connection parameters.""" - pool_name: NotRequired[str] - pool_size: NotRequired[int] - pool_reset_session: NotRequired[bool] - class MysqlConnectorDriverFeatures(TypedDict): """MysqlConnector driver feature flags. @@ -280,7 +275,7 @@ class MysqlConnectorAsyncConnectionContext(AsyncPoolConnectionContext): __slots__ = () async def __aenter__(self) -> MysqlConnectorAsyncConnection: - self._connection = await self._config._acquire_async_connection() + self._connection = await self._config.create_connection() return cast("MysqlConnectorAsyncConnection", self._connection) async def __aexit__( @@ -296,7 +291,7 @@ class _MysqlConnectorAsyncSessionConnectionHandler(AsyncPoolSessionFactory): __slots__ = () async def acquire_connection(self) -> MysqlConnectorAsyncConnection: - self._connection = await self._config._acquire_async_connection() + self._connection = await self._config.create_connection() return cast("MysqlConnectorAsyncConnection", self._connection) async def release_connection(self, _conn: MysqlConnectorAsyncConnection, **kwargs: Any) -> None: @@ -438,9 +433,7 @@ def get_event_runtime_hints(self) -> "EventRuntimeHints": return EventRuntimeHints(poll_interval=0.25, lease_seconds=5, select_for_update=True, skip_locked=True) -class MysqlConnectorAsyncConfig( - AsyncDatabaseConfig[MysqlConnectorAsyncConnection, "MysqlConnectorAsyncPool", MysqlConnectorAsyncDriver] -): +class MysqlConnectorAsyncConfig(NoPoolAsyncConfig[MysqlConnectorAsyncConnection, MysqlConnectorAsyncDriver]): """Configuration for mysql-connector async MySQL connections.""" driver_type: ClassVar[type[MysqlConnectorAsyncDriver]] = MysqlConnectorAsyncDriver @@ -467,7 +460,7 @@ def __init__( self, *, connection_config: "MysqlConnectorAsyncConnectionParams | dict[str, Any] | None" = None, - connection_instance: "MysqlConnectorAsyncPool | None" = None, + connection_instance: Any = None, migration_config: "dict[str, Any] | None" = None, statement_config: "StatementConfig | None" = None, driver_features: "MysqlConnectorDriverFeatures | dict[str, Any] | None" = None, @@ -476,7 +469,7 @@ def __init__( observability_config: "ObservabilityConfig | None" = None, **kwargs: Any, ) -> None: - connection_config = build_connection_config(connection_config) + self.connection_config = build_connection_config(connection_config) statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, driver_features) @@ -484,15 +477,14 @@ def __init__( self._user_connection_hook: Callable[[MysqlConnectorAsyncConnection], Awaitable[None]] | None = ( features_dict.pop("on_connection_create", None) ) - self._initialized_connections: WeakSet[Any] = WeakSet() - features_dict.setdefault("enable_local_infile_bulk_load", connection_config["allow_local_infile"]) - if features_dict.get("enable_local_infile_bulk_load") and not connection_config.get("allow_local_infile"): + features_dict.setdefault("enable_local_infile_bulk_load", self.connection_config["allow_local_infile"]) + if features_dict.get("enable_local_infile_bulk_load") and not self.connection_config.get("allow_local_infile"): msg = "enable_local_infile_bulk_load requires local_infile=True or allow_local_infile=True in connection_config." raise ImproperConfigurationError(msg) super().__init__( - connection_config=connection_config, + connection_config=self.connection_config, connection_instance=connection_instance, migration_config=migration_config, statement_config=statement_config, @@ -503,47 +495,19 @@ def __init__( **kwargs, ) - async def _create_pool(self) -> "MysqlConnectorAsyncPool": - config = dict(self.connection_config) - pool_name = config.pop("pool_name", None) - pool_size = config.pop("pool_size", None) - pool_reset = config.pop("pool_reset_session", True) - pool = MysqlConnectorAsyncPool( - pool_name=pool_name, pool_size=pool_size or 5, pool_reset_session=pool_reset, **config - ) - await pool.initialize_pool() - return pool - - async def _close_pool(self) -> None: - if self.connection_instance is not None: - with contextlib.suppress(Exception): - await self.connection_instance.close_pool() - self.connection_instance = None - - async def _ensure_connection(self, connection: "MysqlConnectorAsyncConnection") -> None: - """Ensure connection callback has been called exactly once for this connection.""" - if self._user_connection_hook is None: - return - underlying = getattr(connection, "_cnx", None) or connection - if underlying not in self._initialized_connections: - await self._user_connection_hook(connection) - self._initialized_connections.add(underlying) - - async def _acquire_async_connection(self) -> MysqlConnectorAsyncConnection: - """Acquire and initialize an async mysql-connector connection from pool.""" - pool = await self.provide_pool() - connection = cast("MysqlConnectorAsyncConnection", await pool.get_connection()) - await self._ensure_connection(connection) - return connection - async def create_connection(self) -> MysqlConnectorAsyncConnection: - """Open a standalone connection owned by the caller.""" - config = {key: value for key, value in self.connection_config.items() if key not in _POOL_ONLY_CONFIG_KEYS} - connection = await mysqlconnector_aio.connect(**config) - autocommit = config.get("autocommit") - if autocommit is not None: - await connection.set_autocommit(bool(autocommit)) - await self._ensure_connection(connection) + """Open and initialize a standalone connection owned by the caller.""" + connection = await mysqlconnector_aio.connect(**self.connection_config) + try: + autocommit = self.connection_config.get("autocommit") + if autocommit is not None: + await connection.set_autocommit(bool(autocommit)) + if self._user_connection_hook is not None: + await self._user_connection_hook(connection) + except BaseException: + with contextlib.suppress(Exception): + await connection.close() + raise return connection def provide_connection(self, *args: Any, **kwargs: Any) -> "MysqlConnectorAsyncConnectionContext": diff --git a/sqlspec/adapters/oracledb/_json_handlers.py b/sqlspec/adapters/oracledb/_json_handlers.py index 61ead2cc4..cbd9360b1 100644 --- a/sqlspec/adapters/oracledb/_json_handlers.py +++ b/sqlspec/adapters/oracledb/_json_handlers.py @@ -119,35 +119,16 @@ def json_output_type_handler(cursor: "Cursor | AsyncCursor", metadata: Any) -> A return _output_type_handler(cursor, metadata) -def _has_input_handler(handler: Any, target_inner: Any) -> bool: +def _has_handler(handler: Any, target_inner: Any, chain_function: Any) -> bool: current = handler while current is not None: if current is target_inner: return True - if isinstance(current, partial): - args = current.args - if args and args[0] is target_inner: - return True - if len(args) > 1: - current = args[1] - continue - break - return False - - -def _has_output_handler(handler: Any, target_inner: Any) -> bool: - current = handler - while current is not None: - if current is target_inner: + if not isinstance(current, partial) or current.func is not chain_function: + return False + inner, current = current.args + if inner is target_inner: return True - if isinstance(current, partial): - args = current.args - if args and args[0] is target_inner: - return True - if len(args) > 1: - current = args[1] - continue - break return False @@ -171,15 +152,19 @@ def register_json_handlers(connection: "Connection | AsyncConnection") -> None: def chain_input_handler(inner: Any, fallback: "Any | None") -> Any: - """Build an input type handler that chains ``inner`` to ``fallback``.""" - if fallback is not None and _has_input_handler(fallback, inner): + """Build an input type handler that chains ``inner`` to ``fallback``. + + A partial keeps the signature introspectable when compiled with mypyc; + python-oracledb uses it to select the handler calling convention. + """ + if fallback is not None and _has_handler(fallback, inner, _chained_input_handler): return fallback return partial(_chained_input_handler, inner, fallback) def chain_output_handler(inner: Any, fallback: "Any | None") -> Any: """Build an output type handler that chains ``inner`` to ``fallback``.""" - if fallback is not None and _has_output_handler(fallback, inner): + if fallback is not None and _has_handler(fallback, inner, _chained_output_handler): return fallback return partial(_chained_output_handler, inner, fallback) diff --git a/sqlspec/adapters/oracledb/_typing.py b/sqlspec/adapters/oracledb/_typing.py index 9f414dcfa..9ac2e396e 100644 --- a/sqlspec/adapters/oracledb/_typing.py +++ b/sqlspec/adapters/oracledb/_typing.py @@ -22,28 +22,43 @@ DB_TYPE_VECTOR, DEQ_IMMEDIATE, DEQ_ON_COMMIT, + AsyncConnection, + AsyncCursor, + Connection, + Cursor, DatabaseError, Error, ) -from oracledb import AsyncConnection as OracleAsyncConnection -from oracledb import AsyncCursor as OracleAsyncRawCursor from oracledb import AuthMode as OracleAuthMode -from oracledb import Connection as OracleSyncConnection -from oracledb import Cursor as OracleSyncRawCursor from oracledb import PoolGetMode as OraclePoolGetMode from oracledb import Purity as OraclePurity from oracledb import SparseVector as OracleSparseVector from oracledb import create_pipeline as oracledb_create_pipeline -from oracledb.pool import AsyncConnectionPool as OracleAsyncConnectionPool -from oracledb.pool import ConnectionPool as OracleSyncConnectionPool +from oracledb.pool import AsyncConnectionPool, ConnectionPool if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType + from typing import TypeAlias from sqlspec.adapters.oracledb.driver import OracleAsyncDriver, OracleSyncDriver from sqlspec.core import StatementConfig + OracleSyncConnection: TypeAlias = Connection + OracleAsyncConnection: TypeAlias = AsyncConnection + OracleSyncConnectionPool: TypeAlias = ConnectionPool + OracleAsyncConnectionPool: TypeAlias = AsyncConnectionPool + OracleSyncRawCursor: TypeAlias = Cursor + OracleAsyncRawCursor: TypeAlias = AsyncCursor + +if not TYPE_CHECKING: + OracleSyncConnection = Connection + OracleAsyncConnection = AsyncConnection + OracleSyncConnectionPool = ConnectionPool + OracleAsyncConnectionPool = AsyncConnectionPool + OracleSyncRawCursor = Cursor + OracleAsyncRawCursor = AsyncCursor + __all__ = ( "DB_TYPE_BLOB", "DB_TYPE_CHAR", diff --git a/sqlspec/adapters/oracledb/config.py b/sqlspec/adapters/oracledb/config.py index f391d9ea8..00d6c0820 100644 --- a/sqlspec/adapters/oracledb/config.py +++ b/sqlspec/adapters/oracledb/config.py @@ -25,12 +25,7 @@ from sqlspec.adapters.oracledb._typing import oracledb_module as oracledb from sqlspec.adapters.oracledb._uuid_handlers import register_uuid_handlers from sqlspec.adapters.oracledb._vector_handlers import register_numpy_handlers -from sqlspec.adapters.oracledb.core import ( - apply_driver_features, - build_connection_config, - client_is_thin_mode, - default_statement_config, -) +from sqlspec.adapters.oracledb.core import apply_driver_features, build_connection_config, default_statement_config from sqlspec.adapters.oracledb.data_dictionary import OracleVersionCache, resolve_oracle_connection_major from sqlspec.adapters.oracledb.driver import ( OracleAsyncDriver, @@ -49,14 +44,12 @@ ) from sqlspec.extensions.events import EventRuntimeHints from sqlspec.utils.config_tools import normalize_connection_config -from sqlspec.utils.logging import get_logger if TYPE_CHECKING: from types import TracebackType from sqlspec.core import StatementConfig -logger = get_logger("sqlspec.adapters.oracledb.config") __all__ = ( "OracleAsyncConfig", @@ -286,20 +279,6 @@ def release_connection(self, _conn: "OracleSyncConnection", **kwargs: Any) -> No self._conn = None -def _ensure_thick_mode(*, lib_dir: "str | None" = None, config_dir: "str | None" = None) -> None: - """Initialize python-oracledb thick mode automatically if requested.""" - if not client_is_thin_mode(): - return - init_oracle_client = getattr(oracledb, "init_oracle_client", None) - if callable(init_oracle_client): - kwargs: dict[str, Any] = {} - if lib_dir is not None: - kwargs["lib_dir"] = lib_dir - if config_dir is not None: - kwargs["config_dir"] = config_dir - init_oracle_client(**kwargs) - - class OracleSyncConfig(SyncDatabaseConfig[OracleSyncConnection, "OracleSyncConnectionPool", OracleSyncDriver]): """Configuration for Oracle synchronous database connections.""" @@ -429,17 +408,6 @@ def _create_pool(self) -> "OracleSyncConnectionPool": """Create the actual connection pool.""" config = dict(self.connection_config) - thick_mode = config.pop("thick_mode", False) - lib_dir = config.pop("lib_dir", None) - config_dir = config.get("config_dir") - if thick_mode or lib_dir is not None or config.get("soda_metadata_cache"): - _ensure_thick_mode(lib_dir=lib_dir, config_dir=config_dir) - - if config.get("soda_metadata_cache") and client_is_thin_mode(): - logger.warning( - "soda_metadata_cache requires python-oracledb Thick mode; SODA operations are unsupported in Thin mode." - ) - config.pop("threaded", None) config["session_callback"] = self._init_connection @@ -488,10 +456,7 @@ def _prepare_driver(self, driver: "OracleSyncDriver") -> "OracleSyncDriver": def _close_pool(self) -> None: """Close the actual connection pool.""" if self.connection_instance: - try: - self.connection_instance.close(force=True) - except TypeError: - self.connection_instance.close() + self.connection_instance.close() self.connection_instance = None self._oracle_version_cache.reset() @@ -635,17 +600,6 @@ async def _create_pool(self) -> "OracleAsyncConnectionPool": """Create the actual async connection pool.""" config = dict(self.connection_config) - thick_mode = config.pop("thick_mode", False) - lib_dir = config.pop("lib_dir", None) - config_dir = config.get("config_dir") - if thick_mode or lib_dir is not None or config.get("soda_metadata_cache"): - _ensure_thick_mode(lib_dir=lib_dir, config_dir=config_dir) - - if config.get("soda_metadata_cache") and client_is_thin_mode(): - logger.warning( - "soda_metadata_cache requires python-oracledb Thick mode; SODA operations are unsupported in Thin mode." - ) - config.pop("threaded", None) config["session_callback"] = self._init_connection @@ -695,9 +649,6 @@ def _prepare_driver(self, driver: "OracleAsyncDriver") -> "OracleAsyncDriver": async def _close_pool(self) -> None: """Close the actual async connection pool.""" if self.connection_instance: - try: - await self.connection_instance.close(force=True) - except TypeError: - await self.connection_instance.close() + await self.connection_instance.close() self.connection_instance = None self._oracle_version_cache.reset() diff --git a/sqlspec/adapters/oracledb/core.py b/sqlspec/adapters/oracledb/core.py index 954ceb131..44498f617 100644 --- a/sqlspec/adapters/oracledb/core.py +++ b/sqlspec/adapters/oracledb/core.py @@ -74,7 +74,6 @@ "build_profile", "build_statement_config", "build_truncate_statement", - "client_is_thin_mode", "coerce_large_parameters_async", "coerce_large_parameters_sync", "coerce_many_parameters_async", @@ -86,6 +85,8 @@ "default_statement_config", "driver_profile", "normalize_column_names", + "normalize_execute_many_parameters_async", + "normalize_execute_many_parameters_sync", "resolve_row_metadata", "resolve_rowcount", "supports_df_batches", @@ -156,14 +157,6 @@ def connection_is_thin(connection: object) -> bool: return bool(thin) -def client_is_thin_mode() -> bool: - """Return whether python-oracledb is currently in Thin mode.""" - is_thin = getattr(oracledb_module, "is_thin_mode", None) - if callable(is_thin): - return bool(is_thin()) - return True - - def supports_direct_path_load(connection: object) -> bool: """Return whether a connection supports direct path load. @@ -232,7 +225,7 @@ def normalize_column_names(column_names: "list[str]", driver_features: "dict[str return normalized -def _normalize_execute_many_parameters_sync(parameters: Any) -> Any: +def normalize_execute_many_parameters_sync(parameters: Any) -> Any: """Normalize parameters for Oracle executemany calls. Args: @@ -246,7 +239,7 @@ def _normalize_execute_many_parameters_sync(parameters: Any) -> Any: return parameters -def _normalize_execute_many_parameters_async(parameters: Any) -> Any: +def normalize_execute_many_parameters_async(parameters: Any) -> Any: """Normalize parameters for Oracle async executemany calls. Args: @@ -370,7 +363,7 @@ def coerce_many_parameters_sync( version_cache: Any = None, ) -> Any: """Coerce every parameter row prepared for synchronous ``executemany``.""" - normalized = _normalize_execute_many_parameters_sync(parameters) + normalized = normalize_execute_many_parameters_sync(parameters) if not normalized: return normalized json_binding_state = _OracleJsonBindingState(connection, version_cache) @@ -404,7 +397,7 @@ async def coerce_many_parameters_async( version_cache: Any = None, ) -> Any: """Coerce every parameter row prepared for asynchronous ``executemany``.""" - normalized = _normalize_execute_many_parameters_async(parameters) + normalized = normalize_execute_many_parameters_async(parameters) if not normalized: return normalized json_binding_state = _OracleJsonBindingState(connection, version_cache) @@ -870,7 +863,10 @@ def collect_sync_rows( if requires_lob_coercion is None: requires_lob_coercion = _description_requires_lob_coercion(description) if not requires_lob_coercion: - return cast("list[tuple[Any, ...]]", fetched_data), resolved_column_names + first_row = fetched_data[0] + first_row_tuple = first_row if isinstance(first_row, tuple) else tuple(first_row) + if not _row_requires_lob_coercion(first_row_tuple): + return cast("list[tuple[Any, ...]]", fetched_data), resolved_column_names data: list[tuple[Any, ...]] = [] for row in fetched_data: @@ -918,7 +914,10 @@ async def collect_async_rows( if requires_lob_coercion is None: requires_lob_coercion = _description_requires_lob_coercion(description) if not requires_lob_coercion: - return cast("list[tuple[Any, ...]]", fetched_data), resolved_column_names + first_row = fetched_data[0] + first_row_tuple = first_row if isinstance(first_row, tuple) else tuple(first_row) + if not _row_requires_lob_coercion(first_row_tuple): + return cast("list[tuple[Any, ...]]", fetched_data), resolved_column_names data: list[tuple[Any, ...]] = [] for row in fetched_data: diff --git a/sqlspec/adapters/oracledb/migrations.py b/sqlspec/adapters/oracledb/migrations.py index 52ce59aee..5437b0633 100644 --- a/sqlspec/adapters/oracledb/migrations.py +++ b/sqlspec/adapters/oracledb/migrations.py @@ -62,11 +62,7 @@ def _qualify_version_table(self, version_table_name: str, version_table_schema: def _tracking_table_builder(self) -> CreateTable: """Return an Oracle CREATE TABLE builder for the tracker table.""" - table_name = self._normalize_oracle_identifier(self.version_table_name) - builder = sql.create_table(table_name) - if self.version_table_schema: - builder.in_schema(self._normalize_oracle_identifier(self.version_table_schema)) - return builder + return sql.create_table(self.version_table) def _tracking_table_ddl(self) -> CreateTable: """Get Oracle-specific SQL builder for creating the tracking table. diff --git a/sqlspec/adapters/pymysql/_typing.py b/sqlspec/adapters/pymysql/_typing.py index 088b7a1d4..e3debdc90 100644 --- a/sqlspec/adapters/pymysql/_typing.py +++ b/sqlspec/adapters/pymysql/_typing.py @@ -8,11 +8,9 @@ from typing import TYPE_CHECKING, Any import pymysql -import pymysql.constants -from pymysql import MySQLError as PyMysqlMySQLError -from pymysql.connections import Connection as PyMysqlConnection +from pymysql.constants import FIELD_TYPE as _PYMYSQL_FIELD_TYPE +from pymysql.constants import SERVER_STATUS as _PYMYSQL_SERVER_STATUS from pymysql.cursors import RE_INSERT_VALUES as PYMYSQL_INSERT_VALUES_PATTERN -from pymysql.cursors import Cursor as PyMysqlRawCursor from pymysql.cursors import DictCursor as PyMysqlDictCursor from pymysql.cursors import SSCursor as PyMysqlSSCursor @@ -34,14 +32,20 @@ class PyMysqlFieldTypeProtocol(Protocol): class PyMysqlServerStatusProtocol(Protocol): SERVER_STATUS_IN_TRANS: int - PyMysqlConnect: TypeAlias = type[PyMysqlConnection] + PyMysqlConnect: TypeAlias = type["PyMysqlConnection"] + PyMysqlConnection: TypeAlias = pymysql.connections.Connection PyMysqlFieldType: TypeAlias = PyMysqlFieldTypeProtocol + PyMysqlMySQLError: TypeAlias = pymysql.MySQLError + PyMysqlRawCursor: TypeAlias = pymysql.cursors.Cursor PyMysqlServerStatus: TypeAlias = PyMysqlServerStatusProtocol if not TYPE_CHECKING: PyMysqlConnect = pymysql.connect - PyMysqlFieldType = pymysql.constants.FIELD_TYPE - PyMysqlServerStatus = pymysql.constants.SERVER_STATUS + PyMysqlConnection = pymysql.connections.Connection + PyMysqlFieldType = _PYMYSQL_FIELD_TYPE + PyMysqlMySQLError = pymysql.MySQLError + PyMysqlRawCursor = pymysql.cursors.Cursor + PyMysqlServerStatus = _PYMYSQL_SERVER_STATUS __all__ = ( diff --git a/sqlspec/adapters/pymysql/pool.py b/sqlspec/adapters/pymysql/pool.py index be1ffef80..08906a6a4 100644 --- a/sqlspec/adapters/pymysql/pool.py +++ b/sqlspec/adapters/pymysql/pool.py @@ -207,8 +207,14 @@ def acquire(self) -> PyMysqlConnection: def release(self, connection: PyMysqlConnection) -> None: """Release connection back to the pool, sanitizing transactions.""" if bool(getattr(connection, "server_status", 0) & PyMysqlServerStatus.SERVER_STATUS_IN_TRANS): - with contextlib.suppress(Exception): + try: connection.rollback() + except Exception: + if getattr(self._thread_local, "connection", None) is connection: + self._close_thread_connection() + else: + self._retire_connection(connection) + raise def size(self) -> int: """Report total active connections managed by this pool.""" diff --git a/sqlspec/adapters/spanner/data_dictionary.py b/sqlspec/adapters/spanner/data_dictionary.py index c20f17463..f6b0986c7 100644 --- a/sqlspec/adapters/spanner/data_dictionary.py +++ b/sqlspec/adapters/spanner/data_dictionary.py @@ -1,6 +1,5 @@ """Spanner metadata queries using INFORMATION_SCHEMA.""" -from collections.abc import Sequence from typing import TYPE_CHECKING, Any, ClassVar, cast from mypy_extensions import mypyc_attr @@ -30,6 +29,8 @@ from sqlspec.driver import SyncDataDictionaryBase if TYPE_CHECKING: + from collections.abc import Sequence + from sqlspec.adapters.spanner.driver import SpannerSyncDriver __all__ = ("SpannerDataDictionary",) @@ -80,15 +81,12 @@ def __init__(self, mode: str = "googlesql") -> None: super().__init__() self.mode = _normalize_spanner_metadata_mode(mode) - def get_query(self, domain: str, operation: str, *, mode: str | None = None) -> SQL: + def get_query(self, domain: str, operation: str, *, mode: "str | None" = None) -> SQL: """Return an exact domain query for this dialect.""" resolved_mode = self.mode if mode is None else _normalize_spanner_metadata_mode(mode) - query = super().get_query(domain, operation, mode=resolved_mode) - if not isinstance(query, SQL): - query = cast("SQL", query) - return query + return super().get_query(domain, operation, mode=resolved_mode) - def get_version(self, driver: "SpannerSyncDriver | None" = None) -> VersionInfo | None: + def get_version(self, driver: "SpannerSyncDriver | None" = None) -> "VersionInfo | None": """Get Spanner version information. Args: @@ -126,7 +124,7 @@ def get_optimal_type(self, driver: "SpannerSyncDriver | None" = None, type_categ _ = driver return self.get_dialect_config().get_optimal_type(type_category) - def get_tables(self, driver: "SpannerSyncDriver", schema: str | None = None) -> list[TableMetadata]: + def get_tables(self, driver: "SpannerSyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": """Get tables using INFORMATION_SCHEMA.""" schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -135,8 +133,8 @@ def get_tables(self, driver: "SpannerSyncDriver", schema: str | None = None) -> ) def get_columns( - self, driver: "SpannerSyncDriver", table: str | None = None, schema: str | None = None - ) -> list[ColumnMetadata]: + self, driver: "SpannerSyncDriver", table: "str | None" = None, schema: "str | None" = None + ) -> "list[ColumnMetadata]": """Get column information for a table or schema.""" schema_name = self.resolve_schema(schema) if table is None: @@ -156,8 +154,8 @@ def get_columns( ) def get_indexes( - self, driver: "SpannerSyncDriver", table: str | None = None, schema: str | None = None - ) -> list[IndexMetadata]: + self, driver: "SpannerSyncDriver", table: "str | None" = None, schema: "str | None" = None + ) -> "list[IndexMetadata]": """Get index metadata for a table or schema.""" schema_name = self.resolve_schema(schema) if table is None: @@ -177,8 +175,8 @@ def get_indexes( ) def get_foreign_keys( - self, driver: "SpannerSyncDriver", table: str | None = None, schema: str | None = None - ) -> list[ForeignKeyMetadata]: + self, driver: "SpannerSyncDriver", table: "str | None" = None, schema: "str | None" = None + ) -> "list[ForeignKeyMetadata]": """Get foreign key metadata.""" schema_name = self.resolve_schema(schema) if table is None: @@ -202,7 +200,7 @@ def _build_domain_result(self, domain: str, rows: list[Any]) -> MetadataResult: return MetadataResult(domain, capability=capability, items=tuple(rows), warnings=capability.warnings) def get_constraints( - self, driver: "SpannerSyncDriver", table: str | None = None, schema: str | None = None + self, driver: "SpannerSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> MetadataResult: """Get constraint metadata for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -219,7 +217,7 @@ def get_constraints( return self._build_domain_result("constraints", rows) def get_sequences( - self, driver: "SpannerSyncDriver", schema: str | None = None, sequence_name: str | None = None + self, driver: "SpannerSyncDriver", schema: "str | None" = None, sequence_name: "str | None" = None ) -> MetadataResult: """Get sequence metadata.""" schema_name = self.resolve_schema(schema) @@ -232,7 +230,7 @@ def get_sequences( return self._build_domain_result("sequences", rows) def get_change_streams( - self, driver: "SpannerSyncDriver", schema: str | None = None, stream_name: str | None = None + self, driver: "SpannerSyncDriver", schema: "str | None" = None, stream_name: "str | None" = None ) -> MetadataResult: """Get change stream metadata.""" schema_name = self.resolve_schema(schema) @@ -279,7 +277,7 @@ def get_system_metadata( ) def get_metadata_capabilities( - self, driver: Any, domains: Sequence[str] | None = None, *, mode: str | None = None + self, driver: Any, domains: "Sequence[str] | None" = None, *, mode: "str | None" = None ) -> MetadataCapabilityProfile: """Get Spanner data-dictionary capability profile.""" _ = driver @@ -294,7 +292,7 @@ def get_ddl( self, driver: Any, object_name: str, - schema: str | None = None, + schema: "str | None" = None, *, object_type: str = "table", include_dependencies: bool = True, @@ -364,7 +362,7 @@ def _spanner_capability_for_domain(domain: str, *, mode: str) -> MetadataCapabil return MetadataCapability.unsupported(domain) -def _normalize_spanner_metadata_mode(mode: str | None) -> str: +def _normalize_spanner_metadata_mode(mode: "str | None") -> str: normalized = (mode or "googlesql").lower() if normalized in {"googlesql", "google_sql", "spanner_googlesql"}: return "googlesql" @@ -373,7 +371,7 @@ def _normalize_spanner_metadata_mode(mode: str | None) -> str: return normalized -def _get_spanner_ddl_statements(driver: Any) -> tuple[str, ...]: +def _get_spanner_ddl_statements(driver: Any) -> "tuple[str, ...]": database = _get_spanner_database(driver) if database is None: return () @@ -407,7 +405,7 @@ def _get_spanner_database_admin_api(database: Any) -> Any: return getattr(client, "database_admin_api", None) -def _select_spanner_ddl_for_object(statements: tuple[str, ...], object_name: str) -> str: +def _select_spanner_ddl_for_object(statements: "tuple[str, ...]", object_name: str) -> str: normalized_name = object_name.lower() for statement in statements: if normalized_name in statement.lower(): diff --git a/sqlspec/builder/_base.py b/sqlspec/builder/_base.py index 8be712735..7fc1bd7ad 100644 --- a/sqlspec/builder/_base.py +++ b/sqlspec/builder/_base.py @@ -39,6 +39,7 @@ MAX_PARAMETER_COLLISION_ATTEMPTS = 1000 PARAMETER_INDEX_PATTERN = re.compile(r"^param_(?P\d+)$") _UPPER_FOLDING_DIALECTS: Final[frozenset[str]] = frozenset({"oracle", "db2"}) +_UNQUOTED_IDENTIFIER_PATTERN = re.compile(r"^[A-Za-z_][\w$#]*$") logger = get_logger(__name__) @@ -1093,7 +1094,7 @@ def _folds_unquoted_to_upper(self, dialect: "DialectType | str | None") -> bool: return str(dialect).lower() in _UPPER_FOLDING_DIALECTS def _unquote_identifiers(self, expression: exp.Expr) -> exp.Expr: - """Return a copy of the expression with identifier quoting removed. + """Remove generated quoting while preserving explicit table identifiers. Upper-folding dialects resolve quoted lowercase names case-sensitively, so quoting is removed to keep lookups aligned with how unquoted DDL created the objects. @@ -1260,6 +1261,10 @@ def __call__(self, node: exp.Expr) -> exp.Expr: def _unquote_identifier(node: exp.Expr) -> exp.Expr: - if isinstance(node, exp.Identifier): + if ( + isinstance(node, exp.Identifier) + and not node.meta.get("sqlspec_explicit_table_quote") + and _UNQUOTED_IDENTIFIER_PATTERN.fullmatch(node.name) + ): node.set("quoted", False) return node diff --git a/sqlspec/builder/_ddl.py b/sqlspec/builder/_ddl.py index 122cdda49..bcb08f8fd 100644 --- a/sqlspec/builder/_ddl.py +++ b/sqlspec/builder/_ddl.py @@ -12,7 +12,7 @@ from typing_extensions import Self from sqlspec.builder._base import BuiltQuery, QueryBuilder -from sqlspec.builder._parsing_utils import _normalize_dialect, parse_table_expression +from sqlspec.builder._parsing_utils import _mark_explicit_table_quotes, _normalize_dialect from sqlspec.builder._select import Select from sqlspec.core import SQL, StatementConfig from sqlspec.exceptions import SQLBuilderError @@ -1653,29 +1653,20 @@ def _parse_column_type(name: str | None, dtype: str, dialect: "DialectType | Non raise SQLBuilderError(msg) from exc -_MIN_QUOTED_IDENTIFIER_LENGTH = 2 - - def _parse_ddl_identifier(name: str, dialect: "DialectType | None" = None) -> exp.Identifier: """Parse a DDL identifier, preserving explicit SQL quoting without double-wrapping.""" stripped = name.strip() if not stripped: return exp.to_identifier(name) - parsed = parse_table_expression(stripped, dialect=dialect) - if isinstance(parsed, exp.Table) and isinstance(parsed.this, exp.Identifier) and not parsed.args.get("db"): + parsed = _mark_explicit_table_quotes(exp.to_table(stripped, dialect=dialect), stripped) + if isinstance(parsed.this, exp.Identifier) and not parsed.args.get("db"): return parsed.this - if len(stripped) >= _MIN_QUOTED_IDENTIFIER_LENGTH and stripped[0] == stripped[-1] == '"': - return exp.Identifier(this=stripped[1:-1].replace('""', '"'), quoted=True) return exp.to_identifier(name) def _parse_ddl_table(table_name: str, schema: "str | None" = None, dialect: "DialectType | None" = None) -> exp.Table: """Parse a DDL table reference, preserving quoted table and schema identifiers.""" - parsed = parse_table_expression(table_name, dialect=dialect) - if isinstance(parsed, exp.Table): - table = parsed.copy() - else: - table = exp.Table(this=_parse_ddl_identifier(table_name, dialect=dialect)) + table = _mark_explicit_table_quotes(exp.to_table(table_name, dialect=dialect), table_name) if schema: table.set("db", _parse_ddl_identifier(schema, dialect=dialect)) return table diff --git a/sqlspec/builder/_dml.py b/sqlspec/builder/_dml.py index 5077c8c3c..1af81575b 100644 --- a/sqlspec/builder/_dml.py +++ b/sqlspec/builder/_dml.py @@ -8,7 +8,11 @@ from typing_extensions import Self from sqlspec.builder._base import BuiltQuery, QueryBuilder -from sqlspec.builder._parsing_utils import extract_expression, extract_sql_object_expression +from sqlspec.builder._parsing_utils import ( + _mark_explicit_table_quotes, + extract_expression, + extract_sql_object_expression, +) from sqlspec.exceptions import SQLBuilderError from sqlspec.protocols import SQLBuilderProtocol from sqlspec.utils.serializers import schema_dump @@ -54,7 +58,7 @@ def from_(self, table: str) -> Self: raise SQLBuilderError(msg) assert current_expr is not None - current_expr.set("this", exp.to_table(table)) + current_expr.set("this", _mark_explicit_table_quotes(exp.to_table(table), table)) return self @@ -76,7 +80,7 @@ def into(self, table: str) -> Self: raise SQLBuilderError(msg) assert current_expr is not None - current_expr.set("this", exp.to_table(table)) + current_expr.set("this", _mark_explicit_table_quotes(exp.to_table(table), table)) return self @@ -251,7 +255,7 @@ def table(self, table_name: str, alias: str | None = None) -> Self: assert current_expr is not None - table_expr: exp.Expr = exp.to_table(table_name, alias=alias) + table_expr: exp.Expr = _mark_explicit_table_quotes(exp.to_table(table_name, alias=alias), table_name) current_expr.set("this", table_expr) return self diff --git a/sqlspec/builder/_parsing_utils.py b/sqlspec/builder/_parsing_utils.py index f0d17d950..12d387a36 100644 --- a/sqlspec/builder/_parsing_utils.py +++ b/sqlspec/builder/_parsing_utils.py @@ -143,6 +143,38 @@ def parse_column_expression(column_input: str | exp.Expr | Any, builder: Any | N return exp.maybe_parse(column_input) or exp.column(str(column_input)) # pyright: ignore[reportArgumentType] +def _mark_explicit_table_quotes(table: exp.Table, source: str) -> exp.Table: + """Distinguish supplied table quotes from quoting added by optimization.""" + if not any(quote in source for quote in ('"', "`", "[")): + return table + remaining = source.strip() + for identifier in table.parts: + if not isinstance(identifier, exp.Identifier) or not remaining: + break + name = identifier.name + opening = remaining[0] + if opening in ('"', "`", "["): + closing = "]" if opening == "[" else opening + spelling = opening + name.replace(closing, closing * 2) + closing + if not remaining.startswith(spelling): + # SQLGlot's permissive fallback can retain literal quote delimiters. + if not (name.startswith(opening) and name.endswith(closing) and remaining.startswith(name)): + break + spelling = name + identifier.set("this", name[1:-1].replace(closing * 2, closing)) + identifier.set("quoted", True) + identifier.meta["sqlspec_explicit_table_quote"] = True + else: + spelling = name + if not remaining.startswith(spelling): + break + remaining = remaining[len(spelling) :].lstrip() + if not remaining.startswith("."): + break + remaining = remaining[1:].lstrip() + return table + + def parse_table_expression( table_input: str, explicit_alias: "str | None" = None, dialect: "DialectType | None" = None ) -> exp.Expr: @@ -156,7 +188,8 @@ def parse_table_expression( parts = table_input.strip().split(None, 1) if len(parts) == ALIAS_PARTS_EXPECTED_COUNT: base_table, alias = parts - return exp.to_table(base_table, alias=alias, dialect=dialect) + if _is_simple_identifier(base_table): + return exp.to_table(base_table, alias=alias, dialect=dialect) if _is_simple_identifier(table_input): return exp.to_table(table_input, alias=explicit_alias, dialect=dialect) @@ -167,11 +200,13 @@ def parse_table_expression( from_clause = parsed.find(exp.From) if from_clause is not None: table_expr = from_clause.this + if isinstance(table_expr, exp.Table): + _mark_explicit_table_quotes(table_expr, table_input) if explicit_alias: return exp.alias_(table_expr, explicit_alias) return table_expr # type: ignore[no-any-return] - return exp.to_table(table_input, alias=explicit_alias, dialect=dialect) + return _mark_explicit_table_quotes(exp.to_table(table_input, alias=explicit_alias, dialect=dialect), table_input) def parse_order_expression(order_input: str | exp.Expr) -> exp.Expr: diff --git a/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/columns.sql b/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/columns.sql index 8cffc294d..57e9ee3ac 100644 --- a/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/columns.sql +++ b/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/columns.sql @@ -19,7 +19,6 @@ SELECT IDENTITY_START_WITH_COUNTER AS identity_start_with_counter FROM INFORMATION_SCHEMA.COLUMNS WHERE (CAST(:schema_name AS STRING) IS NULL OR TABLE_SCHEMA = :schema_name) - AND (CAST(:table_name AS STRING) IS NULL OR TABLE_NAME = :table_name) ORDER BY TABLE_SCHEMA, TABLE_NAME, ORDINAL_POSITION; -- name: by_table diff --git a/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/indexes.sql b/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/indexes.sql index 65e5a5b86..14a45dd54 100644 --- a/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/indexes.sql +++ b/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/indexes.sql @@ -19,7 +19,6 @@ LEFT JOIN INFORMATION_SCHEMA.INDEX_COLUMNS AS ic AND i.TABLE_NAME = ic.TABLE_NAME AND i.INDEX_NAME = ic.INDEX_NAME WHERE (CAST(:schema_name AS STRING) IS NULL OR i.TABLE_SCHEMA = :schema_name) - AND (CAST(:table_name AS STRING) IS NULL OR i.TABLE_NAME = :table_name) GROUP BY i.TABLE_CATALOG, i.TABLE_SCHEMA, diff --git a/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/tables.sql b/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/tables.sql index 82643f41c..7007497fc 100644 --- a/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/tables.sql +++ b/sqlspec/data_dictionary/dialects/spanner/sql/googlesql/tables.sql @@ -11,7 +11,6 @@ SELECT ROW_DELETION_POLICY_EXPRESSION AS row_deletion_policy_expression FROM INFORMATION_SCHEMA.TABLES WHERE (CAST(:schema_name AS STRING) IS NULL OR TABLE_SCHEMA = :schema_name) - AND (CAST(:table_name AS STRING) IS NULL OR TABLE_NAME = :table_name) ORDER BY TABLE_SCHEMA, TABLE_NAME; -- name: by_table diff --git a/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/columns.sql b/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/columns.sql index 9923d21e8..4a581d8c4 100644 --- a/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/columns.sql +++ b/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/columns.sql @@ -13,7 +13,6 @@ SELECT generation_expression FROM information_schema.columns WHERE (:schema_name::text IS NULL OR table_schema = :schema_name) - AND (:table_name::text IS NULL OR table_name = :table_name) ORDER BY table_schema, table_name, ordinal_position; -- name: by_table diff --git a/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/indexes.sql b/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/indexes.sql index 844f08293..512c97d1b 100644 --- a/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/indexes.sql +++ b/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/indexes.sql @@ -15,7 +15,6 @@ LEFT JOIN information_schema.index_columns AS ic AND i.table_name = ic.table_name AND i.index_name = ic.index_name WHERE (:schema_name::text IS NULL OR i.table_schema = :schema_name) - AND (:table_name::text IS NULL OR i.table_name = :table_name) GROUP BY i.table_catalog, i.table_schema, diff --git a/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/tables.sql b/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/tables.sql index 5495b0b59..e9b1fe1f8 100644 --- a/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/tables.sql +++ b/sqlspec/data_dictionary/dialects/spanner/sql/postgresql/tables.sql @@ -7,7 +7,6 @@ SELECT table_type FROM information_schema.tables WHERE (:schema_name::text IS NULL OR table_schema = :schema_name) - AND (:table_name::text IS NULL OR table_name = :table_name) ORDER BY table_schema, table_name; -- name: by_table diff --git a/sqlspec/dialects/db2/_generators.py b/sqlspec/dialects/db2/_generators.py index 1ddbb1e24..65fb68515 100644 --- a/sqlspec/dialects/db2/_generators.py +++ b/sqlspec/dialects/db2/_generators.py @@ -259,7 +259,7 @@ def date_add_sql( def str_position_sql(generator: "generator.Generator", expression: exp.StrPosition) -> str: - """Render a string position search with Db2 POSSTR.""" + """Render a string position search with the matching native Db2 function.""" return render_posstr(generator, expression) diff --git a/sqlspec/dialects/db2/_transforms.py b/sqlspec/dialects/db2/_transforms.py index ad9930586..51dbe73fb 100644 --- a/sqlspec/dialects/db2/_transforms.py +++ b/sqlspec/dialects/db2/_transforms.py @@ -131,9 +131,15 @@ def render_ilike(generator: Any, expression: exp.ILike) -> str: def render_posstr(generator: Any, expression: exp.StrPosition) -> str: - """Map string position functions to Db2 POSSTR(haystack, needle).""" + """Render string searches without discarding start or occurrence arguments.""" this = generator.sql(expression, "this") substr = generator.sql(expression, "substr") + position = generator.sql(expression, "position") + occurrence = generator.sql(expression, "occurrence") + if occurrence: + return f"LOCATE_IN_STRING({this}, {substr}, {position or '1'}, {occurrence})" + if position: + return f"LOCATE({substr}, {this}, {position})" return f"POSSTR({this}, {substr})" diff --git a/sqlspec/dialects/spanner/_generators.py b/sqlspec/dialects/spanner/_generators.py index cec38eaf1..5d64c4762 100644 --- a/sqlspec/dialects/spanner/_generators.py +++ b/sqlspec/dialects/spanner/_generators.py @@ -56,12 +56,12 @@ def _is_post_schema_spanner_property(expression: exp.Expr) -> bool: return expression.this.name.upper() in _SPANNER_PROPERTY_NAMES -def _get_dialect_name(generator: Any) -> str | None: +def _get_dialect_name(generator: Any) -> "str | None": dialect_class = getattr(generator.dialect, "__class__", None) return dialect_class.__name__ if dialect_class else None -def _interval_days(expression: exp.Expr) -> int | None: +def _interval_days(expression: exp.Expr) -> "int | None": """Extract a whole-day count from an interval expression when possible.""" if isinstance(expression, exp.Interval): unit = expression.args.get("unit") @@ -114,7 +114,7 @@ def _render_pg_interval_spec(generator: Any, expression: exp.Expr) -> str: return cast("str", generator.sql(expression)) -def _render_interleave_sql(generator: Any, expression: exp.Property) -> str | None: +def _render_interleave_sql(generator: Any, expression: exp.Property) -> "str | None": """Render INTERLEAVE IN [PARENT] for either dialect, or None if not interleave.""" if not isinstance(expression.this, exp.Literal): return None @@ -137,7 +137,7 @@ def _render_interleave_sql(generator: Any, expression: exp.Property) -> str | No return sql -def _row_deletion_components(expression: exp.Property) -> tuple[exp.Expr, exp.Expr] | None: +def _row_deletion_components(expression: exp.Property) -> "tuple[exp.Expr, exp.Expr] | None": if not isinstance(expression.this, exp.Literal) or expression.this.name.upper() != _ROW_DELETION_NAME: return None values = expression.args.get("value") @@ -350,12 +350,13 @@ def _spanner_create_transform(generator: Any, expression: exp.Create) -> str: """Transform CREATE statements for Spanner including TABLE, SEQUENCE, and CHANGE STREAM.""" if expression.kind == "SEQUENCE": name = generator.sql(expression, "this") + exists = " IF NOT EXISTS" if expression.args.get("exists") else "" props = expression.args.get("properties") if props: opts_list = [f"{generator.sql(p.this)} = {generator.sql(p.args.get('value'))}" for p in props.expressions] opts_str = ", ".join(opts_list) - return f"CREATE SEQUENCE {name} OPTIONS ({opts_str})" - return f"CREATE SEQUENCE {name}" + return f"CREATE SEQUENCE{exists} {name} OPTIONS ({opts_str})" + return f"CREATE SEQUENCE{exists} {name}" if expression.kind == "CHANGE STREAM": name = generator.sql(expression, "this") parts = [f"CREATE CHANGE STREAM {name}"] @@ -689,6 +690,16 @@ def _spangres_anonymous_transform(generator: Any, expression: exp.Anonymous) -> return str(generator.anonymous_sql(expression)) +def _spanner_join_sql(generator: Any, expression: exp.Join) -> str: + """Preserve native join hints in their position after JOIN.""" + sql = str(generator.join_sql(expression)) + hint = expression.args.get("spanner_hint") + if hint is not None and _get_dialect_name(generator) in {"Spanner", "Spangres"}: + prefix, separator, suffix = sql.partition("JOIN ") + return f"{prefix}{separator}{generator.sql(hint)} {suffix}" + return sql + + def _build_function_fallback_transform(expected_dialect: str, original: Any) -> Any: def _transform(generator: Any, expression: exp.Expr) -> str: if _get_dialect_name(generator) == expected_dialect: @@ -719,6 +730,7 @@ def _transform(generator: Any, expression: exp.Expr) -> str: BigQueryGenerator.TRANSFORMS[exp.Hint] = _bq_hint_transform BigQueryGenerator.TRANSFORMS[exp.Select] = _bq_select_transform BigQueryGenerator.TRANSFORMS[exp.Table] = _bq_table_transform +BigQueryGenerator.TRANSFORMS[exp.Join] = _spanner_join_sql BigQueryGenerator.TRANSFORMS[exp.Anonymous] = _spanner_anonymous_transform BigQueryGenerator.TRANSFORMS[CosineDistance] = _build_function_fallback_transform( "Spanner", _original_bq_cosine_distance_transform @@ -736,6 +748,7 @@ def _transform(generator: Any, expression: exp.Expr) -> str: PostgresGenerator.TRANSFORMS[exp.Hint] = _pg_hint_transform PostgresGenerator.TRANSFORMS[exp.Select] = _pg_select_transform PostgresGenerator.TRANSFORMS[exp.Table] = _pg_table_transform +PostgresGenerator.TRANSFORMS[exp.Join] = _spanner_join_sql PostgresGenerator.TRANSFORMS[exp.Anonymous] = _spangres_anonymous_transform PostgresGenerator.TRANSFORMS[CosineDistance] = _build_function_fallback_transform( "Spangres", _original_pg_cosine_distance_transform diff --git a/sqlspec/dialects/spanner/_parsers.py b/sqlspec/dialects/spanner/_parsers.py index 855cf4b58..3283210aa 100644 --- a/sqlspec/dialects/spanner/_parsers.py +++ b/sqlspec/dialects/spanner/_parsers.py @@ -49,7 +49,7 @@ _SPANNER_DIALECT_NAMES: Final[frozenset[str]] = frozenset({"Spangres", "Spanner"}) -def build_interleave_property(parent: exp.Expr, on_delete: str | None = None, in_parent: bool = True) -> exp.Property: +def build_interleave_property(parent: exp.Expr, on_delete: "str | None" = None, in_parent: bool = True) -> exp.Property: """Build the canonical interleave property node.""" if not in_parent: return exp.Property(this=exp.Literal.string(_INTERLEAVE_IN_NAME), value=exp.Tuple(expressions=[parent])) @@ -72,7 +72,7 @@ def _is_spanner_parser(parser: Any) -> bool: return dialect is not None and type(dialect).__name__ in _SPANNER_DIALECT_NAMES -def _parse_interleave(parser: Any) -> exp.Property | None: +def _parse_interleave(parser: Any) -> "exp.Property | None": """Parse ``INTERLEAVE IN [PARENT] table [ON DELETE {CASCADE | NO ACTION}]``. The INTERLEAVE token is already consumed by sqlglot's property dispatch. @@ -94,7 +94,7 @@ def _parse_interleave(parser: Any) -> exp.Property | None: return build_interleave_property(parent, on_delete, in_parent=in_parent) -def _parse_row_deletion_policy(parser: Any) -> exp.Property | None: +def _parse_row_deletion_policy(parser: Any) -> "exp.Property | None": """Parse ``ROW DELETION POLICY (OLDER_THAN(column, INTERVAL n DAY))``. The ROW token is already consumed by sqlglot's property dispatch. @@ -116,7 +116,7 @@ def _parse_row_deletion_policy(parser: Any) -> exp.Property | None: return _build_row_deletion_property(column, interval) -def _parse_ttl(parser: Any) -> exp.Property | None: +def _parse_ttl(parser: Any) -> "exp.Property | None": """Parse ``TTL INTERVAL interval_spec ON column`` into the canonical policy node. The TTL token is already consumed by sqlglot's property dispatch. @@ -176,7 +176,7 @@ def _parse_get_next_sequence_value(parser: Any) -> exp.Anonymous: """Parse GET_NEXT_SEQUENCE_VALUE(SEQUENCE sequence_name).""" if _is_spanner_parser(parser): parser._match_text_seq("SEQUENCE") - seq_name = cast("exp.Expr", parser._parse_id_var()) + seq_name = cast("exp.Expr", parser._parse_table_parts(schema=True)) return get_next_sequence_value(seq_name) return exp.Anonymous(this="GET_NEXT_SEQUENCE_VALUE", expressions=parser._parse_csv(parser._parse_lambda)) @@ -276,14 +276,15 @@ def _parse_create_search_index(parser: Any) -> exp.Index: def _parse_create_sequence(parser: Any) -> exp.Create: """Parse CREATE SEQUENCE name [OPTIONS (...)].""" - name = parser._parse_id_var() + exists = parser._parse_exists(not_=True) + name = parser._parse_table_parts(schema=True) options = _parse_options_properties(parser) - return exp.Create(this=name, kind="SEQUENCE", properties=options) + return exp.Create(this=name, kind="SEQUENCE", exists=exists, properties=options) def _parse_alter_sequence(parser: Any) -> exp.Alter: """Parse ALTER SEQUENCE name SET OPTIONS (...).""" - name = parser._parse_id_var() + name = parser._parse_table_parts(schema=True) options: exp.Properties | None = None if parser._match_text_seq("SET", "OPTIONS") or parser._match_text_seq("OPTIONS"): parser._retreat(parser._index - 1) @@ -425,7 +426,6 @@ def attach_hints(expression: exp.Expr) -> None: continue hint_comments = [c for c in comments if c.strip().startswith("@")] for hc in hint_comments: - comments.remove(hc) hint = parse_hint_expression(hc) target_table: exp.Table | None = None if isinstance(node, exp.Table): @@ -433,11 +433,16 @@ def attach_hints(expression: exp.Expr) -> None: elif isinstance(node, exp.TableAlias) and isinstance(node.parent, exp.Table): target_table = node.parent if target_table is not None: + comments.remove(hc) existing_hints = list(target_table.args.get("hints") or []) existing_hints.append(hint) target_table.set("hints", existing_hints) elif isinstance(node, (exp.Select, exp.Query)): + comments.remove(hc) node.set("hint", hint) + elif isinstance(node, exp.Join): + comments.remove(hc) + node.set("spanner_hint", hint) _original_bq_statement_create: Any = BigQueryParser.STATEMENT_PARSERS.get(TokenType.CREATE) diff --git a/sqlspec/dialects/spanner/_spangres.py b/sqlspec/dialects/spanner/_spangres.py index 9838558d3..508562a7b 100644 --- a/sqlspec/dialects/spanner/_spangres.py +++ b/sqlspec/dialects/spanner/_spangres.py @@ -7,14 +7,16 @@ node so generation always emits the valid PostgreSQL-dialect ``TTL`` form. """ -from typing import Any +from typing import TYPE_CHECKING, Any -from sqlglot import exp from sqlglot.dialects.postgres import Postgres from sqlspec.dialects.spanner._generators import SpangresGenerator from sqlspec.dialects.spanner._parsers import SpangresParser, attach_hints +if TYPE_CHECKING: + from sqlglot import exp + __all__ = ("Spangres",) @@ -24,7 +26,7 @@ class Spangres(Postgres): Parser = SpangresParser Generator = SpangresGenerator - def parse(self, sql: str, **opts: Any) -> list[exp.Expr | None]: + def parse(self, sql: str, **opts: Any) -> "list[exp.Expr | None]": """Parse Spangres SQL statements and attach hints.""" expressions = super().parse(sql, **opts) for expression in expressions: @@ -32,7 +34,7 @@ def parse(self, sql: str, **opts: Any) -> list[exp.Expr | None]: attach_hints(expression) return expressions - def parse_into(self, expression_type: Any, sql: str, **opts: Any) -> list[exp.Expr | None]: + def parse_into(self, expression_type: Any, sql: str, **opts: Any) -> "list[exp.Expr | None]": """Parse into specific expression type with attached hints.""" expressions = super().parse_into(expression_type, sql, **opts) for expression in expressions: diff --git a/sqlspec/dialects/spanner/_spanner.py b/sqlspec/dialects/spanner/_spanner.py index 2234ced3a..9a83f78f6 100644 --- a/sqlspec/dialects/spanner/_spanner.py +++ b/sqlspec/dialects/spanner/_spanner.py @@ -9,7 +9,6 @@ from typing import TYPE_CHECKING, Any -from sqlglot import exp from sqlglot.dialects.bigquery import BigQuery from sqlglot.tokenizer_core import TokenType @@ -17,6 +16,7 @@ from sqlspec.dialects.spanner._parsers import SpannerParser, attach_hints, normalize_spanner_tokens if TYPE_CHECKING: + from sqlglot import exp from sqlglot.tokenizer_core import Token __all__ = ("Spanner",) @@ -40,7 +40,7 @@ class Spanner(BigQuery): Parser = SpannerParser Generator = SpannerGenerator - def parse(self, sql: str, **opts: Any) -> list[exp.Expr | None]: + def parse(self, sql: str, **opts: Any) -> "list[exp.Expr | None]": """Parse Spanner SQL statements and attach hints.""" expressions = super().parse(sql, **opts) for expression in expressions: @@ -48,7 +48,7 @@ def parse(self, sql: str, **opts: Any) -> list[exp.Expr | None]: attach_hints(expression) return expressions - def parse_into(self, expression_type: Any, sql: str, **opts: Any) -> list[exp.Expr | None]: + def parse_into(self, expression_type: Any, sql: str, **opts: Any) -> "list[exp.Expr | None]": """Parse into specific expression type with attached hints.""" expressions = super().parse_into(expression_type, sql, **opts) for expression in expressions: diff --git a/sqlspec/migrations/base.py b/sqlspec/migrations/base.py index a97cc84b3..bf42950e8 100644 --- a/sqlspec/migrations/base.py +++ b/sqlspec/migrations/base.py @@ -114,10 +114,8 @@ def __init__(self, version_table_name: str = "ddl_migrations", version_table_sch else: resolved_schema_identifier = embedded_schema_identifier - bare_name, embedded_schema = self._split_version_table(version_table_name) - resolved_schema = resolved_schema_identifier.name if resolved_schema_identifier is not None else embedded_schema - self.version_table_name = bare_name - self.version_table_schema = resolved_schema + self.version_table_name = table_identifier.name + self.version_table_schema = resolved_schema_identifier.name if resolved_schema_identifier is not None else None self.version_table = self._qualify_version_table( table_identifier.sql(), resolved_schema_identifier.sql() if resolved_schema_identifier is not None else None ) @@ -175,7 +173,9 @@ def _split_version_table(version_table_name: str) -> "tuple[str, str | None]": def _qualify_version_table(self, version_table_name: str, version_table_schema: str | None) -> str: """Return the tracker table name, qualified with schema when configured.""" - return _parse_ddl_table(version_table_name, schema=version_table_schema).sql() + if version_table_schema: + return f"{version_table_schema}.{version_table_name}" + return version_table_name def _tracking_table_builder(self) -> "CreateTable": """Return a CREATE TABLE builder for the tracker table.""" diff --git a/sqlspec/migrations/tracker.py b/sqlspec/migrations/tracker.py index bcefa21ac..7778c6f0a 100644 --- a/sqlspec/migrations/tracker.py +++ b/sqlspec/migrations/tracker.py @@ -223,7 +223,7 @@ def _migrate_schema_if_needed(self, driver: "SyncDriverAdapterBase") -> None: from sqlspec.migrations.schema import SchemaTarget, ensure_schema_sync try: - target = SchemaTarget(self.version_table_name, self._tracking_table_ddl(), self.version_table_schema) + target = SchemaTarget(self.version_table, self._tracking_table_ddl()) result = ensure_schema_sync(driver, [target], manage_schema=True, create_schema=False, assume_existing=True) added_columns = result.added_columns.get(target.identity, []) if not added_columns: @@ -449,7 +449,7 @@ async def _migrate_schema_if_needed(self, driver: "AsyncDriverAdapterBase") -> N from sqlspec.migrations.schema import SchemaTarget, ensure_schema_async try: - target = SchemaTarget(self.version_table_name, self._tracking_table_ddl(), self.version_table_schema) + target = SchemaTarget(self.version_table, self._tracking_table_ddl()) result = await ensure_schema_async( driver, [target], manage_schema=True, create_schema=False, assume_existing=True ) diff --git a/sqlspec/utils/fixtures.py b/sqlspec/utils/fixtures.py index cfd6ca2c7..f9baeb15b 100644 --- a/sqlspec/utils/fixtures.py +++ b/sqlspec/utils/fixtures.py @@ -53,7 +53,7 @@ _POSTGRES_DIALECTS: Final["frozenset[str]"] = frozenset({"postgres", "postgresql"}) _MYSQL_DIALECTS: Final["frozenset[str]"] = frozenset({"mariadb", "mysql"}) _ON_CONFLICT_FAMILIES: Final["frozenset[str]"] = frozenset({"duckdb", "postgres", "sqlite"}) -_JSON_VALUE_FAMILIES: Final["frozenset[str]"] = frozenset({"duckdb", "mysql", "postgres"}) +_JSON_VALUE_FAMILIES: Final["frozenset[str]"] = frozenset({"duckdb", "mysql", "postgres", "sqlite"}) _JSON_TYPE_NAMES: Final["frozenset[str]"] = frozenset({"json", "jsonb"}) _METADATA_FAMILIES: Final["frozenset[str]"] = frozenset({"duckdb", "mysql", "postgres", "sqlite"}) _BINARY_TYPES: Final["frozenset[str]"] = frozenset({ @@ -259,9 +259,9 @@ def load_table_fixtures_sync( DuckDB ``[months, days, nanoseconds]`` interval lists to interval text; strings and numbers in numeric and decimal columns to ``Decimal``; strings in uuid columns to ``UUID``; base64 strings in bytea, blob, binary, and varbinary columns to bytes; and - PostgreSQL, MySQL, and DuckDB json and jsonb values to JSON text. SQLite only converts - binary columns. All other values, including array elements and durations with year or - month parts such as ``P1M``, are passed to the driver as decoded from JSON. Values + PostgreSQL, MySQL, DuckDB, and SQLite json and jsonb values to JSON text. SQLite + only converts binary and JSON columns. All other values, including array elements + and durations with year or month parts such as ``P1M``, are passed to the driver as decoded from JSON. Values of generated columns are ignored. On PostgreSQL, ``GENERATED ALWAYS`` identity columns receive the loaded values through ``OVERRIDING SYSTEM VALUE`` and are never updated by upserts. @@ -363,9 +363,9 @@ async def load_table_fixtures_async( DuckDB ``[months, days, nanoseconds]`` interval lists to interval text; strings and numbers in numeric and decimal columns to ``Decimal``; strings in uuid columns to ``UUID``; base64 strings in bytea, blob, binary, and varbinary columns to bytes; and - PostgreSQL, MySQL, and DuckDB json and jsonb values to JSON text. SQLite only converts - binary columns. All other values, including array elements and durations with year or - month parts such as ``P1M``, are passed to the driver as decoded from JSON. Values + PostgreSQL, MySQL, DuckDB, and SQLite json and jsonb values to JSON text. SQLite + only converts binary and JSON columns. All other values, including array elements + and durations with year or month parts such as ``P1M``, are passed to the driver as decoded from JSON. Values of generated columns are ignored. On PostgreSQL, ``GENERATED ALWAYS`` identity columns receive the loaded values through ``OVERRIDING SYSTEM VALUE`` and are never updated by upserts. @@ -961,8 +961,8 @@ def _drop_generated_values( generated = _generated_column_names(columns) unknown: set[str] = set() if ignore_unknown_columns: - known_lower = {column.name.lower() for column in columns} - unknown = {key for key in rows[0] if key.lower() not in known_lower} + known = {column.name for column in columns} + unknown = set(rows[0]) - known to_drop = (generated & set(rows[0])) | unknown if not to_drop: return @@ -1077,29 +1077,16 @@ def _decode_uuid(value: Any) -> Any: return convert_uuid(value) if isinstance(value, str) else value -def _encode_json_column_value(value: Any) -> str: - """Encode a JSON column value to JSON text, preserving legacy pre-encoded JSON object/array strings.""" - if isinstance(value, str) and value.strip()[:1] in {"{", "["}: - try: - decoded = decode_json(value) - except (ValueError, TypeError): - pass - else: - if isinstance(decoded, (dict, list)): - return encode_json(decoded) - return encode_json(value) - - def _column_value_decoder(family: str, data_type: str) -> "Callable[[Any], Any] | None": """Return the converter from a JSON value to the driver value for a column type, if any.""" if data_type.endswith("]"): return None if data_type in _BINARY_TYPES or data_type.startswith(("binary(", "varbinary(")): return _decode_bytes + if data_type in _JSON_TYPE_NAMES: + return encode_json if family in _JSON_VALUE_FAMILIES else None if family == "sqlite": return None - if data_type in _JSON_TYPE_NAMES: - return _encode_json_column_value if family in _JSON_VALUE_FAMILIES else None if data_type.startswith(("timestamp", "datetime")): return _decode_datetime if data_type == "date": @@ -1261,7 +1248,26 @@ def _table_export_query(dialect: "DialectType", table: str, table_columns: "list column.name for column in table_columns if _is_orderable_type(column.data_type) ] order_by = [_quoted_identifier_sql(name) for name in order_columns] or ["1"] - return Select("*", dialect=dialect).from_(_quoted_table_name(table)).order_by(*order_by) + projections: list[exp.Expr] = [] + if any(column.data_type in _JSON_TYPE_NAMES for column in table_columns): + for column in table_columns: + expression = exp.Column(this=_quoted_identifier(column.name)) + if column.data_type in _JSON_TYPE_NAMES: + # Normalize native JSON decoding across drivers before reading scalar strings. + projections.append( + exp.Alias( + this=exp.Cast(this=expression, to=exp.DataType(this=exp.DataType.Type.TEXT)), + alias=_quoted_identifier(column.name), + ) + ) + else: + projections.append(expression) + return ( + Select(dialect=dialect) + .select(*(projections or [exp.Star()])) + .from_(_quoted_table_name(table)) + .order_by(*order_by) + ) def _sequence_resync_enabled(resync_sequences: bool, dialect_name: str, family: str) -> bool: diff --git a/tests/unit/adapters/test_db2/test_async_pool.py b/tests/unit/adapters/test_db2/test_async_pool.py index dac9899bf..994a510f2 100644 --- a/tests/unit/adapters/test_db2/test_async_pool.py +++ b/tests/unit/adapters/test_db2/test_async_pool.py @@ -324,3 +324,50 @@ async def test_async_connection_create_hook_is_awaited( assert seen == [connection] finally: await pool.close() + + +async def test_cancelled_health_check_closes_detached_connection( + fake_ibm_db: FakeModules, monkeypatch: pytest.MonkeyPatch +) -> None: + pool = Db2AsyncConnectionPool({"database": "TESTDB"}, max_size=1) + connection = _as_fake(await pool.acquire()) + await pool.release(connection) + checking = asyncio.Event() + + async def check(_pool: object, _record: object) -> bool: + checking.set() + await asyncio.Event().wait() + return True + + monkeypatch.setattr(Db2AsyncConnectionPool, "_is_reusable", check) + task = asyncio.create_task(pool.acquire()) + await checking.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert connection.sync_connection.closed is True + assert pool.size() == 0 + replacement = await pool.acquire() + await pool.release(replacement) + await pool.close() + + +async def test_pool_closed_while_opening_rejects_and_closes_connection(fake_ibm_db: FakeModules) -> None: + opening = asyncio.Event() + finish = asyncio.Event() + opened: list[FakeDb2AsyncConnection] = [] + + async def hook(connection: Any) -> None: + opened.append(_as_fake(connection)) + opening.set() + await finish.wait() + + pool = Db2AsyncConnectionPool({"database": "TESTDB"}, max_size=1, on_connection_create=hook) + task = asyncio.create_task(pool.acquire()) + await opening.wait() + await pool.close() + finish.set() + with pytest.raises(DatabaseConnectionError, match="closed"): + await task + assert opened[0].sync_connection.closed is True + assert pool.size() == 0 diff --git a/tests/unit/adapters/test_db2/test_exceptions.py b/tests/unit/adapters/test_db2/test_exceptions.py index 4fc86a7a1..a31b350fe 100644 --- a/tests/unit/adapters/test_db2/test_exceptions.py +++ b/tests/unit/adapters/test_db2/test_exceptions.py @@ -233,7 +233,7 @@ class MockCursor: class NoRowcountCursor: rowcount = -1 - assert resolve_many_rowcount(NoRowcountCursor(), [(1,), (2,)]) == 2 + assert resolve_many_rowcount(NoRowcountCursor(), [(1,), (2,)]) == -1 assert resolve_many_rowcount(NoRowcountCursor(), None) == 0 diff --git a/tests/unit/adapters/test_db2/test_pool.py b/tests/unit/adapters/test_db2/test_pool.py index 6c72fc80f..183a65a38 100644 --- a/tests/unit/adapters/test_db2/test_pool.py +++ b/tests/unit/adapters/test_db2/test_pool.py @@ -206,3 +206,18 @@ def test_pool_is_connection_alive_failure() -> None: result = pool._is_connection_alive(broken_conn) assert result is False + + +def test_failed_creation_hook_closes_unregistered_connection(fake_ibm_db: FakeModules) -> None: + _, module = fake_ibm_db + connection = FakeDb2Connection() + module.pending_connections.append(connection) + + def fail_hook(_connection: object) -> None: + raise RuntimeError("hook failed") + + pool = Db2SyncConnectionPool({"database": "TESTDB"}, on_connection_create=fail_hook) + with pytest.raises(RuntimeError, match="hook failed"): + pool.acquire() + assert connection.closed is True + assert pool.size() == 0 diff --git a/tests/unit/adapters/test_db2/test_transactions.py b/tests/unit/adapters/test_db2/test_transactions.py index 4086e1230..c7fec7382 100644 --- a/tests/unit/adapters/test_db2/test_transactions.py +++ b/tests/unit/adapters/test_db2/test_transactions.py @@ -339,3 +339,21 @@ async def test_begin_requires_ibm_db(db2_mode: DriverMode, monkeypatch: pytest.M with pytest.raises(MissingDependencyError): await db2_mode.call(driver.begin) + + +@pytest.mark.anyio +async def test_session_setup_failure_rolls_back_and_releases(fake_ibm_db: FakeModules, db2_mode: DriverMode) -> None: + connection = _registered_connection(fake_ibm_db) + released: list[tuple[object, object]] = [] + context = db2_mode.session_context(connection, released, begin_transaction=True) + + def fail_prepare(_driver: object) -> None: + raise RuntimeError("driver preparation failed") + + context._prepare_driver = fail_prepare + with pytest.raises(RuntimeError, match="driver preparation failed"): + async with db2_mode.enter(context): + pytest.fail("Session setup must propagate the preparation error") + assert connection.rollbacks == 1 + assert connection.autocommit is True + assert released == [(connection, RuntimeError)] diff --git a/tests/unit/adapters/test_mysql_pool_lifecycle.py b/tests/unit/adapters/test_mysql_pool_lifecycle.py new file mode 100644 index 000000000..444e8ca4e --- /dev/null +++ b/tests/unit/adapters/test_mysql_pool_lifecycle.py @@ -0,0 +1,40 @@ +"""MySQL pool acquisition and shutdown lifecycle regressions.""" + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from sqlspec.adapters.aiomysql.config import AiomysqlConfig +from sqlspec.adapters.asyncmy.config import AsyncmyConfig + + +@pytest.mark.parametrize("config_type", [AiomysqlConfig, AsyncmyConfig]) +async def test_cancelled_connection_hook_releases_pool_checkout(config_type: Any) -> None: + connection = MagicMock() + context = MagicMock() + context.__aenter__ = AsyncMock(return_value=connection) + context.__aexit__ = AsyncMock(return_value=None) + pool = MagicMock() + pool.acquire.return_value = context + hook = AsyncMock(side_effect=asyncio.CancelledError) + config = config_type(connection_instance=pool, driver_features={"on_connection_create": hook}) + + with pytest.raises(asyncio.CancelledError): + async with config.provide_connection(): + pytest.fail("Cancelled hook must prevent connection delivery") + + context.__aexit__.assert_awaited_once() + + +@pytest.mark.parametrize("config_type", [AiomysqlConfig, AsyncmyConfig]) +async def test_pool_shutdown_failure_is_reported_and_pool_retained(config_type: Any) -> None: + pool = MagicMock() + pool.wait_closed = AsyncMock(side_effect=RuntimeError("shutdown failed")) + config = config_type(connection_instance=pool) + + with pytest.raises(RuntimeError, match="shutdown failed"): + await config.close_pool() + + assert config.connection_instance is pool diff --git a/tests/unit/adapters/test_mysqlconnector/test_config.py b/tests/unit/adapters/test_mysqlconnector/test_config.py index e86b23940..815572cac 100644 --- a/tests/unit/adapters/test_mysqlconnector/test_config.py +++ b/tests/unit/adapters/test_mysqlconnector/test_config.py @@ -1,5 +1,6 @@ """Unit tests for mysql-connector configuration modernization.""" +import asyncio from types import SimpleNamespace from typing import TYPE_CHECKING, Any from unittest.mock import AsyncMock, MagicMock @@ -392,3 +393,21 @@ def test_build_connection_config_normalizes_aliases() -> None: assert "db" not in cfg assert cfg["user"] == "alias_user" assert cfg["database"] == "alias_db" + + +@pytest.mark.parametrize("failure", [RuntimeError("hook failed"), asyncio.CancelledError()]) +async def test_async_connection_closes_when_initialization_fails( + monkeypatch: pytest.MonkeyPatch, failure: BaseException +) -> None: + from sqlspec.adapters.mysqlconnector import config as config_module + + connection = MagicMock() + connection.set_autocommit = AsyncMock() + connection.close = AsyncMock() + monkeypatch.setattr(config_module.mysqlconnector_aio, "connect", AsyncMock(return_value=connection)) + config = MysqlConnectorAsyncConfig(driver_features={"on_connection_create": AsyncMock(side_effect=failure)}) + + with pytest.raises(type(failure)): + await config.create_connection() + + connection.close.assert_awaited_once() diff --git a/tests/unit/adapters/test_oracledb/test_config.py b/tests/unit/adapters/test_oracledb/test_config.py index 6db0855e6..da2a08be4 100644 --- a/tests/unit/adapters/test_oracledb/test_config.py +++ b/tests/unit/adapters/test_oracledb/test_config.py @@ -1,8 +1,10 @@ """OracleDB configuration tests covering driver kwargs and typed options.""" from collections.abc import Awaitable, Callable +from inspect import isawaitable from ssl import TLSVersion from typing import Any, cast, get_args, get_origin, get_type_hints +from unittest.mock import AsyncMock, Mock import pytest from oracledb import AuthMode, PoolGetMode, Purity @@ -249,3 +251,39 @@ def test_oracle_config_normalizes_aliases() -> None: assert async_config.connection_config["user"] == "scott" assert "connection_string" not in async_config.connection_config assert "username" not in async_config.connection_config + + +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_pool_close_preserves_native_borrowed_connection_guard(asynchronous: bool) -> None: + pool = Mock() + close = AsyncMock if asynchronous else Mock + pool.close = close(side_effect=RuntimeError("connections remain checked out")) + config = OracleAsyncConfig(connection_instance=pool) if asynchronous else OracleSyncConfig(connection_instance=pool) + + with pytest.raises(RuntimeError, match="checked out"): + result = config._close_pool() + if isawaitable(result): + await result + + pool.close.assert_called_once_with() + assert config.connection_instance is pool + + +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_soda_pool_option_does_not_initialize_global_client_mode( + asynchronous: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + native_pool = Mock() + initialize = Mock(side_effect=AssertionError("client mode must be selected by the application")) + create_pool = Mock(return_value=native_pool) + monkeypatch.setattr(oracle_config_module.oracledb, "init_oracle_client", initialize) + monkeypatch.setattr(oracle_config_module.oracledb, "is_thin_mode", lambda: True) + method = "create_pool_async" if asynchronous else "create_pool" + monkeypatch.setattr(oracle_config_module.oracledb, method, create_pool) + config_type = OracleAsyncConfig if asynchronous else OracleSyncConfig + config = config_type(connection_config={"soda_metadata_cache": True}) + result = config._create_pool() + result = await result if isawaitable(result) else result + + assert result is native_pool + assert create_pool.call_args.kwargs["soda_metadata_cache"] is True diff --git a/tests/unit/adapters/test_oracledb/test_core_row_materialization.py b/tests/unit/adapters/test_oracledb/test_core_row_materialization.py index 3f774240e..f0059597a 100644 --- a/tests/unit/adapters/test_oracledb/test_core_row_materialization.py +++ b/tests/unit/adapters/test_oracledb/test_core_row_materialization.py @@ -1,6 +1,10 @@ # pyright: reportArgumentType=false """Unit tests for Oracle row materialization helpers.""" +from inspect import isawaitable + +import pytest + from sqlspec.adapters.oracledb.core import collect_async_rows, collect_sync_rows, resolve_row_metadata @@ -133,3 +137,15 @@ def test_resolve_row_metadata_cache_contract_still_holds_after_single_pass_rewri assert second_names is first_names assert first_requires_lob is second_requires_lob assert cache[id(description)][0] is description + + +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_collect_rows_coerces_lob_returned_for_non_lob_metadata(asynchronous: bool) -> None: + rows = [(_ReadableValue("custom output"),)] + description = [("PAYLOAD", _TypeCode("DB_TYPE_VARCHAR"))] + collect = collect_async_rows if asynchronous else collect_sync_rows + result = collect(rows, description, {}) + data, column_names = await result if isawaitable(result) else result + + assert data == [("custom output",)] + assert column_names == ["PAYLOAD"] diff --git a/tests/unit/adapters/test_oracledb/test_json_handlers.py b/tests/unit/adapters/test_oracledb/test_json_handlers.py index 67d3fe60b..9912f6080 100644 --- a/tests/unit/adapters/test_oracledb/test_json_handlers.py +++ b/tests/unit/adapters/test_oracledb/test_json_handlers.py @@ -1,6 +1,7 @@ """Unit tests for Oracle native JSON type handlers.""" import inspect +from functools import partial from unittest.mock import Mock import pytest @@ -565,3 +566,29 @@ def test_chained_handlers_are_signature_introspectable() -> None: input_handler = chain_input_handler(Mock(return_value=None), None) assert len(inspect.signature(input_handler).parameters) == 3 + + +@pytest.mark.parametrize("direction", ["input", "output"]) +def test_handler_registration_preserves_unrelated_partial(direction: str) -> None: + inner = Mock(return_value="claimed") + fallback = partial(Mock(return_value="fallback"), inner) + factory = chain_input_handler if direction == "input" else chain_output_handler + arguments = (Mock(), "value", 1) if direction == "input" else (Mock(), Mock()) + + handler = factory(inner, fallback) + + assert handler(*arguments) == "claimed" + inner.assert_called_once_with(*arguments) + + +@pytest.mark.parametrize("direction", ["input", "output"]) +def test_handler_registration_is_idempotent(direction: str) -> None: + inner = Mock(return_value=None) + fallback = Mock(return_value="fallback") + factory = chain_input_handler if direction == "input" else chain_output_handler + arguments = (Mock(), "value", 1) if direction == "input" else (Mock(), Mock()) + handler = factory(inner, factory(inner, fallback)) + + assert handler(*arguments) == "fallback" + inner.assert_called_once_with(*arguments) + fallback.assert_called_once_with(*arguments) diff --git a/tests/unit/adapters/test_oracledb/test_pipeline_helpers.py b/tests/unit/adapters/test_oracledb/test_pipeline_helpers.py index 79c4037fd..1a2c91110 100644 --- a/tests/unit/adapters/test_oracledb/test_pipeline_helpers.py +++ b/tests/unit/adapters/test_oracledb/test_pipeline_helpers.py @@ -7,10 +7,10 @@ from sqlspec import StatementStack from sqlspec.adapters.oracledb._typing import OracleAsyncConnection, OracleSyncConnection from sqlspec.adapters.oracledb.core import ( - _normalize_execute_many_parameters_async, - _normalize_execute_many_parameters_sync, build_pipeline_stack_result, default_statement_config, + normalize_execute_many_parameters_async, + normalize_execute_many_parameters_sync, ) from sqlspec.adapters.oracledb.data_dictionary import OracledbAsyncDataDictionary from sqlspec.adapters.oracledb.driver import OracleAsyncDriver, OracleSyncDriver @@ -242,20 +242,20 @@ async def test_async_pipeline_gate_accepts_thin_26ai(monkeypatch: pytest.MonkeyP def test_normalize_execute_many_normalize_execute_many_parameters_sync_tuple_to_list() -> None: parameters = ({"x": 1}, {"x": 2}) - result = _normalize_execute_many_parameters_sync(parameters) + result = normalize_execute_many_parameters_sync(parameters) assert result == [{"x": 1}, {"x": 2}] assert isinstance(result, list) def test_normalize_execute_many_normalize_execute_many_parameters_async_tuple_to_list() -> None: parameters = ({"x": 1}, {"x": 2}) - result = _normalize_execute_many_parameters_async(parameters) + result = normalize_execute_many_parameters_async(parameters) assert result == [{"x": 1}, {"x": 2}] assert isinstance(result, list) @pytest.mark.parametrize( - "normalizer", [_normalize_execute_many_parameters_sync, _normalize_execute_many_parameters_async] + "normalizer", [normalize_execute_many_parameters_sync, normalize_execute_many_parameters_async] ) def test_normalize_execute_many_normalize_execute_many_parameters_passes_empty_through( normalizer: Callable[[object], object], diff --git a/tests/unit/adapters/test_pymysql/test_config.py b/tests/unit/adapters/test_pymysql/test_config.py index 5f9873a52..1cde06fc6 100644 --- a/tests/unit/adapters/test_pymysql/test_config.py +++ b/tests/unit/adapters/test_pymysql/test_config.py @@ -241,3 +241,19 @@ def test_pymysql_config_with_dsn_applies_to_connection_parameters() -> None: assert pool._connection_parameters["password"] == "pass1" assert pool._connection_parameters["host"] == "dbhost" assert pool._connection_parameters["database"] == "production" + + +def test_pool_discards_connection_after_rollback_failure() -> None: + broken = MagicMock(server_status=1) + broken.rollback.side_effect = RuntimeError("rollback failed") + replacement = MagicMock(server_status=0) + factory = MagicMock(side_effect=[broken, replacement]) + pool = PyMysqlConnectionPool({}, connection_factory=factory) + try: + connection = pool.acquire() + with pytest.raises(RuntimeError, match="rollback failed"): + pool.release(connection) + broken.close.assert_called_once() + assert pool.acquire() is replacement + finally: + pool.close() diff --git a/tests/unit/adapters/test_spanner/test_data_dictionary_routing.py b/tests/unit/adapters/test_spanner/test_data_dictionary_routing.py index 8de5c5478..fbce9f8ab 100644 --- a/tests/unit/adapters/test_spanner/test_data_dictionary_routing.py +++ b/tests/unit/adapters/test_spanner/test_data_dictionary_routing.py @@ -2,6 +2,8 @@ from typing import Any, cast +import pytest + from sqlspec.adapters.spanner.data_dictionary import SpannerDataDictionary from sqlspec.data_dictionary import ColumnMetadata, ForeignKeyMetadata, IndexMetadata, TableMetadata @@ -120,3 +122,15 @@ def test_get_query_explicit_mode_override() -> None: index_query = dictionary.get_query("indexes", "by_schema", mode="postgresql") assert "i.index_name AS index_name" in index_query.raw_sql assert "AS columns" in index_query.raw_sql + + +@pytest.mark.parametrize("mode", ["googlesql", "postgresql"]) +@pytest.mark.parametrize("domain", ["tables", "columns", "indexes"]) +def test_schema_metadata_binds_only_schema(mode: str, domain: str) -> None: + dictionary = SpannerDataDictionary(mode=mode) + driver = MockSpannerDriver() + getattr(dictionary, f"get_{domain}")(cast("Any", driver), schema="public") + + statement = driver.last_query.copy(parameters={"schema_name": driver.last_params["schema_name"]}) + _sql, parameters = statement.compile() + assert parameters == ("public", "public") diff --git a/tests/unit/builder/test_ddl_builder.py b/tests/unit/builder/test_ddl_builder.py index 133d4ed1a..a5cd6d07f 100644 --- a/tests/unit/builder/test_ddl_builder.py +++ b/tests/unit/builder/test_ddl_builder.py @@ -1,7 +1,7 @@ """Regression tests for DDL builder Wave 1 fixes.""" import pytest -from sqlglot import exp +from sqlglot import exp, parse_one from sqlspec import sql from sqlspec.builder._ddl import ( @@ -238,3 +238,73 @@ def test_ddl_builders_preserve_quoted_and_schema_qualified_identifiers() -> None assert alter_from_in_schema.sql.startswith('ALTER TABLE "App"."Tracker" ADD COLUMN') for stmt in (create_from_qualified, create_from_in_schema, drop_stmt, alter_from_qualified, alter_from_in_schema): assert '""' not in stmt.sql + + +def test_ddl_identifiers_with_spaces_and_escaped_quotes() -> None: + config = StatementConfig(dialect="postgres") + schema = '"App Schema"' + table = '"Migration ""Tracker"""' + qualified = f"{schema}.{table}" + + create = sql.create_table(table).in_schema(schema).column("id", "INT").to_statement(config) + alter = sql.alter_table(qualified).add_column("name", "TEXT").to_statement(config) + drop = sql.drop_table(qualified).to_statement(config) + + assert create.sql.startswith(f"CREATE TABLE {qualified} (") + assert alter.sql.startswith(f"ALTER TABLE {qualified} ADD COLUMN") + assert drop.sql == f"DROP TABLE {qualified}" + + +@pytest.mark.parametrize("dialect", ["oracle", "db2", "snowflake"]) +@pytest.mark.parametrize("table", ['"History"."Migration Tracker"', "history.tracker"]) +def test_ddl_quoting_survives_dialect_rendering(dialect: str, table: str) -> None: + expected = '"history"."tracker"' if dialect == "snowflake" and table == "history.tracker" else table + config = StatementConfig(dialect=dialect) + builders = ( + (sql.create_table(table).column("id", "INT"), "CREATE TABLE"), + (sql.alter_table(table).add_column("name", "TEXT"), "ALTER TABLE"), + (sql.drop_table(table), "DROP TABLE"), + ) + for builder, command in builders: + assert builder.build(dialect=dialect).sql.startswith(f"{command} {expected}") + assert builder.to_statement(config).sql.startswith(f"{command} {expected}") + + +@pytest.mark.parametrize("dialect", ["oracle", "db2"]) +@pytest.mark.parametrize("quoted", [True, False]) +def test_tracking_table_dml_preserves_table_quotes(dialect: str, quoted: bool) -> None: + table = '"History"."Migration Tracker"' if quoted else "history.tracker" + config = StatementConfig(dialect=dialect) + builders = ( + sql.select("version_num").from_(table), + sql.insert(table).columns("version_num").values("0001"), + sql.update(table).set("version_num", "0002"), + sql.delete().from_(table), + ) + for builder in builders: + statement = builder.to_statement(config) + assert builder.to_statement(config).sql == statement.sql + parsed = parse_one(statement.sql, read=dialect) + target = parsed.find(exp.Table) + assert isinstance(target, exp.Table) + assert target.name == ("Migration Tracker" if quoted else "tracker") + assert target.db == ("History" if quoted else "history") + assert target.this.quoted is quoted + assert target.args["db"].quoted is quoted + + +@pytest.mark.parametrize("table", ["tracker$log", "tracker#log"]) +def test_oracle_bare_table_special_characters_keep_folding(table: str) -> None: + statement = sql.select("version_num").from_(table).to_statement(StatementConfig(dialect="oracle")) + assert f"FROM {table} {table}" in statement.sql + + +@pytest.mark.parametrize("schema", ["App", "tracker#log"]) +def test_oracle_mixed_table_quote_provenance(schema: str) -> None: + table = f'"{schema}".tracker#log' + config = StatementConfig(dialect="oracle") + query = sql.select("version_num").from_(table).to_statement(config) + ddl = sql.create_table(table).column("id", "INT").to_statement(config) + + assert f"FROM {table} tracker#log" in query.sql + assert ddl.sql.startswith(f"CREATE TABLE {table} (") diff --git a/tests/unit/dialects/test_db2.py b/tests/unit/dialects/test_db2.py index 8162e92a8..047d33e78 100644 --- a/tests/unit/dialects/test_db2.py +++ b/tests/unit/dialects/test_db2.py @@ -174,3 +174,18 @@ def test_set_operation_pagination() -> None: "SELECT id FROM a UNION ALL SELECT id FROM b ORDER BY id LIMIT 3 OFFSET 2", read="postgres", write="db2" )[0] assert result == "SELECT id FROM a UNION ALL SELECT id FROM b ORDER BY id OFFSET 2 ROWS FETCH NEXT 3 ROWS ONLY" + + +def test_locate_preserves_native_start_argument() -> None: + sql = "SELECT LOCATE('x', 'x-x', 3) FROM SYSIBM.SYSDUMMY1" + assert parse_one(sql, dialect="db2").sql(dialect="db2") == sql + + +def test_string_position_preserves_occurrence() -> None: + expression = exp.StrPosition( + this=exp.Literal.string("x-x-x"), + substr=exp.Literal.string("x"), + position=exp.Literal.number(2), + occurrence=exp.Literal.number(2), + ) + assert expression.sql(dialect="db2") == "LOCATE_IN_STRING('x-x-x', 'x', 2, 2)" diff --git a/tests/unit/dialects/test_spanner_hints.py b/tests/unit/dialects/test_spanner_hints.py index b3c4cd3c2..14e0f7226 100644 --- a/tests/unit/dialects/test_spanner_hints.py +++ b/tests/unit/dialects/test_spanner_hints.py @@ -97,3 +97,15 @@ def test_normalize_spanner_tokens_without_sql_fallback() -> None: normalized = normalize_spanner_tokens(raw_tokens) assert len(normalized) == 2 assert normalized[0].comments == ["@ LOCK_SCANNED_RANGES=exclusive, OPTIMIZER_VERSION=6"] + + +def test_join_hint_round_trip() -> None: + sql = "SELECT * FROM t JOIN @{JOIN_METHOD=HASH_JOIN} u ON t.id = u.id" + parsed = parse_one(sql, dialect="spanner") + assert parsed.sql(dialect="spanner") == sql + assert "JOIN /*@ JOIN_METHOD=HASH_JOIN */ u" in parsed.sql(dialect="spangres") + + +def test_at_comment_on_expression_is_preserved() -> None: + sql = "SELECT 1 /* @ ordinary comment */" + assert "@ ordinary comment" in parse_one(sql, dialect="spanner").sql(dialect="spanner") diff --git a/tests/unit/dialects/test_spanner_sequences.py b/tests/unit/dialects/test_spanner_sequences.py index f6d17d7c0..46b85e969 100644 --- a/tests/unit/dialects/test_spanner_sequences.py +++ b/tests/unit/dialects/test_spanner_sequences.py @@ -1,5 +1,6 @@ """Unit tests for Cloud Spanner Sequences (CREATE, ALTER, DROP, GET_NEXT_SEQUENCE_VALUE).""" +import pytest from sqlglot import exp, parse_one @@ -82,3 +83,14 @@ def test_get_next_sequence_value_in_parenthesized_expressions() -> None: spangres_sql = "SELECT GET_NEXT_SEQUENCE_VALUE(SEQUENCE customer_seq)" parsed_spangres = parse_one(spangres_sql, dialect="spangres") assert "GET_NEXT_SEQUENCE_VALUE(SEQUENCE customer_seq)" in parsed_spangres.sql(dialect="spangres") + + +@pytest.mark.parametrize("dialect", ["spanner", "spangres"]) +def test_qualified_sequence_value(dialect: str) -> None: + sql = "SELECT GET_NEXT_SEQUENCE_VALUE(SEQUENCE catalog.seq)" + assert parse_one(sql, dialect=dialect).sql(dialect=dialect) == sql + + +def test_create_sequence_if_not_exists() -> None: + sql = "CREATE SEQUENCE IF NOT EXISTS catalog.seq OPTIONS (sequence_kind = 'bit_reversed_positive')" + assert parse_one(sql, dialect="spanner").sql(dialect="spanner") == sql diff --git a/tests/unit/migrations/test_tracker_idempotency.py b/tests/unit/migrations/test_tracker_idempotency.py index a756fc42d..56bcecfc4 100644 --- a/tests/unit/migrations/test_tracker_idempotency.py +++ b/tests/unit/migrations/test_tracker_idempotency.py @@ -7,6 +7,7 @@ import pytest from sqlspec.adapters.oracledb.migrations import OracleAsyncMigrationTracker, OracleSyncMigrationTracker +from sqlspec.core import StatementConfig from sqlspec.driver import AsyncDriverAdapterBase, SyncDriverAdapterBase from sqlspec.migrations.tracker import AsyncMigrationTracker, SyncMigrationTracker @@ -110,6 +111,19 @@ def test_oracle_tracker_preserves_mixed_case_qualified_tracking_table() -> None: assert tracker.version_table_name == "DdlMigrations" +@pytest.mark.parametrize("tracker_type", [OracleSyncMigrationTracker, OracleAsyncMigrationTracker]) +def test_oracle_tracking_ddl_preserves_explicit_case(tracker_type: Any) -> None: + tracker = tracker_type(version_table_name='"app"."Migration Tracker"') + ddl = tracker._tracking_table_ddl().to_statement(StatementConfig(dialect="oracle")) + + assert tracker.version_table == '"app"."Migration Tracker"' + assert tracker.version_table_name == "Migration Tracker" + assert tracker.version_table_schema == "app" + assert ddl.sql.startswith('CREATE TABLE "app"."Migration Tracker" (') + query = tracker._current_version_query().to_statement(StatementConfig(dialect="oracle")) + assert 'FROM "app"."Migration Tracker"' in query.sql + + def test_oracle_sync_tracker_introspects_unmodified_table_name_with_schema() -> None: """Oracle data-dictionary calls should preserve mixed-case configured identifiers.""" tracker = OracleSyncMigrationTracker(version_table_name="DdlMigrations", version_table_schema="AppOwner") diff --git a/tests/unit/migrations/test_version.py b/tests/unit/migrations/test_version.py index cd4082f44..9da6148ac 100644 --- a/tests/unit/migrations/test_version.py +++ b/tests/unit/migrations/test_version.py @@ -364,7 +364,8 @@ def test_parse_extension_stem_round_trips_through_parse_version() -> None: assert parsed.sequence == 1 -def test_quoted_and_schema_qualified_version_table_ddl_and_tracker() -> None: +@pytest.mark.parametrize("table_name,schema_name", [("Tracker", "App"), ("Migration Tracker", "App Schema")]) +def test_quoted_and_schema_qualified_version_table_ddl_and_tracker(table_name: str, schema_name: str) -> None: """Quoted and schema-qualified version_table identifiers render valid SQL without doubled quotes.""" pg_config = StatementConfig(dialect="postgres") @@ -388,18 +389,21 @@ def test_quoted_and_schema_qualified_version_table_ddl_and_tracker() -> None: assert '""' not in alter_stmt.sql for tracker in ( - SyncMigrationTracker(version_table_name='"App"."Tracker"'), - SyncMigrationTracker(version_table_name='"Tracker"', version_table_schema='"App"'), + SyncMigrationTracker(version_table_name=f'"{schema_name}"."{table_name}"'), + SyncMigrationTracker(version_table_name=f'"{table_name}"', version_table_schema=f'"{schema_name}"'), ): - assert tracker.version_table_name == "Tracker" - assert tracker.version_table_schema == "App" - assert tracker.version_table == '"App"."Tracker"' + qualified = f'"{schema_name}"."{table_name}"' + assert tracker.version_table_name == table_name + assert tracker.version_table_schema == schema_name + assert tracker.version_table == qualified + version_query = tracker._current_version_query().to_statement(pg_config) # pyright: ignore[reportPrivateUsage] + assert f"FROM {qualified}" in version_query.sql ddl_builder = tracker._tracking_table_ddl() # pyright: ignore[reportPrivateUsage] ddl_sql = str(ddl_builder) pg_ddl_sql = ddl_builder.to_statement(pg_config).sql - assert 'CREATE TABLE IF NOT EXISTS "App"."Tracker"' in ddl_sql - assert 'CREATE TABLE IF NOT EXISTS "App"."Tracker"' in pg_ddl_sql + assert f"CREATE TABLE IF NOT EXISTS {qualified}" in ddl_sql + assert f"CREATE TABLE IF NOT EXISTS {qualified}" in pg_ddl_sql assert '""' not in ddl_sql assert '""' not in pg_ddl_sql @@ -411,11 +415,11 @@ def test_quoted_and_schema_qualified_version_table_ddl_and_tracker() -> None: tracker.ensure_tracking_table(driver) - driver.data_dictionary.get_columns.assert_called_once_with(driver, "Tracker", schema="App") + driver.data_dictionary.get_columns.assert_called_once_with(driver, table_name, schema=schema_name) assert driver.execute.call_count == 2 alter_builder = driver.execute.call_args_list[1].args[0] alter_rendered = alter_builder.to_statement(pg_config).sql - assert 'ALTER TABLE "App"."Tracker" ADD COLUMN' in alter_rendered + assert f"ALTER TABLE {qualified} ADD COLUMN" in alter_rendered assert '""' not in alter_rendered diff --git a/tests/unit/utils/test_fixtures.py b/tests/unit/utils/test_fixtures.py index 9f4db6aab..a71a9847f 100644 --- a/tests/unit/utils/test_fixtures.py +++ b/tests/unit/utils/test_fixtures.py @@ -1386,6 +1386,9 @@ async def test_export_decodes_json_and_jsonb_string_cells_with_fallback_async(tm await export_table_fixtures_async(driver, tmp_path, ["events"], compress=False) + query = driver.select.call_args.args[0].build().sql + assert 'CAST("events"."payload" AS TEXT) AS "payload"' in query + assert 'CAST("events"."raw_json" AS TEXT) AS "raw_json"' in query exported = json.loads((tmp_path / "events.json").read_text()) assert exported == [ {"id": 1, "payload": {"a": 1}, "raw_json": "scalar"}, @@ -1402,16 +1405,16 @@ async def test_export_decodes_json_and_jsonb_string_cells_with_fallback_async(tm (42, "42"), ({"k": "v"}, '{"k":"v"}'), ([1, 2], "[1,2]"), - ('{"k": "v"}', '{"k":"v"}'), - (" [1, 2] ", "[1,2]"), + ('{"k": "v"}', json.dumps('{"k": "v"}')), + (" [1, 2] ", json.dumps(" [1, 2] ")), ("{not valid json", '"{not valid json"'), ], ) -def test_encode_json_column_value_handles_scalars_and_legacy_pre_encoded_strings( - raw_value: Any, expected_json_text: str -) -> None: - """JSON column encoding serializes scalars and structures without double-encoding legacy JSON object/array strings.""" - encoded = fixture_module._encode_json_column_value(raw_value) +def test_fixture_json_conversion_preserves_scalar_types(raw_value: Any, expected_json_text: str) -> None: + """Fixture values represent JSON values, including strings that look like objects or arrays.""" + decoder = fixture_module._column_value_decoder("duckdb", "json") + assert callable(decoder) + encoded = decoder(raw_value) assert json.loads(encoded) == json.loads(expected_json_text) @@ -1585,3 +1588,42 @@ async def test_exclude_update_columns_async_and_validation(tmp_path: Path) -> No with pytest.raises(ValueError, match="Invalid column name"): await load_table_fixtures_async(driver, tmp_path, exclude_update_columns=["bad col"]) await config.close_pool() + + +@pytest.mark.parametrize("config_type", [SqliteConfig, DuckDBConfig]) +def test_json_fixture_roundtrip_preserves_string_scalar_types( + tmp_path: Path, config_type: type[SqliteConfig] | type[DuckDBConfig] +) -> None: + values = ["true", "42", "null", "[1, 2]", '{"key": 1}', True, 42, [1, 2], {"key": 1}, None] + config = config_type(connection_config={"database": ":memory:"}) + try: + with config.provide_session() as driver: + driver.execute("CREATE TABLE json_values (id INTEGER PRIMARY KEY, payload JSON)") + driver.execute_many( + "INSERT INTO json_values (id, payload) VALUES (:id, :payload)", + [ + {"id": index, "payload": None if value is None else json.dumps(value)} + for index, value in enumerate(values) + ], + ) + export_table_fixtures_sync(driver, tmp_path, ["json_values"], compress=False) + assert [row["payload"] for row in json.loads((tmp_path / "json_values.json").read_text())] == values + driver.execute("DELETE FROM json_values") + load_table_fixtures_sync(driver, tmp_path) + loaded = driver.select("SELECT CAST(payload AS TEXT) AS payload FROM json_values ORDER BY id") + assert [None if row["payload"] is None else json.loads(row["payload"]) for row in loaded] == values + finally: + config.close_pool() + + +def test_ignore_unknown_fixture_columns_uses_exact_quoted_names(tmp_path: Path) -> None: + (tmp_path / "users.json").write_text(json.dumps([{"id": 1, "Name": "wrong column"}]), encoding="utf-8") + driver = MagicMock() + driver.statement_config.dialect = "postgres" + driver.data_dictionary.get_columns.return_value = [ + {"column_name": "id", "data_type": "integer", "is_primary": True}, + {"column_name": "name", "data_type": "text", "is_primary": False}, + ] + + assert load_table_fixtures_sync(driver, tmp_path, ignore_unknown_columns=True) == {"users": 1} + assert driver.execute_many.call_args.args[1] == [{"id": 1}] From d417e4cb939cb809f79363e3d51570ed3266da78 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:29:37 +0000 Subject: [PATCH 02/10] fix: preserve synchronous Oracle Thick mode configuration --- sqlspec/adapters/oracledb/config.py | 17 +++++ .../adapters/test_oracledb/test_config.py | 75 ++++++++++++++++--- 2 files changed, 81 insertions(+), 11 deletions(-) diff --git a/sqlspec/adapters/oracledb/config.py b/sqlspec/adapters/oracledb/config.py index 00d6c0820..6a1ef2fda 100644 --- a/sqlspec/adapters/oracledb/config.py +++ b/sqlspec/adapters/oracledb/config.py @@ -42,6 +42,7 @@ SyncPoolConnectionContext, SyncPoolSessionFactory, ) +from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.events import EventRuntimeHints from sqlspec.utils.config_tools import normalize_connection_config @@ -134,6 +135,8 @@ class OracleConnectionParams(TypedDict): class OraclePoolParams(OracleConnectionParams): """OracleDB pool parameters.""" + thick_mode: NotRequired[bool] + lib_dir: NotRequired[str] pool_class: NotRequired[type[Any]] params: NotRequired[oracledb.PoolParams] min: NotRequired[int] @@ -407,6 +410,15 @@ def get_event_runtime_hints(self) -> "EventRuntimeHints": def _create_pool(self) -> "OracleSyncConnectionPool": """Create the actual connection pool.""" config = dict(self.connection_config) + thick_mode = config.pop("thick_mode", False) + lib_dir = config.pop("lib_dir", None) + if (thick_mode or lib_dir is not None or config.get("soda_metadata_cache")) and oracledb.is_thin_mode(): + client_config = {} + if lib_dir is not None: + client_config["lib_dir"] = lib_dir + if config.get("config_dir") is not None: + client_config["config_dir"] = config["config_dir"] + oracledb.init_oracle_client(**client_config) config.pop("threaded", None) config["session_callback"] = self._init_connection @@ -599,6 +611,11 @@ def get_event_runtime_hints(self) -> "EventRuntimeHints": async def _create_pool(self) -> "OracleAsyncConnectionPool": """Create the actual async connection pool.""" config = dict(self.connection_config) + thick_mode = config.pop("thick_mode", False) + lib_dir = config.pop("lib_dir", None) + if thick_mode or lib_dir is not None: + msg = "OracleAsyncConfig only supports Thin mode; use OracleSyncConfig for Thick mode." + raise ImproperConfigurationError(msg) config.pop("threaded", None) config["session_callback"] = self._init_connection diff --git a/tests/unit/adapters/test_oracledb/test_config.py b/tests/unit/adapters/test_oracledb/test_config.py index da2a08be4..f98acde0c 100644 --- a/tests/unit/adapters/test_oracledb/test_config.py +++ b/tests/unit/adapters/test_oracledb/test_config.py @@ -4,7 +4,7 @@ from inspect import isawaitable from ssl import TLSVersion from typing import Any, cast, get_args, get_origin, get_type_hints -from unittest.mock import AsyncMock, Mock +from unittest.mock import AsyncMock, Mock, call import pytest from oracledb import AuthMode, PoolGetMode, Purity @@ -19,6 +19,7 @@ OraclePoolParams, OracleSyncConfig, ) +from sqlspec.exceptions import ImproperConfigurationError class _StubConnection: @@ -269,21 +270,73 @@ async def test_pool_close_preserves_native_borrowed_connection_guard(asynchronou assert config.connection_instance is pool -@pytest.mark.parametrize("asynchronous", [False, True]) -async def test_soda_pool_option_does_not_initialize_global_client_mode( - asynchronous: bool, monkeypatch: pytest.MonkeyPatch +@pytest.mark.parametrize("options", [{"thick_mode": True}, {"lib_dir": "/oracle/lib"}, {"soda_metadata_cache": True}]) +@pytest.mark.parametrize("thin_mode", [False, True]) +def test_sync_pool_initializes_requested_thick_mode( + options: dict[str, Any], thin_mode: bool, monkeypatch: pytest.MonkeyPatch ) -> None: - native_pool = Mock() - initialize = Mock(side_effect=AssertionError("client mode must be selected by the application")) - create_pool = Mock(return_value=native_pool) + initialize = Mock() + create_pool = Mock() + monkeypatch.setattr(oracle_config_module.oracledb, "init_oracle_client", initialize) + monkeypatch.setattr(oracle_config_module.oracledb, "is_thin_mode", lambda: thin_mode) + monkeypatch.setattr(oracle_config_module.oracledb, "create_pool", create_pool) + config = OracleSyncConfig(connection_config={**options, "config_dir": "/oracle/config"}) + + assert config._create_pool() is create_pool.return_value + + expected = {"config_dir": "/oracle/config"} + if "lib_dir" in options: + expected["lib_dir"] = options["lib_dir"] + assert initialize.call_args_list == ([call(**expected)] if thin_mode else []) + assert create_pool.call_args.kwargs == { + **{key: value for key, value in options.items() if key not in {"thick_mode", "lib_dir"}}, + "config_dir": "/oracle/config", + "session_callback": config._init_connection, + } + + +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_explicit_thin_mode_is_consumed(asynchronous: bool, monkeypatch: pytest.MonkeyPatch) -> None: + initialize = Mock() + create_pool = Mock() monkeypatch.setattr(oracle_config_module.oracledb, "init_oracle_client", initialize) monkeypatch.setattr(oracle_config_module.oracledb, "is_thin_mode", lambda: True) - method = "create_pool_async" if asynchronous else "create_pool" - monkeypatch.setattr(oracle_config_module.oracledb, method, create_pool) + monkeypatch.setattr( + oracle_config_module.oracledb, "create_pool_async" if asynchronous else "create_pool", create_pool + ) config_type = OracleAsyncConfig if asynchronous else OracleSyncConfig - config = config_type(connection_config={"soda_metadata_cache": True}) + config = config_type(connection_config={"thick_mode": False}) result = config._create_pool() result = await result if isawaitable(result) else result - assert result is native_pool + assert result is create_pool.return_value + initialize.assert_not_called() + assert create_pool.call_args.kwargs == {"session_callback": config._init_connection} + + +@pytest.mark.parametrize("options", [{"thick_mode": True}, {"lib_dir": "/oracle/lib"}]) +async def test_async_pool_rejects_thick_mode(options: dict[str, Any], monkeypatch: pytest.MonkeyPatch) -> None: + initialize = Mock() + create_pool = Mock() + monkeypatch.setattr(oracle_config_module.oracledb, "init_oracle_client", initialize) + monkeypatch.setattr(oracle_config_module.oracledb, "create_pool_async", create_pool) + config = OracleAsyncConfig(connection_config=options) + + with pytest.raises(ImproperConfigurationError, match="only supports Thin mode"): + await config._create_pool() + + initialize.assert_not_called() + create_pool.assert_not_called() + + +async def test_async_soda_option_does_not_initialize_thick_mode(monkeypatch: pytest.MonkeyPatch) -> None: + initialize = Mock() + create_pool = Mock() + monkeypatch.setattr(oracle_config_module.oracledb, "init_oracle_client", initialize) + monkeypatch.setattr(oracle_config_module.oracledb, "create_pool_async", create_pool) + config = OracleAsyncConfig(connection_config={"soda_metadata_cache": True}) + + assert await config._create_pool() is create_pool.return_value + + initialize.assert_not_called() assert create_pool.call_args.kwargs["soda_metadata_cache"] is True From 6f5139000b5f24dcf3285ac1b98d4ec9471277c0 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:31:16 +0000 Subject: [PATCH 03/10] fix(mysqlconnector): retain native async pooling with compatibility fallback --- sqlspec/adapters/mysqlconnector/_typing.py | 6 ++ sqlspec/adapters/mysqlconnector/config.py | 80 +++++++++++++++++-- .../test_mysqlconnector/test_config.py | 42 ++++++++++ 3 files changed, 121 insertions(+), 7 deletions(-) diff --git a/sqlspec/adapters/mysqlconnector/_typing.py b/sqlspec/adapters/mysqlconnector/_typing.py index 3d0dc3568..1ad16a74d 100644 --- a/sqlspec/adapters/mysqlconnector/_typing.py +++ b/sqlspec/adapters/mysqlconnector/_typing.py @@ -17,11 +17,15 @@ from mysql.connector.constants import FieldType as _MysqlConnectorFieldType from mysql.connector.cursor import MySQLCursor as _MysqlConnectorSyncRawCursor +from sqlspec.typing import import_optional_attr + if TYPE_CHECKING: from collections.abc import Awaitable, Callable from types import TracebackType from typing import ClassVar, Protocol, TypeAlias + from mysql.connector.aio.pooling import MySQLConnectionPool as MysqlConnectorAsyncPool + from sqlspec.adapters.mysqlconnector.driver import MysqlConnectorAsyncDriver, MysqlConnectorSyncDriver from sqlspec.core import StatementConfig @@ -59,6 +63,7 @@ class MysqlConnectorMysqlModuleProtocol(Protocol): MysqlConnectorAsyncRawCursor: TypeAlias = _MysqlConnectorAsyncRawCursor if not TYPE_CHECKING: + MysqlConnectorAsyncPool = import_optional_attr("mysql.connector.aio.pooling", "MySQLConnectionPool") MysqlConnectorAio = _mysql_connector_aio MysqlConnectorSyncConnection = _MysqlConnectorSyncConnection MysqlConnectorAsyncConnection = _MysqlConnectorAsyncConnection @@ -73,6 +78,7 @@ class MysqlConnectorMysqlModuleProtocol(Protocol): "MysqlConnectorAio", "MysqlConnectorAsyncConnection", "MysqlConnectorAsyncCursor", + "MysqlConnectorAsyncPool", "MysqlConnectorAsyncRawCursor", "MysqlConnectorAsyncSessionContext", "MysqlConnectorConnectionPool", diff --git a/sqlspec/adapters/mysqlconnector/config.py b/sqlspec/adapters/mysqlconnector/config.py index 8a64d48bf..58a0f8d85 100644 --- a/sqlspec/adapters/mysqlconnector/config.py +++ b/sqlspec/adapters/mysqlconnector/config.py @@ -11,6 +11,7 @@ MysqlConnectorAio, MysqlConnectorAsyncConnection, MysqlConnectorAsyncCursor, + MysqlConnectorAsyncPool, MysqlConnectorAsyncSessionContext, MysqlConnectorConnectionPool, MysqlConnectorMysqlModule, @@ -25,7 +26,7 @@ MysqlConnectorSyncDriver, MysqlConnectorSyncExceptionHandler, ) -from sqlspec.config import ExtensionConfigs, NoPoolAsyncConfig, SyncDatabaseConfig +from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs, SyncDatabaseConfig from sqlspec.core import TypeCoercionCapabilities from sqlspec.driver import ( AsyncPoolConnectionContext, @@ -172,6 +173,10 @@ class MysqlConnectorPoolParams(MysqlConnectorSyncConnectionParams): class MysqlConnectorAsyncConnectionParams(_MysqlConnectorBaseConnectionParams): """MysqlConnector async connection parameters.""" + pool_name: NotRequired[str] + pool_size: NotRequired[int] + pool_reset_session: NotRequired[bool] + class MysqlConnectorDriverFeatures(TypedDict): """MysqlConnector driver feature flags. @@ -183,7 +188,8 @@ class MysqlConnectorDriverFeatures(TypedDict): on_connection_create: Callback executed when a connection is acquired. For sync: Callable[[MysqlConnectorSyncConnection], None] For async: Callable[[MysqlConnectorAsyncConnection], Awaitable[None]] - Called exactly once per physical connection using WeakSet tracking. + Called once per physical connection for sync connections. Async pooled + connections invoke the callback again after native session reset. enable_events: Enable database event channel support. Defaults to True when extension_config["events"] is configured. events_backend: Event channel backend selection. @@ -275,7 +281,7 @@ class MysqlConnectorAsyncConnectionContext(AsyncPoolConnectionContext): __slots__ = () async def __aenter__(self) -> MysqlConnectorAsyncConnection: - self._connection = await self._config.create_connection() + self._connection = await self._config._acquire_async_connection() return cast("MysqlConnectorAsyncConnection", self._connection) async def __aexit__( @@ -291,7 +297,7 @@ class _MysqlConnectorAsyncSessionConnectionHandler(AsyncPoolSessionFactory): __slots__ = () async def acquire_connection(self) -> MysqlConnectorAsyncConnection: - self._connection = await self._config.create_connection() + self._connection = await self._config._acquire_async_connection() return cast("MysqlConnectorAsyncConnection", self._connection) async def release_connection(self, _conn: MysqlConnectorAsyncConnection, **kwargs: Any) -> None: @@ -370,7 +376,7 @@ def __init__( ) def _ensure_connection(self, connection: "MysqlConnectorSyncConnection") -> None: - """Ensure connection callback has been called exactly once for this connection.""" + """Initialize connection state after creation or a native session reset.""" if self._user_connection_hook is None: return underlying = getattr(connection, "_cnx", None) or connection @@ -433,7 +439,9 @@ def get_event_runtime_hints(self) -> "EventRuntimeHints": return EventRuntimeHints(poll_interval=0.25, lease_seconds=5, select_for_update=True, skip_locked=True) -class MysqlConnectorAsyncConfig(NoPoolAsyncConfig[MysqlConnectorAsyncConnection, MysqlConnectorAsyncDriver]): +class MysqlConnectorAsyncConfig( + AsyncDatabaseConfig[MysqlConnectorAsyncConnection, "MysqlConnectorAsyncPool", MysqlConnectorAsyncDriver] +): """Configuration for mysql-connector async MySQL connections.""" driver_type: ClassVar[type[MysqlConnectorAsyncDriver]] = MysqlConnectorAsyncDriver @@ -483,6 +491,8 @@ def __init__( msg = "enable_local_infile_bulk_load requires local_infile=True or allow_local_infile=True in connection_config." raise ImproperConfigurationError(msg) + self._initialized_connections: WeakSet[Any] = WeakSet() + super().__init__( connection_config=self.connection_config, connection_instance=connection_instance, @@ -495,9 +505,65 @@ def __init__( **kwargs, ) + async def _create_pool(self) -> "MysqlConnectorAsyncPool": + if MysqlConnectorAsyncPool is None: + msg = "Async connection pooling requires mysql-connector-python 9.4 or later" + raise ImproperConfigurationError(msg) + config = dict(self.connection_config) + pool_name = config.pop("pool_name", None) + pool_size = config.pop("pool_size", 5) + pool_reset = config.pop("pool_reset_session", True) + pool = MysqlConnectorAsyncPool( + pool_name=pool_name, pool_size=pool_size, pool_reset_session=pool_reset, **config + ) + try: + await pool.initialize_pool() + except BaseException: + with contextlib.suppress(Exception): + await pool.close_pool() + raise + return pool + + async def _close_pool(self) -> None: + if self.connection_instance is not None: + await self.connection_instance.close_pool() + self.connection_instance = None + + async def _ensure_connection(self, connection: "MysqlConnectorAsyncConnection") -> None: + """Initialize connection state after creation or a native session reset.""" + if self._user_connection_hook is None: + return + underlying = getattr(connection, "_cnx", None) + if underlying is not None and self.connection_config.get("pool_reset_session", True): + await self._user_connection_hook(connection) + return + underlying = underlying or connection + if underlying not in self._initialized_connections: + await self._user_connection_hook(connection) + self._initialized_connections.add(underlying) + + async def _acquire_async_connection(self) -> MysqlConnectorAsyncConnection: + """Acquire and initialize an async mysql-connector connection from pool.""" + if ( + MysqlConnectorAsyncPool is None + and self.connection_instance is None + and not any(key in self.connection_config for key in _POOL_ONLY_CONFIG_KEYS) + ): + return await self.create_connection() + pool = await self.provide_pool() + connection = cast("MysqlConnectorAsyncConnection", await pool.get_connection()) + try: + await self._ensure_connection(connection) + except BaseException: + with contextlib.suppress(Exception): + await connection.close() + raise + return connection + async def create_connection(self) -> MysqlConnectorAsyncConnection: """Open and initialize a standalone connection owned by the caller.""" - connection = await mysqlconnector_aio.connect(**self.connection_config) + config = {key: value for key, value in self.connection_config.items() if key not in _POOL_ONLY_CONFIG_KEYS} + connection = await mysqlconnector_aio.connect(**config) try: autocommit = self.connection_config.get("autocommit") if autocommit is not None: diff --git a/tests/unit/adapters/test_mysqlconnector/test_config.py b/tests/unit/adapters/test_mysqlconnector/test_config.py index 815572cac..78f211844 100644 --- a/tests/unit/adapters/test_mysqlconnector/test_config.py +++ b/tests/unit/adapters/test_mysqlconnector/test_config.py @@ -411,3 +411,45 @@ async def test_async_connection_closes_when_initialization_fails( await config.create_connection() connection.close.assert_awaited_once() + + +@pytest.mark.parametrize("reset_session, calls", [(True, 2), (False, 1)]) +async def test_async_pool_reinitializes_reset_sessions(reset_session: bool, calls: int) -> None: + physical = MagicMock() + connection = MagicMock(_cnx=physical) + connection.close = AsyncMock() + pool = MagicMock(get_connection=AsyncMock(return_value=connection)) + hook = AsyncMock() + config = MysqlConnectorAsyncConfig( + connection_config={"pool_reset_session": reset_session}, + connection_instance=pool, + driver_features={"on_connection_create": hook}, + ) + assert await config._acquire_async_connection() is connection + assert await config._acquire_async_connection() is connection + assert hook.await_count == calls + + +@pytest.mark.parametrize("failure", [RuntimeError("hook failed"), asyncio.CancelledError()]) +async def test_async_pool_returns_connection_after_hook_failure(failure: BaseException) -> None: + connection = MagicMock() + connection.close = AsyncMock() + config = MysqlConnectorAsyncConfig( + connection_instance=MagicMock(get_connection=AsyncMock(return_value=connection)), + driver_features={"on_connection_create": AsyncMock(side_effect=failure)}, + ) + with pytest.raises(type(failure)): + await config._acquire_async_connection() + connection.close.assert_awaited_once() + + +async def test_async_pool_unavailable_preserves_standalone_connections(monkeypatch: pytest.MonkeyPatch) -> None: + from sqlspec.adapters.mysqlconnector import config as config_module + from sqlspec.exceptions import ImproperConfigurationError + + monkeypatch.setattr(config_module, "MysqlConnectorAsyncPool", None) + connection = MagicMock(close=AsyncMock()) + monkeypatch.setattr(config_module.mysqlconnector_aio, "connect", AsyncMock(return_value=connection)) + assert await MysqlConnectorAsyncConfig()._acquire_async_connection() is connection + with pytest.raises(ImproperConfigurationError, match=r"9\.4"): + await MysqlConnectorAsyncConfig(connection_config={"pool_size": 2})._acquire_async_connection() From 5f4e14d9b37362e8fb595e7d7a3e8572bcbe9053 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:31:33 +0000 Subject: [PATCH 04/10] docs: clarify Oracle modes and native MySQL async pooling --- docs/changelog.rst | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 27bdaf2ce..66793b029 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -57,10 +57,11 @@ Unreleased Names with spaces and mixed-case Oracle names retain their quotes. * MySQL pools release connections when setup fails. They discard connections - that fail to roll back. MySQL Connector's async config keeps its direct - connection path. + that fail to roll back. MySQL Connector keeps native async pooling on + Connector 9.4 and later, plus direct connections on older versions. -* Oracle keeps the user's Thin/Thick mode choice. Pool shutdown waits for +* Oracle keeps Thick-mode options for sync pools. Async pools reject Thick + mode before they open. Pool shutdown waits for borrowed connections. Custom handlers still convert LOBs, and JSON handlers preserve the user's callbacks. From 94c1ea6b33403b06c1df7711006a78c4c5c7c850 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:39:36 +0000 Subject: [PATCH 05/10] fix(asyncmy): preserve native local infile loader binding --- sqlspec/adapters/asyncmy/_typing.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/sqlspec/adapters/asyncmy/_typing.py b/sqlspec/adapters/asyncmy/_typing.py index 48912c24d..6f67a481e 100644 --- a/sqlspec/adapters/asyncmy/_typing.py +++ b/sqlspec/adapters/asyncmy/_typing.py @@ -11,7 +11,7 @@ import asyncmy as _asyncmy # pyright: ignore from asyncmy import Connection # pyright: ignore from asyncmy import errors as _asyncmy_errors # pyright: ignore -from asyncmy.connection import LoadLocalFile as _LoadLocalFile # pyright: ignore +from asyncmy.connection import LoadLocalFile # pyright: ignore from asyncmy.connection import MySQLResult as _AsyncmyResult # pyright: ignore from asyncmy.constants import FIELD_TYPE as _ASYNCMY_FIELD_TYPE # pyright: ignore from asyncmy.cursors import RE_INSERT_VALUES as ASYNCMY_INSERT_VALUES_PATTERN @@ -224,7 +224,7 @@ async def _read_load_local_packet(self, first_packet: Any) -> None: if not self.connection._local_infile or os.fsdecode(request) != self._filename: msg = "MySQL requested an unexpected LOCAL INFILE payload." raise SQLSpecError(msg) - await _LoadLocalFile(self._filename, self.connection).send_data() # type: ignore[no-untyped-call] + await LoadLocalFile(self._filename, self.connection).send_data() # type: ignore[no-untyped-call] packet = await self.connection.read_packet() if not packet.is_ok_packet(): msg = "MySQL did not acknowledge the LOCAL INFILE payload." From 5810235b38b4ae9c398d6a9de8e2d25bb6111135 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:05:47 +0000 Subject: [PATCH 06/10] docs: reconcile unreleased adapter changelog after rebase --- docs/changelog.rst | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 66793b029..c09405f46 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -59,9 +59,10 @@ Unreleased * MySQL pools release connections when setup fails. They discard connections that fail to roll back. MySQL Connector keeps native async pooling on Connector 9.4 and later, plus direct connections on older versions. + Asyncmy retains native ``LOAD DATA LOCAL INFILE`` support. * Oracle keeps Thick-mode options for sync pools. Async pools reject Thick - mode before they open. Pool shutdown waits for + mode before they open. Pool shutdown preserves native checks for borrowed connections. Custom handlers still convert LOBs, and JSON handlers preserve the user's callbacks. From 65748b82dc1c5c597c0d24046b2e1d15552bcb26 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:28:27 +0000 Subject: [PATCH 07/10] perf(sqlglot): avoid redundant copies with isolated expression ownership --- docs/changelog.rst | 4 +++ sqlspec/adapters/adbc/core.py | 2 +- sqlspec/adapters/bigquery/core.py | 2 +- sqlspec/adapters/duckdb/core.py | 6 ++-- sqlspec/adapters/psqlpy/core.py | 2 +- sqlspec/builder/_base.py | 16 +++++----- sqlspec/builder/_dml.py | 13 ++++++-- sqlspec/builder/_factory.py | 2 +- sqlspec/builder/_select.py | 2 +- sqlspec/core/query_modifiers.py | 4 +-- sqlspec/core/statement.py | 17 +++++----- sqlspec/dialects/db2/_generators.py | 13 +++++--- sqlspec/dialects/db2/_transforms.py | 15 +++++++-- sqlspec/driver/_common.py | 14 ++++---- sqlspec/migrations/base.py | 3 +- sqlspec/migrations/schema.py | 2 +- sqlspec/utils/fixtures.py | 4 ++- .../builder/test_cte_parameter_collisions.py | 32 +++++++++++++++++++ tests/unit/builder/test_parsing_utils.py | 26 ++++++++++++++- tests/unit/core/test_query_modifiers.py | 20 ++++++++++++ tests/unit/core/test_statement.py | 23 ++++++++++++- tests/unit/dialects/test_db2.py | 22 +++++++++++++ 22 files changed, 196 insertions(+), 48 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index c09405f46..ef7bc9300 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -45,6 +45,10 @@ Unreleased **Fixed:** +* Builder results keep CTE trees independent, and column pruning no longer + exposes its cached expression to mutation. SQL generation avoids redundant + copies of temporary trees while preserving caller and cache ownership. + * Asyncpg stack telemetry reports sequential prepared execution rather than native pipelining. Each statement still returns its own result. diff --git a/sqlspec/adapters/adbc/core.py b/sqlspec/adapters/adbc/core.py index eca51f939..713d9f29b 100644 --- a/sqlspec/adapters/adbc/core.py +++ b/sqlspec/adapters/adbc/core.py @@ -1162,7 +1162,7 @@ def wrap_parameter(node: exp.Expression) -> exp.Expression: for expression in expressions: expression.transform(wrap_parameter, copy=False) - rewritten = "; ".join(expression.sql(dialect=dialect) for expression in expressions) + rewritten = "; ".join(expression.sql(dialect=dialect, copy=False) for expression in expressions) return rewritten, effective_scalars, effective_arrays diff --git a/sqlspec/adapters/bigquery/core.py b/sqlspec/adapters/bigquery/core.py index b70da5da9..1ea8205f1 100644 --- a/sqlspec/adapters/bigquery/core.py +++ b/sqlspec/adapters/bigquery/core.py @@ -924,7 +924,7 @@ def _build_multi_row_insert_script( if not isinstance(statement_values, exp.Values): return None statement_values.set("expressions", chunk) - statements.append(str(statement.sql(dialect="bigquery"))) + statements.append(str(statement.sql(dialect="bigquery", copy=False))) return ";\n".join(statements) diff --git a/sqlspec/adapters/duckdb/core.py b/sqlspec/adapters/duckdb/core.py index a4bfc9858..abef4c75a 100644 --- a/sqlspec/adapters/duckdb/core.py +++ b/sqlspec/adapters/duckdb/core.py @@ -335,7 +335,7 @@ def _build_storage_copy_sql(sql: str, file_format: str) -> str | None: if file_format == "csv": options.append(exp.CopyParameter(this=exp.Var(this="HEADER"), expression=exp.Boolean(this=True))) return exp.Copy(this=exp.Subquery(this=statements[0]), kind=False, files=[exp.Placeholder()], params=options).sql( - dialect="duckdb" + dialect="duckdb", copy=False ) @@ -366,7 +366,9 @@ def _build_storage_read_sql(table: str, uri: str, file_format: str) -> str | Non exp.Kwarg(this=exp.Var(this="hive_partitioning"), expression=exp.Boolean(this=False)), ], ) - return exp.Insert(this=target, expression=exp.select("*").from_(reader)).sql(dialect="duckdb") + return exp.Insert(this=target, expression=exp.select("*").from_(reader, copy=False)).sql( + dialect="duckdb", copy=False + ) def _resolve_native_storage_target( diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index 9f764ef47..602482c92 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -680,7 +680,7 @@ def _dml_count_query(sql: str) -> str | None: count_expression = exp.alias_(exp.Count(this=exp.Star()), _DML_COUNT_COLUMN) count_query = exp.select(count_expression).from_(cte_alias, copy=False) count_query.with_(cte_alias, as_=expression, copy=False) - return count_query.sql(dialect="postgres") + return count_query.sql(dialect="postgres", copy=False) def _create_postgres_error(error: Any, error_class: type[SQLSpecError], description: str) -> SQLSpecError: diff --git a/sqlspec/builder/_base.py b/sqlspec/builder/_base.py index 7fc1bd7ad..278ba1396 100644 --- a/sqlspec/builder/_base.py +++ b/sqlspec/builder/_base.py @@ -259,11 +259,11 @@ def _build_final_expression(self, *, copy: bool = False) -> exp.Expr: final_expression: exp.Expr = base_expression existing_with = final_expression.args.get("with_") if existing_with is None: - final_expression.set("with_", exp.With(expressions=list(self._with_ctes.values()))) + final_expression.set("with_", exp.With(expressions=[cte.copy() for cte in self._with_ctes.values()])) else: for cte_node in self._with_ctes.values(): if cte_node not in existing_with.expressions: - existing_with.append("expressions", cte_node) + existing_with.append("expressions", cte_node.copy()) if any(cte.meta.get("recursive") for cte in self._with_ctes.values()): final_expression.args["with_"].set("recursive", True) @@ -292,9 +292,8 @@ def _resolve_cte_query( self._raise_cte_query_error( alias, f"expression must be a Select or Values, got {type(query_expr).__name__}" ) - cte_select_expression: exp.Expr = query_expr.copy() + cte_select_expression: exp.Expr = query_expr if isinstance(cte_select_expression, exp.Values) and cte_select_expression.args.get("alias"): - cte_select_expression = cte_select_expression.copy() cte_select_expression.set("alias", None) param_mapping = self._merge_cte_parameters(alias, query.parameters) if param_mapping: @@ -316,7 +315,6 @@ def _resolve_cte_query( ) cte_duck_expression: exp.Expr = raw_query_expr.copy() if isinstance(cte_duck_expression, exp.Values) and cte_duck_expression.args.get("alias"): - cte_duck_expression = cte_duck_expression.copy() cte_duck_expression.set("alias", None) if hasattr(raw_query, "parameters"): param_mapping = self._merge_cte_parameters(alias, raw_query.parameters) @@ -629,7 +627,7 @@ def build(self, dialect: DialectType = None) -> "BuiltQuery": Returns: BuiltQuery: A dataclass containing the SQL string and parameters. """ - final_expression = self._build_final_expression() + final_expression = self._build_final_expression(copy=True) self._validate_update_from(final_expression, _resolve_dialect(dialect, self.dialect)) if self.enable_optimization and isinstance(final_expression, exp.Expr): @@ -648,7 +646,9 @@ def build(self, dialect: DialectType = None) -> "BuiltQuery": identify = self._should_identify(target_dialect) if normalized_expression.find(exp.Lock) and target_dialect != "db2": register_lock_generator(target_dialect) - sql_string = normalized_expression.sql(dialect=target_dialect, pretty=True, identify=identify) + sql_string = normalized_expression.sql( + dialect=target_dialect, pretty=True, identify=identify, copy=False + ) sql_string = self._strip_merge_target_quotes(sql_string) else: sql_string = str(final_expression) @@ -1187,7 +1187,7 @@ def build_static_expression( target_dialect = str(dialect) if dialect else self.dialect_name identify = self._should_identify(target_dialect) - sql_string = expr.sql(dialect=target_dialect, pretty=True, identify=identify) + sql_string = expr.sql(dialect=target_dialect, pretty=True, identify=identify, copy=not copy) return BuiltQuery( sql=sql_string, parameters=parameters.copy() if parameters else {}, diff --git a/sqlspec/builder/_dml.py b/sqlspec/builder/_dml.py index 1af81575b..6fa95732b 100644 --- a/sqlspec/builder/_dml.py +++ b/sqlspec/builder/_dml.py @@ -392,7 +392,13 @@ def from_(self, table: str | exp.Expr | Any, alias: str | None = None) -> Self: msg = "Subquery builder has no expression to include in FROM clause." raise SQLBuilderError(msg) - subquery_copy = raw_expression.copy() if hasattr(raw_expression, "copy") else raw_expression + subquery_copy = ( + raw_expression + if isinstance(table, QueryBuilder) + else raw_expression.copy() + if hasattr(raw_expression, "copy") + else raw_expression + ) base_builder = cast("QueryBuilder", self) builder_alias = getattr(table, "alias_name", None) or getattr(table, "alias", None) if not isinstance(builder_alias, str): @@ -411,10 +417,11 @@ def from_(self, table: str | exp.Expr | Any, alias: str | None = None) -> Self: if alias: cols: list[str] = [] existing_alias = subquery_copy.args.get("alias") + source_columns = getattr(table, "columns", None) if existing_alias and existing_alias.args.get("columns"): cols = [c.name for c in existing_alias.args["columns"]] - elif hasattr(table, "columns") and isinstance(table.columns, (list, tuple)): - cols = [str(c) for c in table.columns] + elif isinstance(source_columns, (list, tuple)): + cols = [str(c) for c in source_columns] table_expr = exp.alias_(subquery_copy, alias, table=cols or False) else: table_expr = subquery_copy diff --git a/sqlspec/builder/_factory.py b/sqlspec/builder/_factory.py index 66336432d..42fdc3c17 100644 --- a/sqlspec/builder/_factory.py +++ b/sqlspec/builder/_factory.py @@ -118,7 +118,7 @@ def build_copy_statement( expression = _build_copy_expression( direction=direction, table=table, location=location, columns=columns, options=options ) - rendered = expression.sql(dialect=_normalize_copy_dialect(dialect)) + rendered = expression.sql(dialect=_normalize_copy_dialect(dialect), copy=False) return SQL(rendered) diff --git a/sqlspec/builder/_select.py b/sqlspec/builder/_select.py index c09ae3741..6585b4ac2 100644 --- a/sqlspec/builder/_select.py +++ b/sqlspec/builder/_select.py @@ -290,7 +290,7 @@ def from_( msg = "Subquery builder has no expression to include in FROM clause." raise SQLBuilderError(msg) - subquery_copy = subquery_expression.copy() + subquery_copy = subquery_expression if isinstance(table, QueryBuilder) else subquery_expression.copy() base_builder = cast("QueryBuilder", builder) param_mapping = base_builder._merge_cte_parameters(alias or "subquery", table.parameters) if param_mapping: diff --git a/sqlspec/core/query_modifiers.py b/sqlspec/core/query_modifiers.py index 8fbc602e2..57f9b0e55 100644 --- a/sqlspec/core/query_modifiers.py +++ b/sqlspec/core/query_modifiers.py @@ -440,7 +440,7 @@ def wrap_as_subquery(expression: exp.Expr, alias: str = "filtered") -> exp.Selec right.set("order", None) subquery = working if isinstance(working, exp.Subquery) and not bounded else exp.Subquery(this=working) subquery.set("alias", exp.TableAlias(this=exp.to_identifier(alias))) - outer = exp.Select().select("*").from_(subquery) + outer = exp.Select().select("*", copy=False).from_(subquery, copy=False) if with_ is not None: outer.set("with_", with_) if order is not None: @@ -503,7 +503,7 @@ def apply_column_pruning( # Cache the result if cache_key is not None: cache = get_cache() - cache.put_optimized(cache_key, pruned, dialect) + cache.put_optimized(cache_key, pruned.copy(), dialect) if isinstance(pruned, exp.Expr): return pruned diff --git a/sqlspec/core/statement.py b/sqlspec/core/statement.py index 6b0754b60..66f67419e 100644 --- a/sqlspec/core/statement.py +++ b/sqlspec/core/statement.py @@ -1607,25 +1607,24 @@ def builder(self, dialect: "DialectType | None" = None) -> "QueryBuilder": builder: QueryBuilder if isinstance(base_expression, (exp.Select, exp.Union, exp.Except, exp.Intersect, exp.Values)): builder = Select(dialect=builder_dialect) - builder.set_expression(base_expression.copy()) + builder.set_expression(base_expression) elif isinstance(base_expression, exp.Insert): builder = Insert(dialect=builder_dialect) - builder.set_expression(base_expression.copy()) + builder.set_expression(base_expression) elif isinstance(base_expression, exp.Update): builder = Update(dialect=builder_dialect) - builder.set_expression(base_expression.copy()) + builder.set_expression(base_expression) elif isinstance(base_expression, exp.Delete): builder = Delete(dialect=builder_dialect) - builder.set_expression(base_expression.copy()) + builder.set_expression(base_expression) elif isinstance(base_expression, exp.Merge): builder = Merge(dialect=builder_dialect) - builder.set_expression(base_expression.copy()) + builder.set_expression(base_expression) else: - copied = base_expression.copy() - if not isinstance(copied, exp.Expression): - msg = f"Unsupported expression type for builder: {type(copied).__name__}" + if not isinstance(base_expression, exp.Expression): + msg = f"Unsupported expression type for builder: {type(base_expression).__name__}" raise sqlspec.exceptions.SQLBuilderError(msg) - builder = ExpressionBuilder(copied, dialect=builder_dialect) + builder = ExpressionBuilder(base_expression, dialect=builder_dialect) if ctes: builder.load_ctes(ctes) diff --git a/sqlspec/dialects/db2/_generators.py b/sqlspec/dialects/db2/_generators.py index 65fb68515..2467a5e18 100644 --- a/sqlspec/dialects/db2/_generators.py +++ b/sqlspec/dialects/db2/_generators.py @@ -161,9 +161,8 @@ def _convert_limit_to_fetch(query: exp.Expr) -> exp.Expr: """Convert an exp.Limit clause to an exp.Fetch clause for Db2.""" limit = query.args.get("limit") if isinstance(limit, exp.Limit): - query = query.copy() direction = "NEXT" if query.args.get("offset") else "FIRST" - fetch = exp.Fetch(direction=direction, count=exp.maybe_copy(limit.expression)) + fetch = exp.Fetch(direction=direction, count=limit.expression) query.set("limit", fetch) return query @@ -171,14 +170,20 @@ def _convert_limit_to_fetch(query: exp.Expr) -> exp.Expr: def select_sql(generator: "generator.Generator", expression: exp.Select) -> str: """Render a SELECT with a Db2 dummy table, FETCH pagination and statement tail.""" detached, locks, tail_args = _detach_statement_tail(expression) - select = cast("exp.Select", _convert_limit_to_fetch(add_sysibm_dual(cast("exp.Select", detached)))) + if detached is expression and ( + expression.args.get("from_") is None or isinstance(expression.args.get("limit"), exp.Limit) + ): + detached = expression.copy() + select = cast("exp.Select", _convert_limit_to_fetch(add_sysibm_dual(cast("exp.Select", detached), copy=False))) return _with_statement_tail(generator.select_sql(select), render_statement_tail(generator, tail_args, locks)) def set_operation_sql(generator: "generator.Generator", expression: exp.SetOperation) -> str: """Render UNION, INTERSECT or EXCEPT followed by the Db2 statement tail.""" detached, locks, tail_args = _detach_statement_tail(expression) - operation = cast("exp.SetOperation", _convert_limit_to_fetch(cast("exp.SetOperation", detached))) + if detached is expression and isinstance(expression.args.get("limit"), exp.Limit): + detached = expression.copy() + operation = cast("exp.SetOperation", _convert_limit_to_fetch(detached)) return _with_statement_tail(generator.set_operations(operation), render_statement_tail(generator, tail_args, locks)) diff --git a/sqlspec/dialects/db2/_transforms.py b/sqlspec/dialects/db2/_transforms.py index 51dbe73fb..b9a3dc4e6 100644 --- a/sqlspec/dialects/db2/_transforms.py +++ b/sqlspec/dialects/db2/_transforms.py @@ -18,10 +18,19 @@ ) -def add_sysibm_dual(expression: exp.Select) -> exp.Select: - """Add FROM SYSIBM.SYSDUMMY1 to SELECT statements lacking a FROM clause.""" +def add_sysibm_dual(expression: exp.Select, *, copy: bool = True) -> exp.Select: + """Add FROM SYSIBM.SYSDUMMY1 to SELECT statements lacking a FROM clause. + + Args: + expression: SELECT expression to render. + copy: Copy before modification unless the caller owns the expression. + + Returns: + The SELECT expression with a FROM clause. + """ if expression.args.get("from_") is None: - expression = expression.copy() + if copy: + expression = expression.copy() expression.set( "from_", exp.From(this=exp.Table(this=exp.to_identifier("SYSDUMMY1"), db=exp.to_identifier("SYSIBM"))) ) diff --git a/sqlspec/driver/_common.py b/sqlspec/driver/_common.py index f96b23d70..dffae51d0 100644 --- a/sqlspec/driver/_common.py +++ b/sqlspec/driver/_common.py @@ -1708,8 +1708,8 @@ def _count_query(self, original_sql: "SQL") -> "SQL": subquery_expr.set("order", None) subquery_expr.set("limit", None) subquery_expr.set("offset", None) - subquery = subquery_expr.subquery(alias="grouped_data") - count_expr = exp.select(exp.Count(this=exp.Star())).from_(subquery) + subquery = subquery_expr.subquery(alias="grouped_data", copy=False) + count_expr = exp.select(exp.Count(this=exp.Star())).from_(subquery, copy=False) else: # Direct count from source count_expr = exp.select(exp.Count(this=exp.Star())) @@ -1726,7 +1726,7 @@ def _count_query(self, original_sql: "SQL") -> "SQL": if tables: first_table = tables[0] # Create new FROM clause with the found table - count_expr = count_expr.from_(first_table.copy()) + count_expr = count_expr.from_(first_table.copy(), copy=False) # Copy JOIN clauses joins = expr.args.get("joins") @@ -1760,8 +1760,8 @@ def _count_query(self, original_sql: "SQL") -> "SQL": count_source.set("order", None) count_source.set("limit", None) count_source.set("offset", None) - subquery = count_source.subquery(alias="total_query") - count_expr = exp.select(exp.Count(this=exp.Star())).from_(subquery) + subquery = count_source.subquery(alias="total_query", copy=False) + count_expr = exp.select(exp.Count(this=exp.Star())).from_(subquery, copy=False) if cte is not None: count_expr.set("with_", cte.copy()) # Filter out pagination parameters (limit/offset) captured before compile() @@ -1808,8 +1808,8 @@ def _with_total_count(self, original_sql: "SQL", alias: str = "_total_count") -> if cte: expr_copy.set("with_", None) # Wrap set operation in subquery, then select all columns + count window - subquery = expr_copy.subquery(alias="__set_op_subq") - modified_expr = exp.select(exp.Column(this=exp.Star())).from_(subquery) + subquery = expr_copy.subquery(alias="__set_op_subq", copy=False) + modified_expr = exp.select(exp.Column(this=exp.Star())).from_(subquery, copy=False) count_window = exp.Window(this=exp.Count(this=exp.Star())) aliased_count = exp.alias_(count_window, alias) modified_expr = modified_expr.select(aliased_count, copy=False) diff --git a/sqlspec/migrations/base.py b/sqlspec/migrations/base.py index bf42950e8..d836efa17 100644 --- a/sqlspec/migrations/base.py +++ b/sqlspec/migrations/base.py @@ -117,7 +117,8 @@ def __init__(self, version_table_name: str = "ddl_migrations", version_table_sch self.version_table_name = table_identifier.name self.version_table_schema = resolved_schema_identifier.name if resolved_schema_identifier is not None else None self.version_table = self._qualify_version_table( - table_identifier.sql(), resolved_schema_identifier.sql() if resolved_schema_identifier is not None else None + table_identifier.sql(copy=False), + resolved_schema_identifier.sql(copy=False) if resolved_schema_identifier is not None else None, ) self._output_policy = {"use_logger": False, "echo": True, "summary_only": False} diff --git a/sqlspec/migrations/schema.py b/sqlspec/migrations/schema.py index e0e799b94..6d1394959 100644 --- a/sqlspec/migrations/schema.py +++ b/sqlspec/migrations/schema.py @@ -322,7 +322,7 @@ def _add_column_statements( ) if not target_table.args.get("db") and target.schema: target_table.set("db", _parse_ddl_identifier(target.schema, dialect=target.create_table.dialect)) - alter_table_name = target_table.sql(dialect=target.create_table.dialect) + alter_table_name = target_table.sql(dialect=target.create_table.dialect, copy=False) statements: list[tuple[str, AlterTable]] = [] for column_name in sorted(missing_columns): diff --git a/sqlspec/utils/fixtures.py b/sqlspec/utils/fixtures.py index f9baeb15b..acb8ef241 100644 --- a/sqlspec/utils/fixtures.py +++ b/sqlspec/utils/fixtures.py @@ -1182,7 +1182,9 @@ def _table_insert_statement( ) if family != "postgres" or always_identity.isdisjoint(columns): return statement - return insert_expression.sql(dialect=dialect).replace(") VALUES (", ") OVERRIDING SYSTEM VALUE VALUES (", 1) + return insert_expression.sql(dialect=dialect, copy=False).replace( + ") VALUES (", ") OVERRIDING SYSTEM VALUE VALUES (", 1 + ) def _insert_table_rows_sync( diff --git a/tests/unit/builder/test_cte_parameter_collisions.py b/tests/unit/builder/test_cte_parameter_collisions.py index 81b71e3db..2a25bd964 100644 --- a/tests/unit/builder/test_cte_parameter_collisions.py +++ b/tests/unit/builder/test_cte_parameter_collisions.py @@ -4,6 +4,8 @@ handled with unique parameter naming to prevent collisions. """ +from typing import cast + from sqlspec import sql @@ -239,3 +241,33 @@ def test_multiple_cte_levels_parameter_isolation() -> None: assert "login" in param_values assert "summary" in param_values assert "monthly" in param_values + + +def test_final_expression_copy_detaches_nested_ctes() -> None: + from sqlglot import exp + + inner = sql.select("id").from_("items") + owner = sql.select("id").with_cte("source", inner).from_("source") + owner.enable_optimization = False + before = owner.build().sql + detached = owner._build_final_expression(copy=True) + cte = next(detached.find_all(exp.CTE)) + cte.this.set("where", exp.Where(this=exp.false())) + + assert owner.build().sql == before + assert "FALSE" in detached.sql() + + +def test_from_subquery_keeps_nested_cte_owner_isolated() -> None: + from sqlglot import exp + + source = sql.select("id").with_cte("nested", sql.select("id").from_("items")).from_("nested") + source.enable_optimization = False + before = source.build().sql + target = sql.select("*").from_(source, alias="derived") + target.enable_optimization = False + cte = next(cast("exp.Expr", target.get_expression()).find_all(exp.CTE)) + cte.this.set("where", exp.Where(this=exp.false())) + + assert source.build().sql == before + assert "FALSE" in target.build().sql diff --git a/tests/unit/builder/test_parsing_utils.py b/tests/unit/builder/test_parsing_utils.py index 63d55ddd8..0bf3f7434 100644 --- a/tests/unit/builder/test_parsing_utils.py +++ b/tests/unit/builder/test_parsing_utils.py @@ -6,7 +6,7 @@ """ import contextlib -from typing import Any +from typing import Any, cast import pytest from sqlglot import exp @@ -319,3 +319,27 @@ def test_cached_static_expression_respects_copy_flag() -> None: assert "tbl" not in result.sql assert "tbl" not in repeat.sql assert repeat.parameters == {"val": 456} + + +@pytest.mark.parametrize("static", [False, True]) +def test_cross_dialect_render_preserves_owned_and_cached_ast(static: bool) -> None: + expression = exp.select("a", "b").from_("items").distinct("a").order_by("a", "b") + builder = Select(dialect="postgres", enable_optimization=False) + builder.set_expression(expression) + expected = expression.sql(dialect="postgres") + cache_key = "cross-dialect-render-isolation" + cache = get_cache() + cache.delete_expression(cache_key) + try: + if static: + builder.build_static_expression(cache_key=cache_key, expression_factory=lambda: expression, dialect="mysql") + cached = cast("exp.Expr", cache.get_expression(cache_key)) + assert cached.sql(dialect="postgres") == expected + rendered = builder.build_static_expression(cache_key=cache_key, dialect="postgres") + else: + builder.build(dialect="mysql") + rendered = builder.build(dialect="postgres") + assert expression.sql(dialect="postgres") == expected + assert "DISTINCT ON" in rendered.sql + finally: + cache.delete_expression(cache_key) diff --git a/tests/unit/core/test_query_modifiers.py b/tests/unit/core/test_query_modifiers.py index ecd037b1a..bc8cc1678 100644 --- a/tests/unit/core/test_query_modifiers.py +++ b/tests/unit/core/test_query_modifiers.py @@ -659,3 +659,23 @@ def test_set_operation_support_apply_offset_rejects_non_select_non_set_operation update_expr = exp.update("users", {"name": exp.Literal.string("test")}) with pytest.raises(SQLSpecError, match="OFFSET only valid for SELECT"): apply_offset(update_expr, 5) + + +@pytest.mark.parametrize("dialect", ["postgres", "sqlite"]) +def test_column_pruning_cache_isolates_returned_expression(dialect: str) -> None: + source = sqlglot.parse_one("SELECT a FROM (SELECT a, b FROM items) AS nested") + original = source.sql() + cache_key = "prune-return-isolation" + cache = get_cache() + cache.delete_optimized(cache_key, dialect) + try: + first = apply_column_pruning(source, dialect=dialect, cache_key=cache_key) + expected = first.sql(dialect=dialect) + first.set("limit", exp.Limit(expression=exp.Literal.number(7))) + second = apply_column_pruning(source, dialect=dialect, cache_key=cache_key) + assert second.sql(dialect=dialect) == expected + second.set("offset", exp.Offset(expression=exp.Literal.number(3))) + assert apply_column_pruning(source, dialect=dialect, cache_key=cache_key).sql(dialect=dialect) == expected + assert source.sql() == original + finally: + cache.delete_optimized(cache_key, dialect) diff --git a/tests/unit/core/test_statement.py b/tests/unit/core/test_statement.py index 848086c31..c0d2ba833 100644 --- a/tests/unit/core/test_statement.py +++ b/tests/unit/core/test_statement.py @@ -18,7 +18,7 @@ import importlib.util import logging import pickle -from typing import Any +from typing import Any, cast from unittest.mock import MagicMock, patch import pytest @@ -1710,3 +1710,24 @@ def test_qmark_escape_survives_unparsed_fallback() -> None: rendered, _ = statement.compile() assert "data ? other_col" in rendered assert "COALESCE" not in rendered + + +@pytest.mark.parametrize("dialect", ["postgres", "sqlite"]) +@pytest.mark.parametrize("compiled", [False, True]) +def test_statement_builder_owns_its_expression(dialect: str, compiled: bool) -> None: + statement = SQL("WITH source AS (SELECT id FROM items) SELECT id FROM source", dialect=dialect) + if compiled: + statement.compile() + builder = statement.builder() + builder.enable_optimization = False + expected = builder.build().sql + expression = cast("exp.Expr", builder.get_expression()) + expression.set("limit", exp.Limit(expression=exp.Literal.number(2))) + nested = next(expression.find_all(exp.CTE)) + nested.this.set("where", exp.Where(this=exp.false())) + + fresh = statement.builder() + fresh.enable_optimization = False + assert fresh.build().sql == expected + assert "FALSE" in builder.build().sql + assert "LIMIT 2" in builder.build().sql diff --git a/tests/unit/dialects/test_db2.py b/tests/unit/dialects/test_db2.py index 047d33e78..361daef62 100644 --- a/tests/unit/dialects/test_db2.py +++ b/tests/unit/dialects/test_db2.py @@ -189,3 +189,25 @@ def test_string_position_preserves_occurrence() -> None: occurrence=exp.Literal.number(2), ) assert expression.sql(dialect="db2") == "LOCATE_IN_STRING('x-x-x', 'x', 2, 2)" + + +def test_db2_direct_generation_preserves_query_and_pagination_ownership() -> None: + from sqlspec.dialects.db2._generators import select_sql, set_operation_sql + from sqlspec.dialects.db2._transforms import add_sysibm_dual + + select = parse_one("SELECT 1 LIMIT 2", read="postgres") + select.set("sqlspec_db2_isolation", "UR") + snapshot = select.copy() + dialect = DB2() + assert select_sql(dialect.generator(), select) == "SELECT 1 FROM SYSIBM.SYSDUMMY1 FETCH FIRST 2 ROWS ONLY WITH UR" + assert select == snapshot + assert select.args["limit"].expression.parent is select.args["limit"] + assert add_sysibm_dual(select) is not select + assert select == snapshot + + union = parse_one("SELECT 1 UNION SELECT 2 LIMIT 3", read="postgres") + union_snapshot = union.copy() + assert set_operation_sql(dialect.generator(), union).endswith("FETCH FIRST 3 ROWS ONLY") + assert union == union_snapshot + assert union.args["limit"].expression.parent is union.args["limit"] + assert union.sql(dialect="db2") == set_operation_sql(dialect.generator(), union) From 1a672aa5e9134bac888a538a555bc241d360b0d4 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:30:48 +0000 Subject: [PATCH 08/10] fix(mssql): preserve schema restoration state after failures --- docs/changelog.rst | 4 +++ docs/usage/migrations.rst | 5 ++- sqlspec/adapters/mssql_python/driver.py | 2 +- sqlspec/adapters/pymssql/driver.py | 2 +- .../test_migration_schema.py | 33 +++++++++++++++++++ 5 files changed, 43 insertions(+), 3 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index ef7bc9300..75206fa39 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -49,6 +49,10 @@ Unreleased exposes its cached expression to mutation. SQL generation avoids redundant copies of temporary trees while preserving caller and cache ownership. +* SQL Server migration drivers retain the previous default schema if restoring + it fails, so cleanup can be retried. The migration guide clarifies that this + setting belongs to the database user rather than one connection. + * Asyncpg stack telemetry reports sequential prepared execution rather than native pipelining. Each statement still returns its own result. diff --git a/docs/usage/migrations.rst b/docs/usage/migrations.rst index 8cf3bba35..241fc8a0b 100644 --- a/docs/usage/migrations.rst +++ b/docs/usage/migrations.rst @@ -278,7 +278,10 @@ raises ``MigrationError`` before any DDL is issued. ``sys.schemas``. The previous default schema is restored after each migration (a failed transactional migration rolls the switch back). SQL Server refuses to alter the default schema of the ``dbo`` database - user (what ``sa`` maps to), so connect with a dedicated login. Because + user (what ``sa`` maps to). Use a dedicated migration login and database + user. ``ALTER USER`` changes the database user's default schema across + connections; it is not a session-local setting. Do not run migrations + alongside other work or migrations that share that database user. Because the restore is committed, a failed non-transactional migration on a connection with autocommit disabled also commits the statements that succeeded before the failure. diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index 2a53da99c..5eddec33a 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -282,10 +282,10 @@ def reset_migration_session_schema(self) -> None: if self._migration_schema_restore is None: return user_name, previous_schema = self._migration_schema_restore - self._migration_schema_restore = None with self.with_cursor(self.connection) as cursor: _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None) self.connection.commit() + self._migration_schema_restore = None def has_schema(self, schema: str) -> bool: """Return whether the specified schema exists.""" diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 10739027f..0655b8052 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -269,10 +269,10 @@ def reset_migration_session_schema(self) -> None: if self._migration_schema_restore is None: return user_name, previous_schema = self._migration_schema_restore - self._migration_schema_restore = None with self.with_cursor(self.connection) as cursor: cursor.execute(_alter_default_schema_sql(user_name, previous_schema)) self.connection.commit() + self._migration_schema_restore = None def has_schema(self, schema: str) -> bool: """Return whether the specified schema exists.""" diff --git a/tests/unit/adapters/test_mssql_python/test_migration_schema.py b/tests/unit/adapters/test_mssql_python/test_migration_schema.py index 115afdcd2..3168a1795 100644 --- a/tests/unit/adapters/test_mssql_python/test_migration_schema.py +++ b/tests/unit/adapters/test_mssql_python/test_migration_schema.py @@ -1,9 +1,13 @@ """Unit coverage for mssql-python migration schema hooks.""" from typing import Any, cast +from unittest.mock import Mock + +import pytest from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver +from sqlspec.adapters.pymssql.driver import PymssqlDriver class FakeCursor: @@ -103,3 +107,32 @@ def test_mssql_python_reset_without_set_is_noop() -> None: driver.reset_migration_session_schema() assert cursor.executed == [] + + +@pytest.mark.parametrize("driver_type", [MssqlPythonDriver, PymssqlDriver]) +@pytest.mark.parametrize("failure_point", ["execute", "commit"]) +def test_schema_restore_retries_original_schema_after_failure( + driver_type: "type[MssqlPythonDriver | PymssqlDriver]", failure_point: str, monkeypatch: pytest.MonkeyPatch +) -> None: + cursor = FakeCursor(current_schema="sales") + connection = FakeConnection(cursor) + driver = driver_type(cast("Any", connection)) + driver.set_migration_session_schema("tenant") + failure = RuntimeError("restore failed") + target = cursor if failure_point == "execute" else connection + original = getattr(target, failure_point) + failing = Mock(side_effect=failure) + monkeypatch.setattr(target, failure_point, failing) + with pytest.raises(RuntimeError, match="restore failed") as caught: + driver.reset_migration_session_schema() + assert caught.value is failure + monkeypatch.setattr(target, failure_point, original) + cursor.user_name = "different_user" + cursor.current_schema = "tenant" + driver.reset_migration_session_schema() + assert cursor.executed[-1] == ("ALTER USER [sqlspec_migrator] WITH DEFAULT_SCHEMA = [sales];", None) + assert connection.commits == 1 + completed = list(cursor.executed) + driver.reset_migration_session_schema() + assert cursor.executed == completed + assert connection.commits == 1 From 5efa1f4be6bbf1959083696980e10c6f6b9e6327 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:35:09 +0000 Subject: [PATCH 09/10] fix(spanner): refresh row plans after deserializer changes --- docs/changelog.rst | 3 ++- sqlspec/adapters/spanner/driver.py | 6 +++++- tests/unit/adapters/test_spanner/test_core.py | 20 +++++++++++++++++++ 3 files changed, 27 insertions(+), 2 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 75206fa39..b3a3a64ac 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -80,7 +80,8 @@ Unreleased * Spanner schema queries no longer require a table name. SQL output keeps JOIN hints and plain comments. Sequence statements keep qualified names and - ``IF NOT EXISTS`` guards. + ``IF NOT EXISTS`` guards. Cached row converters refresh when the configured + JSON deserializer changes. * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index 2867e5e07..322426006 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -93,7 +93,7 @@ class SpannerSyncDriver(SyncDriverAdapterBase): """Synchronous Spanner driver operating on Snapshot or Transaction contexts.""" dialect: "DialectType" = "spanner" - __slots__ = ("_data_dictionary", "_pending_execute_options", "_row_plan_cache") + __slots__ = ("_data_dictionary", "_pending_execute_options", "_row_plan_cache", "_row_plan_deserializer") def __init__( self, @@ -109,6 +109,7 @@ def __init__( self._data_dictionary: SpannerDataDictionary | None = None self._pending_execute_options: _PerCallExecuteOptions | None = None self._row_plan_cache: dict[int, tuple[Any, list[str], tuple[tuple[int, Any], ...] | None]] = {} + self._row_plan_deserializer = cast("Callable[[str], Any]", features.get("json_deserializer", from_json)) def dispatch_execute(self, cursor: "SpannerConnection", statement: "SQL") -> ExecutionResult: sql, params = self._compiled_sql(statement, self.statement_config) @@ -599,6 +600,9 @@ def _infer_param_types(self, params: "dict[str, Any] | list[Any] | tuple[Any, .. def _resolve_row_plan(self, fields: Any) -> "tuple[list[str], tuple[tuple[int, Any], ...] | None]": json_deserializer = cast("Callable[[str], Any]", self.driver_features.get("json_deserializer", from_json)) + if json_deserializer is not self._row_plan_deserializer: + self._row_plan_cache.clear() + self._row_plan_deserializer = json_deserializer return resolve_row_plan(fields, self._row_plan_cache, json_deserializer=json_deserializer) diff --git a/tests/unit/adapters/test_spanner/test_core.py b/tests/unit/adapters/test_spanner/test_core.py index f55cae3de..0ba44654c 100644 --- a/tests/unit/adapters/test_spanner/test_core.py +++ b/tests/unit/adapters/test_spanner/test_core.py @@ -12,6 +12,7 @@ resolve_column_names, resolve_row_plan, ) +from sqlspec.adapters.spanner.driver import SpannerSyncDriver from sqlspec.core import TypedParameter @@ -77,6 +78,25 @@ def test_resolve_row_plan_caches_json_metadata_and_plan() -> None: assert first_column_plan[0][0] == 1 +def test_driver_row_plan_uses_changed_json_deserializer_with_reused_metadata() -> None: + fields = [_field("payload", TypeCode.JSON)] + rows = [('{"value":1}',)] + driver = SpannerSyncDriver( + cast("Any", object()), driver_features={"json_deserializer": lambda value: {"first": value}} + ) + + names, first_plan = driver._resolve_row_plan(fields) + first_rows, _ = collect_rows(rows, fields, column_names=names, column_plan=first_plan) + assert first_rows == [({"first": '{"value":1}'},)] + + driver.driver_features["json_deserializer"] = lambda value: {"second": value} + names, second_plan = driver._resolve_row_plan(fields) + second_rows, _ = collect_rows(rows, fields, column_names=names, column_plan=second_plan) + assert second_rows == [({"second": '{"value":1}'},)] + assert second_plan is not first_plan + assert driver._resolve_row_plan(fields)[1] is second_plan + + def test_collect_rows_returns_original_rows_without_a_plan() -> None: fields = [_field("id", TypeCode.INT64), _field("payload", TypeCode.STRING)] rows = [(1, '{"kind":"string"}')] From 825edc710637669223abd3136f32bd79ca2feded Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 23:34:47 +0000 Subject: [PATCH 10/10] test(db2): narrow parsed query types for ownership checks --- tests/unit/dialects/test_db2.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/dialects/test_db2.py b/tests/unit/dialects/test_db2.py index 361daef62..50847f978 100644 --- a/tests/unit/dialects/test_db2.py +++ b/tests/unit/dialects/test_db2.py @@ -196,6 +196,7 @@ def test_db2_direct_generation_preserves_query_and_pagination_ownership() -> Non from sqlspec.dialects.db2._transforms import add_sysibm_dual select = parse_one("SELECT 1 LIMIT 2", read="postgres") + assert isinstance(select, exp.Select) select.set("sqlspec_db2_isolation", "UR") snapshot = select.copy() dialect = DB2() @@ -206,6 +207,7 @@ def test_db2_direct_generation_preserves_query_and_pagination_ownership() -> Non assert select == snapshot union = parse_one("SELECT 1 UNION SELECT 2 LIMIT 3", read="postgres") + assert isinstance(union, exp.SetOperation) union_snapshot = union.copy() assert set_operation_sql(dialect.generator(), union).endswith("FETCH FIRST 3 ROWS ONLY") assert union == union_snapshot