diff --git a/docs/changelog.rst b/docs/changelog.rst index d8555cc36..b3a3a64ac 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -45,13 +45,43 @@ 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. +* 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. + +* 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. -* 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 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 preserves native checks 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. 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. @@ -60,6 +90,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/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/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/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/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..6f67a481e 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 # 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/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/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/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/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/mysqlconnector/_typing.py b/sqlspec/adapters/mysqlconnector/_typing.py index 066648399..1ad16a74d 100644 --- a/sqlspec/adapters/mysqlconnector/_typing.py +++ b/sqlspec/adapters/mysqlconnector/_typing.py @@ -7,22 +7,25 @@ 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 + +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 @@ -44,25 +47,34 @@ 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 + MysqlConnectorAsyncPool = import_optional_attr("mysql.connector.aio.pooling", "MySQLConnectionPool") + 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", diff --git a/sqlspec/adapters/mysqlconnector/config.py b/sqlspec/adapters/mysqlconnector/config.py index 40f594df0..58a0f8d85 100644 --- a/sqlspec/adapters/mysqlconnector/config.py +++ b/sqlspec/adapters/mysqlconnector/config.py @@ -188,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. @@ -375,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 @@ -467,7 +468,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 +477,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 +485,16 @@ 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) + self._initialized_connections: WeakSet[Any] = WeakSet() + super().__init__( - connection_config=connection_config, + connection_config=self.connection_config, connection_instance=connection_instance, migration_config=migration_config, statement_config=statement_config, @@ -504,46 +506,74 @@ def __init__( ) 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", 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 or 5, pool_reset_session=pool_reset, **config + pool_name=pool_name, pool_size=pool_size, pool_reset_session=pool_reset, **config ) - await pool.initialize_pool() + 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: - with contextlib.suppress(Exception): - await self.connection_instance.close_pool() + 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.""" + """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 + 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()) - await self._ensure_connection(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 a standalone connection owned by the caller.""" + """Open and initialize 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) + 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..6a1ef2fda 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, @@ -47,16 +42,15 @@ SyncPoolConnectionContext, SyncPoolSessionFactory, ) +from sqlspec.exceptions import ImproperConfigurationError 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", @@ -141,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] @@ -286,20 +282,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.""" @@ -428,17 +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) - 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." - ) + 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 @@ -488,10 +468,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() @@ -634,17 +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) - 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." - ) + 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 @@ -695,9 +666,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/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/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/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/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/sqlspec/builder/_base.py b/sqlspec/builder/_base.py index 8be712735..278ba1396 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__) @@ -258,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) @@ -291,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: @@ -315,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) @@ -628,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): @@ -647,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) @@ -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. @@ -1186,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 {}, @@ -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..6fa95732b 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 @@ -388,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): @@ -407,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/_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/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/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..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)) @@ -259,7 +264,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..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"))) ) @@ -131,9 +140,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/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 a97cc84b3..d836efa17 100644 --- a/sqlspec/migrations/base.py +++ b/sqlspec/migrations/base.py @@ -114,12 +114,11 @@ 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 + 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} @@ -175,7 +174,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/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/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..acb8ef241 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": @@ -1195,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( @@ -1261,7 +1250,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_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 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..78f211844 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,63 @@ 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() + + +@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() diff --git a/tests/unit/adapters/test_oracledb/test_config.py b/tests/unit/adapters/test_oracledb/test_config.py index 6db0855e6..f98acde0c 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, call import pytest from oracledb import AuthMode, PoolGetMode, Purity @@ -17,6 +19,7 @@ OraclePoolParams, OracleSyncConfig, ) +from sqlspec.exceptions import ImproperConfigurationError class _StubConnection: @@ -249,3 +252,91 @@ 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("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: + 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) + 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={"thick_mode": False}) + result = config._create_pool() + result = await result if isawaitable(result) else result + + 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 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_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"}')] 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_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_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/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 8162e92a8..50847f978 100644 --- a/tests/unit/dialects/test_db2.py +++ b/tests/unit/dialects/test_db2.py @@ -174,3 +174,42 @@ 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)" + + +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") + assert isinstance(select, exp.Select) + 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") + 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 + assert union.args["limit"].expression.parent is union.args["limit"] + assert union.sql(dialect="db2") == set_operation_sql(dialect.generator(), union) 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}]