From d7fc668b211328c68b9f97ec9a5d646b15919357 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Thu, 24 Sep 2026 23:48:11 +0000 Subject: [PATCH 01/11] feat(adapters): optimize MSSQL family adapters --- sqlspec/adapters/mssql_python/_typing.py | 36 +-- sqlspec/adapters/mssql_python/adk/store.py | 220 ++++++++--------- sqlspec/adapters/mssql_python/config.py | 96 ++++---- sqlspec/adapters/mssql_python/core.py | 95 +++++--- .../adapters/mssql_python/data_dictionary.py | 14 +- sqlspec/adapters/mssql_python/driver.py | 212 ++++++++++------ .../adapters/mssql_python/litestar/store.py | 118 +++------ sqlspec/adapters/mssql_python/pool.py | 31 ++- .../adapters/mssql_python/type_converter.py | 26 +- sqlspec/adapters/pymssql/_typing.py | 39 +-- sqlspec/adapters/pymssql/adk/store.py | 226 ++++++++--------- sqlspec/adapters/pymssql/config.py | 72 +++--- sqlspec/adapters/pymssql/core.py | 121 ++++++---- sqlspec/adapters/pymssql/data_dictionary.py | 54 +++-- sqlspec/adapters/pymssql/driver.py | 227 ++++++++++++++++-- sqlspec/adapters/pymssql/events/store.py | 7 +- sqlspec/adapters/pymssql/litestar/store.py | 118 +++------ .../adapters/test_mssql_python/test_config.py | 26 ++ .../adapters/test_mssql_python/test_core.py | 23 +- .../test_mssql_python/test_data_dictionary.py | 25 ++ .../test_mssql_python/test_load_from_arrow.py | 45 +++- .../test_mssql_python/test_type_converter.py | 6 + tests/unit/adapters/test_pymssql/_fakes.py | 21 ++ .../unit/adapters/test_pymssql/test_config.py | 2 + tests/unit/adapters/test_pymssql/test_core.py | 52 ++++ .../test_pymssql/test_data_dictionary.py | 25 ++ .../unit/adapters/test_pymssql/test_driver.py | 124 +++++++++- 27 files changed, 1309 insertions(+), 752 deletions(-) diff --git a/sqlspec/adapters/mssql_python/_typing.py b/sqlspec/adapters/mssql_python/_typing.py index 27f48317c..f35a4a841 100644 --- a/sqlspec/adapters/mssql_python/_typing.py +++ b/sqlspec/adapters/mssql_python/_typing.py @@ -1,23 +1,23 @@ """mssql-python adapter type definitions and mypyc-excluded context managers.""" import contextlib +from collections.abc import Callable +from types import TracebackType from typing import TYPE_CHECKING, Any -import mssql_python as _mssql_python # pyright: ignore[reportMissingImports] +import mssql_python as _mssql_python from mssql_python import Error as MssqlPythonError -from mssql_python.connection import Connection, TokenProvider # pyright: ignore -from mssql_python.cursor import Cursor # pyright: ignore +from mssql_python.connection import Connection, TokenProvider +from mssql_python.cursor import Cursor + +from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver +from sqlspec.core import StatementConfig MSSQL_PYTHON_MODULE: Any = _mssql_python if TYPE_CHECKING: - from collections.abc import Callable - from types import TracebackType from typing import TypeAlias - from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver - from sqlspec.core import StatementConfig - MssqlPythonConnection: TypeAlias = Connection MssqlPythonRawCursor: TypeAlias = Cursor @@ -41,11 +41,11 @@ class MssqlPythonCursor: __slots__ = ("connection", "cursor") - def __init__(self, connection: "MssqlPythonConnection") -> None: + def __init__(self, connection: MssqlPythonConnection) -> None: self.connection = connection self.cursor: MssqlPythonRawCursor | None = None - def __enter__(self) -> "MssqlPythonRawCursor": + def __enter__(self) -> MssqlPythonRawCursor: self.cursor = self.connection.cursor() return self.cursor @@ -70,11 +70,11 @@ class MssqlPythonSessionContext: def __init__( self, - acquire_connection: "Callable[[], MssqlPythonConnection]", - release_connection: "Callable[..., Any]", - statement_config: "StatementConfig", - driver_features: "dict[str, Any]", - prepare_driver: "Callable[[MssqlPythonDriver], MssqlPythonDriver]", + acquire_connection: Callable[[], MssqlPythonConnection], + release_connection: Callable[..., Any], + statement_config: StatementConfig, + driver_features: dict[str, Any], + prepare_driver: Callable[[MssqlPythonDriver], MssqlPythonDriver], ) -> None: self._acquire_connection = acquire_connection self._release_connection = release_connection @@ -84,7 +84,7 @@ def __init__( self._connection: MssqlPythonConnection | None = None self._driver: MssqlPythonDriver | None = None - def __enter__(self) -> "MssqlPythonDriver": + def __enter__(self) -> MssqlPythonDriver: from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver self._connection = self._acquire_connection() @@ -94,8 +94,8 @@ def __enter__(self) -> "MssqlPythonDriver": return self._prepare_driver(self._driver) def __exit__( - self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" - ) -> "bool | None": + self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None + ) -> bool | None: if exc_type is not None and self._driver is not None: with contextlib.suppress(Exception): self._driver.rollback() diff --git a/sqlspec/adapters/mssql_python/adk/store.py b/sqlspec/adapters/mssql_python/adk/store.py index 7540f0a7d..25b85bb6d 100644 --- a/sqlspec/adapters/mssql_python/adk/store.py +++ b/sqlspec/adapters/mssql_python/adk/store.py @@ -1,32 +1,32 @@ """mssql-python ADK stores for Google Agent Development Kit session storage.""" -import re -from datetime import datetime -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast +from collections.abc import Sequence +from datetime import datetime, timedelta +from typing import Any, ClassVar, Final, Literal, cast from typing_extensions import NotRequired -from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor, MssqlPythonError +from sqlspec.adapters.mssql_python._typing import MssqlPythonError +from sqlspec.adapters.mssql_python.config import MssqlPythonConfig +from sqlspec.adapters.mssql_python.core import extract_error_number from sqlspec.config import ADKConfig -from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options -from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore +from sqlspec.extensions.adk import ( + BaseSyncADKMemoryStore, + BaseSyncADKStore, + SessionOrderBy, + StoredEvent, + StoredMemory, + StoredSession, + normalize_session_list_options, +) from sqlspec.utils.serializers import from_json, to_json -if TYPE_CHECKING: - from collections.abc import Sequence - from datetime import timedelta - - from sqlspec.adapters.mssql_python.config import MssqlPythonConfig - from sqlspec.extensions.adk import SessionOrderBy - from sqlspec.extensions.adk.memory._types import StoredMemory - __all__ = ("MssqlPythonADKConfig", "MssqlPythonADKMemoryStore", "MssqlPythonADKStore") MSSQL_TABLE_NOT_FOUND_ERROR: Final[int] = 208 MSSQL_DUPLICATE_OBJECT_ERROR: Final[int] = 2714 MSSQL_DUPLICATE_INDEX_ERROR: Final[int] = 1913 MSSQL_SCHEMA: Final[str] = "dbo" -MSSQL_ERROR_NUMBER_PATTERN: Final[re.Pattern[str]] = re.compile(r"\(([-]?\d+)\)") JSON_FALLBACK_COLUMN_TYPE: Final[str] = "NVARCHAR(MAX)" JSON_NATIVE_COLUMN_TYPE: Final[str] = "JSON" @@ -44,7 +44,7 @@ class MssqlPythonADKStore(BaseSyncADKStore["MssqlPythonConfig"]): connector_name: ClassVar[str] = "mssql_python" __slots__ = ("_json_column_type",) - def __init__(self, config: "MssqlPythonConfig") -> None: + def __init__(self, config: MssqlPythonConfig) -> None: super().__init__(config) adk_config = _adk_config(config) native_json = adk_config.get("native_json") @@ -71,7 +71,7 @@ def create_tables(self) -> None: driver.commit() def create_session( - self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None + self, session_id: str, app_name: str, user_id: str, state: dict[str, Any], owner_id: Any | None = None ) -> StoredSession: """Create a new ADK session.""" owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" @@ -95,8 +95,8 @@ def create_session( return _session_record_from_row(row) def get_session( - self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None - ) -> "StoredSession | None": + self, app_name: str, user_id: str, session_id: str, *, renew_for: int | timedelta | None = None + ) -> StoredSession | None: """Return a scoped session or ``None`` if absent.""" try: if renew_for is not None and self._calculate_expires_at(renew_for) is not None: @@ -123,7 +123,7 @@ def get_session( raise return _session_record_from_row(row) if row is not None else None - def update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None: + def update_session_state(self, app_name: str, user_id: str, session_id: str, state: dict[str, Any]) -> None: """Replace a session's durable state.""" self._execute( f""" @@ -138,13 +138,13 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta def list_sessions( self, app_name: str, - user_id: "str | None" = None, + user_id: str | None = None, *, - order_by: "SessionOrderBy" = "update_time", + order_by: SessionOrderBy = "update_time", descending: bool = True, - limit: "int | None" = None, - offset: "int | None" = None, - ) -> "list[StoredSession]": + limit: int | None = None, + offset: int | None = None, + ) -> list[StoredSession]: """List ADK sessions for an application, optionally scoped to a user.""" column, direction, page_limit, page_offset = normalize_session_list_options(order_by, descending, limit, offset) if page_limit == 0: @@ -179,10 +179,10 @@ def append_event_and_update_state( app_name: str, user_id: str, session_id: str, - state: "dict[str, Any]", + state: dict[str, Any], *, - app_state: "dict[str, Any] | None" = None, - user_state: "dict[str, Any] | None" = None, + app_state: dict[str, Any] | None = None, + user_state: dict[str, Any] | None = None, ) -> StoredSession: """Atomically append an event and update durable session/scoped state.""" update_sql = f""" @@ -191,21 +191,20 @@ def append_event_and_update_state( OUTPUT inserted.id, inserted.app_name, inserted.user_id, inserted.state, inserted.create_time, inserted.update_time WHERE app_name = ? AND user_id = ? AND id = ? """ - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + with self._config.provide_session() as driver: try: - cursor.execute(update_sql, (to_json(state), app_name, user_id, session_id)) - row = cursor.fetchone() + row = driver.select_one_or_none(update_sql, (to_json(state), app_name, user_id, session_id)) if row is None: _raise_session_not_found(session_id) - cursor.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) + driver.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) if app_state is not None: - cursor.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) + driver.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) if user_state is not None: - cursor.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) + driver.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) except Exception: - conn.rollback() + driver.rollback() raise - conn.commit() + driver.commit() return _session_record_from_row(row) def get_events( @@ -213,9 +212,9 @@ def get_events( app_name: str, user_id: str, session_id: str, - after_timestamp: "datetime | None" = None, - limit: "int | None" = None, - ) -> "list[StoredEvent]": + after_timestamp: datetime | None = None, + limit: int | None = None, + ) -> list[StoredEvent]: """Return events for a scoped session ordered by event timestamp.""" if limit == 0: return [] @@ -228,7 +227,7 @@ def get_events( raise return [_event_record_from_row(row) for row in rows] - def delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: + def delete_expired_events(self, before: datetime, app_name: str | None = None) -> int: """Delete events older than ``before``.""" sql = f"DELETE FROM {_table_ref(self._events_table)} WHERE timestamp < ?" params: list[Any] = [before] @@ -242,7 +241,7 @@ def delete_expired_events(self, before: datetime, app_name: "str | None" = None) return 0 raise - def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: + def delete_idle_sessions(self, updated_before: datetime, app_name: str | None = None) -> int: """Delete sessions whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._session_table)} WHERE update_time < ?" params: list[Any] = [updated_before] @@ -256,7 +255,7 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" return 0 raise - def delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: + def delete_idle_user_states(self, updated_before: datetime, app_name: str | None = None) -> int: """Delete user state rows whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._user_state_table)} WHERE update_time < ?" params: list[Any] = [updated_before] @@ -270,7 +269,7 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: "str | Non return 0 raise - def get_app_state(self, app_name: str) -> "dict[str, Any] | None": + def get_app_state(self, app_name: str) -> dict[str, Any] | None: """Return app-scoped state.""" try: row = self._execute_fetchone( @@ -282,7 +281,7 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": raise return _json_dict(row[0]) if row is not None else None - def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None": + def get_user_state(self, app_name: str, user_id: str) -> dict[str, Any] | None: """Return user-scoped state.""" try: row = self._execute_fetchone( @@ -299,15 +298,15 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" raise return _json_dict(row[0]) if row is not None else None - def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: + def upsert_app_state(self, app_name: str, state: dict[str, Any]) -> None: """Insert or replace app-scoped state.""" self._execute(self._upsert_app_state_sql(), (app_name, to_json(state)), commit=True) - def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: + def upsert_user_state(self, app_name: str, user_id: str, state: dict[str, Any]) -> None: """Insert or replace user-scoped state.""" self._execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(state)), commit=True) - def get_metadata(self, key: str) -> "str | None": + def get_metadata(self, key: str) -> str | None: """Return an ADK metadata value.""" try: row = self._execute_fetchone( @@ -323,7 +322,7 @@ def set_metadata(self, key: str, value: str) -> None: """Set an ADK metadata value.""" self._execute(_upsert_metadata_sql(self._metadata_table), (key, value), commit=True) - def _index_specs(self) -> "list[tuple[str, str, str]]": + def _index_specs(self) -> list[tuple[str, str, str]]: """Return ``(index_name, table, columns)`` specs for session and event indexes.""" return [*_sessions_index_specs(self._session_table), *_events_index_specs(self._events_table)] @@ -356,7 +355,7 @@ def _drop_user_states_table_sql(self) -> str: def _drop_metadata_table_sql(self) -> str: return f"DROP TABLE IF EXISTS {_table_ref(self._metadata_table)}" - def _drop_tables_sql(self) -> "list[str]": + def _drop_tables_sql(self) -> list[str]: return [ self._drop_metadata_table_sql(), self._drop_user_states_table_sql(), @@ -376,33 +375,31 @@ def _events_query( app_name: str, user_id: str, session_id: str, - after_timestamp: "datetime | None" = None, - limit: "int | None" = None, - ) -> "tuple[str, tuple[Any, ...]]": + after_timestamp: datetime | None = None, + limit: int | None = None, + ) -> tuple[str, tuple[Any, ...]]: return _events_query(self._events_table, app_name, user_id, session_id, after_timestamp, limit) def _json_column_type_sync(self) -> str: return self._json_column_type - def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None": - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: - cursor.execute(sql, params) - row = cursor.fetchone() + def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> Any | None: + with self._config.provide_session() as driver: + row = driver.select_one_or_none(sql, params) if commit: - conn.commit() + driver.commit() return row - def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: - cursor.execute(sql, params) - return list(cursor.fetchall()) + def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + with self._config.provide_session() as driver: + return driver.select(sql, params) - def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: - cursor.execute(sql, params) - rowcount = _cursor_rowcount(cursor) + def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + with self._config.provide_session() as driver: + res = driver.execute(sql, params) + rowcount = res.rows_affected if commit: - conn.commit() + driver.commit() return rowcount @@ -411,7 +408,7 @@ class MssqlPythonADKMemoryStore(BaseSyncADKMemoryStore["MssqlPythonConfig"]): __slots__ = () - def __init__(self, config: "MssqlPythonConfig") -> None: + def __init__(self, config: MssqlPythonConfig) -> None: super().__init__(config) def create_tables(self) -> None: @@ -432,7 +429,7 @@ def create_tables(self) -> None: driver.execute(_create_index_sql(index_table, index_name, columns)) driver.commit() - def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: + def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | None = None) -> int: """Bulk insert memory entries with event-id deduplication.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -442,7 +439,6 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" owner_value = ", ?" if self._owner_id_column_name else "" - # Keep the key-range lock and insertion in one statement, including autocommit. sql = f""" INSERT INTO {_table_ref(self._memory_table)} ( id, session_id, app_name, user_id, scope, event_id, author, timestamp, @@ -455,7 +451,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object ); """ inserted = 0 - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + with self._config.provide_session() as driver: for entry in entries: params: tuple[Any, ...] = ( entry["id"], @@ -472,9 +468,9 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object ) if self._owner_id_column_name: params = (*params, owner_id) - cursor.execute(sql, (*params, entry["event_id"])) - inserted += _cursor_rowcount(cursor) - conn.commit() + res = driver.execute(sql, (*params, entry["event_id"])) + inserted += res.rows_affected + driver.commit() return inserted def search_entries( @@ -482,10 +478,10 @@ def search_entries( query: str, app_name: str, user_id: str, - limit: "int | None" = None, + limit: int | None = None, scope_filter: Literal["all", "user", "app"] = "all", - embedding: "Sequence[float] | None" = None, - ) -> "list[StoredMemory]": + embedding: Sequence[float] | None = None, + ) -> list[StoredMemory]: """Search memory entries by text query.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -509,7 +505,7 @@ def delete_entries_by_session(self, session_id: str) -> int: f"DELETE FROM {_table_ref(self._memory_table)} WHERE session_id = ?", (session_id,), commit=True ) - def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: + def delete_entries_older_than(self, days: int, app_name: str | None = None, scope: str | None = None) -> int: """Delete memory entries older than the retention window.""" clauses = ["inserted_at < DATEADD(day, -?, SYSUTCDATETIME())"] params: list[Any] = [days] @@ -550,7 +546,7 @@ def _memory_table_ddl(self) -> str: END; """ - def _memory_index_specs(self) -> "list[tuple[str, str, str]]": + def _memory_index_specs(self) -> list[tuple[str, str, str]]: """Return ``(index_name, table, columns)`` specs for memory-table indexes.""" return [ ( @@ -563,20 +559,19 @@ def _memory_index_specs(self) -> "list[tuple[str, str, str]]": (f"idx_{self._memory_table}_timestamp", self._memory_table, "timestamp DESC"), ] - def _drop_memory_table_sql(self) -> "list[str]": + def _drop_memory_table_sql(self) -> list[str]: return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"] - def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: - cursor.execute(sql, params) - return list(cursor.fetchall()) + def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + with self._config.provide_session() as driver: + return driver.select(sql, params) - def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: - with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: - cursor.execute(sql, params) - rowcount = _cursor_rowcount(cursor) + def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + with self._config.provide_session() as driver: + res = driver.execute(sql, params) + rowcount = res.rows_affected if commit: - conn.commit() + driver.commit() return rowcount @@ -590,7 +585,7 @@ def _adk_config(config: Any) -> MssqlPythonADKConfig: return cast("MssqlPythonADKConfig", adk_config) -def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: "str | None") -> str: +def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: str | None) -> str: owner_line = f",\n {owner_id_column_ddl}" if owner_id_column_ddl else "" return f""" IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo')) @@ -610,7 +605,7 @@ def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: """ -def _sessions_index_specs(table: str) -> "list[tuple[str, str, str]]": +def _sessions_index_specs(table: str) -> list[tuple[str, str, str]]: return [ (f"idx_{table}_app_user", table, "app_name, user_id"), (f"idx_{table}_update_time", table, "update_time DESC"), @@ -639,7 +634,7 @@ def _events_table_ddl(table: str, session_table: str, json_column_type: str) -> """ -def _events_index_specs(table: str) -> "list[tuple[str, str, str]]": +def _events_index_specs(table: str) -> list[tuple[str, str, str]]: return [ (f"idx_{table}_scope", table, "app_name, user_id, session_id, timestamp ASC"), (f"idx_{table}_session", table, "session_id, timestamp ASC"), @@ -695,7 +690,7 @@ def _create_index_sql(table: str, index_name: str, columns: str) -> str: return f"CREATE INDEX {_quote_identifier(index_name)} ON {_table_ref(table)} ({columns})" -def _casefold_names(rows: "list[Any]", key: str) -> "set[str]": +def _casefold_names(rows: list[Any], key: str) -> set[str]: """Collapse data-dictionary rows into a case-folded, schema-stripped name set.""" return {str(row.get(key, "")).rsplit(".", 1)[-1].casefold() for row in rows} @@ -714,7 +709,7 @@ def _insert_event_sql(table: str) -> str: """ -def _upsert_state_sql(table: str, key_columns: "tuple[str, ...]", key_params: "tuple[str, ...]") -> str: +def _upsert_state_sql(table: str, key_columns: tuple[str, ...], key_params: tuple[str, ...]) -> str: source_columns = ", ".join( f"{param} AS {_quote_identifier(column)}" for column, param in zip(key_columns, key_params, strict=False) ) @@ -750,8 +745,8 @@ def _upsert_metadata_sql(table: str) -> str: def _events_query( - table: str, app_name: str, user_id: str, session_id: str, after_timestamp: "datetime | None", limit: "int | None" -) -> "tuple[str, tuple[Any, ...]]": + table: str, app_name: str, user_id: str, session_id: str, after_timestamp: datetime | None, limit: int | None +) -> tuple[str, tuple[Any, ...]]: top_clause = "TOP (?) " if limit is not None else "" params: list[Any] = [limit] if limit is not None else [] params.extend([app_name, user_id, session_id]) @@ -768,7 +763,7 @@ def _events_query( return sql, tuple(params) -def _event_insert_params(event_record: StoredEvent) -> "tuple[Any, ...]": +def _event_insert_params(event_record: StoredEvent) -> tuple[Any, ...]: return ( event_record["id"], event_record["app_name"], @@ -798,7 +793,7 @@ def _event_record_from_row(row: Any) -> StoredEvent: ) -def _memory_record_from_row(row: Any) -> "StoredMemory": +def _memory_record_from_row(row: Any) -> StoredMemory: return cast( "StoredMemory", { @@ -821,7 +816,7 @@ def _memory_record_from_row(row: Any) -> "StoredMemory": def _build_mssql_scope_where( app_name: str, user_id: str, scope_filter: Literal["all", "user", "app"] -) -> "tuple[str, tuple[Any, ...]]": +) -> tuple[str, tuple[Any, ...]]: if scope_filter == "all": return "app_name = ? AND ((scope = 'user' AND user_id = ?) OR scope = 'app')", (app_name, user_id) if scope_filter == "user": @@ -829,7 +824,7 @@ def _build_mssql_scope_where( return "app_name = ? AND scope = 'app'", (app_name,) -def _json_dict(value: Any) -> "dict[str, Any]": +def _json_dict(value: Any) -> dict[str, Any]: if value is None: return {} if isinstance(value, dict): @@ -841,24 +836,9 @@ def _json_dict(value: Any) -> "dict[str, Any]": return cast("dict[str, Any]", from_json(str(value))) -def _cursor_rowcount(cursor: Any) -> int: - rowcount = getattr(cursor, "rowcount", 0) - return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 - - def _is_mssql_table_missing(exc: BaseException) -> bool: text = str(exc).lower() - return "invalid object name" in text or _mssql_error_number(exc) == MSSQL_TABLE_NOT_FOUND_ERROR - - -def _mssql_error_number(exc: BaseException) -> "int | None": - matches = MSSQL_ERROR_NUMBER_PATTERN.findall(str(exc)) - if not matches: - return None - try: - return int(matches[-1]) - except ValueError: - return None + return "invalid object name" in text or extract_error_number(exc) == MSSQL_TABLE_NOT_FOUND_ERROR def _quote_identifier(identifier: str) -> str: @@ -883,14 +863,8 @@ def _raise_session_not_found(session_id: str) -> None: def _session_list_query( - session_table: str, - app_name: str, - user_id: "str | None", - column: str, - direction: str, - limit: "int | None", - offset: int, -) -> "tuple[str, tuple[Any, ...]]": + session_table: str, app_name: str, user_id: str | None, column: str, direction: str, limit: int | None, offset: int +) -> tuple[str, tuple[Any, ...]]: """Return the bounded session-list query and its bound values.""" params: list[Any] = [app_name] where_clause = "app_name = ?" diff --git a/sqlspec/adapters/mssql_python/config.py b/sqlspec/adapters/mssql_python/config.py index 0fda4dc80..f40982ba5 100644 --- a/sqlspec/adapters/mssql_python/config.py +++ b/sqlspec/adapters/mssql_python/config.py @@ -1,6 +1,8 @@ """mssql-python database configuration.""" -from typing import TYPE_CHECKING, Any, ClassVar, TypedDict, cast +from collections.abc import Callable +from types import TracebackType +from typing import Any, ClassVar, TypedDict, cast from typing_extensions import NotRequired @@ -10,18 +12,12 @@ from sqlspec.adapters.mssql_python.migrations import MssqlPythonSyncMigrationTracker from sqlspec.adapters.mssql_python.pool import MssqlPythonConnectionPool from sqlspec.config import ExtensionConfigs, SyncDatabaseConfig -from sqlspec.core import TypeCoercionCapabilities +from sqlspec.core import StatementConfig, TypeCoercionCapabilities from sqlspec.driver import SyncPoolConnectionContext, SyncPoolSessionFactory +from sqlspec.observability import ObservabilityConfig from sqlspec.utils.config_tools import normalize_connection_config from sqlspec.utils.serializers import from_json, to_json -if TYPE_CHECKING: - from collections.abc import Callable - from types import TracebackType - - from sqlspec.core import StatementConfig - from sqlspec.observability import ObservabilityConfig - __all__ = ( "MssqlPythonConfig", "MssqlPythonConnectionParams", @@ -96,9 +92,9 @@ class MssqlPythonDriverFeatures(TypedDict): """mssql-python driver feature flags.""" use_pool: NotRequired[bool] - json_serializer: "NotRequired[Callable[[Any], str]]" - json_deserializer: "NotRequired[Callable[[str], Any]]" - on_connection_create: "NotRequired[Callable[[MssqlPythonConnection], None]]" + json_serializer: NotRequired[Callable[[Any], str]] + json_deserializer: NotRequired[Callable[[str], Any]] + on_connection_create: NotRequired[Callable[[MssqlPythonConnection], None]] enable_events: NotRequired[bool] @@ -111,15 +107,15 @@ def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) self._conn: MssqlPythonConnection | None = None - def __enter__(self) -> "MssqlPythonConnection": + def __enter__(self) -> MssqlPythonConnection: pool = self._config.provide_pool() conn = pool.acquire() self._conn = conn return cast("MssqlPythonConnection", conn) def __exit__( - self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" - ) -> "bool | None": + self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None + ) -> bool | None: if self._conn is not None: self._config.provide_pool().release(self._conn) self._conn = None @@ -133,13 +129,13 @@ def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) self._conn: MssqlPythonConnection | None = None - def acquire_connection(self) -> "MssqlPythonConnection": + def acquire_connection(self) -> MssqlPythonConnection: pool = self._config.provide_pool() conn = pool.acquire() self._conn = conn return cast("MssqlPythonConnection", conn) - def release_connection(self, _conn: "MssqlPythonConnection", **kwargs: Any) -> None: + def release_connection(self, _conn: MssqlPythonConnection, **kwargs: Any) -> None: if self._conn is None: return self._config.provide_pool().release(self._conn) @@ -151,38 +147,38 @@ class MssqlPythonConfig(SyncDatabaseConfig[MssqlPythonConnection, MssqlPythonCon __slots__ = ("_user_connection_hook",) - driver_type: "ClassVar[type[MssqlPythonDriver]]" = MssqlPythonDriver - connection_type: "ClassVar[type[MssqlPythonConnection]]" = MssqlPythonConnection - migration_tracker_type: "ClassVar[type[MssqlPythonSyncMigrationTracker]]" = MssqlPythonSyncMigrationTracker - supports_transactional_ddl: "ClassVar[bool]" = True - supports_migration_schemas: "ClassVar[bool]" = True - supports_native_arrow_export: "ClassVar[bool]" = True - supports_native_arrow_import: "ClassVar[bool]" = True - supports_arrow_streaming: "ClassVar[bool]" = True - supports_native_row_streaming: "ClassVar[bool]" = True - supports_native_parquet_export: "ClassVar[bool]" = False - supports_native_parquet_import: "ClassVar[bool]" = False - type_coercion_capabilities: "ClassVar[TypeCoercionCapabilities]" = TypeCoercionCapabilities( + driver_type: ClassVar[type[MssqlPythonDriver]] = MssqlPythonDriver + connection_type: ClassVar[type[MssqlPythonConnection]] = MssqlPythonConnection + migration_tracker_type: ClassVar[type[MssqlPythonSyncMigrationTracker]] = MssqlPythonSyncMigrationTracker + supports_transactional_ddl: ClassVar[bool] = True + supports_migration_schemas: ClassVar[bool] = True + supports_native_arrow_export: ClassVar[bool] = True + supports_native_arrow_import: ClassVar[bool] = True + supports_arrow_streaming: ClassVar[bool] = True + supports_native_row_streaming: ClassVar[bool] = True + supports_native_parquet_export: ClassVar[bool] = False + supports_native_parquet_import: ClassVar[bool] = False + type_coercion_capabilities: ClassVar[TypeCoercionCapabilities] = TypeCoercionCapabilities( datetime_binding="native", timestamp_precision="microsecond", json_columns_decoded=False, uuid_binding="native" ) - _connection_context_class: "ClassVar[type[MssqlPythonConnectionContext]]" = MssqlPythonConnectionContext - _session_factory_class: "ClassVar[type[_MssqlPythonSyncSessionConnectionHandler]]" = ( + _connection_context_class: ClassVar[type[MssqlPythonConnectionContext]] = MssqlPythonConnectionContext + _session_factory_class: ClassVar[type[_MssqlPythonSyncSessionConnectionHandler]] = ( _MssqlPythonSyncSessionConnectionHandler ) - _session_context_class: "ClassVar[type[MssqlPythonSessionContext]]" = MssqlPythonSessionContext + _session_context_class: ClassVar[type[MssqlPythonSessionContext]] = MssqlPythonSessionContext _default_statement_config = default_statement_config def __init__( self, *, - connection_config: "MssqlPythonPoolParams | dict[str, Any] | None" = None, - connection_instance: "MssqlPythonConnectionPool | None" = None, - migration_config: "dict[str, Any] | None" = None, - statement_config: "StatementConfig | None" = None, - driver_features: "MssqlPythonDriverFeatures | dict[str, Any] | None" = None, - bind_key: "str | None" = None, - extension_config: "ExtensionConfigs | None" = None, - observability_config: "ObservabilityConfig | None" = None, + connection_config: MssqlPythonPoolParams | dict[str, Any] | None = None, + connection_instance: MssqlPythonConnectionPool | None = None, + migration_config: dict[str, Any] | None = None, + statement_config: StatementConfig | None = None, + driver_features: MssqlPythonDriverFeatures | dict[str, Any] | None = None, + bind_key: str | None = None, + extension_config: ExtensionConfigs | None = None, + observability_config: ObservabilityConfig | None = None, **kwargs: Any, ) -> None: normalized, features_dict, user_connection_hook = _normalize_mssql_python_init( @@ -202,11 +198,11 @@ def __init__( **kwargs, ) - def create_connection(self) -> "MssqlPythonConnection": + def create_connection(self) -> MssqlPythonConnection: pool = self.provide_pool() return pool.acquire() - def get_signature_namespace(self) -> "dict[str, Any]": + def get_signature_namespace(self) -> dict[str, Any]: namespace = super().get_signature_namespace() namespace.update({ "MssqlPythonConfig": MssqlPythonConfig, @@ -221,7 +217,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": }) return namespace - def _create_pool(self) -> "MssqlPythonConnectionPool": + def _create_pool(self) -> MssqlPythonConnectionPool: return _create_mssql_python_pool(dict(self.connection_config), self.driver_features, self._user_connection_hook) def _close_pool(self) -> None: @@ -244,10 +240,10 @@ def _apply_json_serializer_override(statement_config: Any, features_dict: dict[s def _create_mssql_python_pool( - connection_config: "dict[str, Any]", - driver_features: "dict[str, Any]", - on_connection_create: "Callable[[MssqlPythonConnection], None] | None" = None, -) -> "MssqlPythonConnectionPool": + connection_config: dict[str, Any], + driver_features: dict[str, Any], + on_connection_create: Callable[[MssqlPythonConnection], None] | None = None, +) -> MssqlPythonConnectionPool: pool_size = int(connection_config.get("pool_size", 100)) pool_idle_timeout = int(connection_config.get("pool_idle_timeout", 600)) pool_enabled = bool(connection_config.get("pool_enabled", driver_features.get("use_pool", True))) @@ -263,9 +259,9 @@ def _create_mssql_python_pool( def _normalize_mssql_python_init( - connection_config: "MssqlPythonPoolParams | dict[str, Any] | None", - driver_features: "MssqlPythonDriverFeatures | dict[str, Any] | None", -) -> "tuple[dict[str, Any], dict[str, Any], Callable[[MssqlPythonConnection], None] | None]": + connection_config: MssqlPythonPoolParams | dict[str, Any] | None, + driver_features: MssqlPythonDriverFeatures | dict[str, Any] | None, +) -> tuple[dict[str, Any], dict[str, Any], Callable[[MssqlPythonConnection], None] | None]: normalized = normalize_connection_config(connection_config) _, features_dict = apply_driver_features(default_statement_config, driver_features) hook = cast("Callable[[MssqlPythonConnection], None] | None", features_dict.pop("on_connection_create", None)) diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index 7d7e26dbf..e5b8e0156 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -1,9 +1,12 @@ """mssql-python adapter core helpers.""" import re +from collections.abc import Callable, Mapping, Sequence from importlib.metadata import PackageNotFoundError, version -from typing import TYPE_CHECKING, Any, Final +from logging import Logger +from typing import Any, Final +from sqlspec.core import StatementConfig from sqlspec.core.parameters import ParameterStyle from sqlspec.core.parameters._registry import build_statement_config_from_profile from sqlspec.core.parameters._types import DriverParameterProfile @@ -25,12 +28,6 @@ from sqlspec.utils.serializers import from_json, to_json from sqlspec.utils.type_converters import build_uuid_coercions -if TYPE_CHECKING: - from collections.abc import Callable, Mapping, Sequence - from logging import Logger - - from sqlspec.core import StatementConfig - __all__ = ( "MSSQL_PYTHON_VERSION", "apply_driver_features", @@ -40,6 +37,7 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "extract_error_number", "materialize_tuple_rows", ) @@ -98,7 +96,49 @@ } -def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) -> SQLSpecError: +def extract_error_number(exc: BaseException | None) -> int | None: + """Extract numeric SQL Server error code using fast string parsing before regex fallback.""" + if exc is None: + return None + ddbc_err = getattr(exc, "ddbc_error", None) + if isinstance(ddbc_err, str) and ddbc_err.startswith("("): + end_idx = ddbc_err.find(",") + if end_idx == -1: + end_idx = ddbc_err.find(")") + if end_idx != -1: + num_str = ddbc_err[1:end_idx].strip() + try: + return int(num_str) + except ValueError: + pass + + if exc.args and isinstance(exc.args[0], str): + msg = exc.args[0] + start_idx = msg.rfind("(") + if start_idx != -1: + end_idx = msg.find(",", start_idx) + if end_idx == -1: + end_idx = msg.find(")", start_idx) + if end_idx != -1: + num_str = msg[start_idx + 1 : end_idx].strip() + try: + return int(num_str) + except ValueError: + pass + + matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) + if not matches: + return None + try: + return int(matches[-1]) + except ValueError: + return None + + +_extract_error_number = extract_error_number + + +def create_mapped_exception(error: Exception, *, logger: Logger | None = None) -> SQLSpecError: """Map a mssql-python exception to SQLSpec's exception hierarchy.""" error_number = _extract_error_number(error) if error_number == _MSSQL_CONSTRAINT_547: @@ -130,22 +170,25 @@ def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) return SQLSpecError(f"SQL Server database error. Original error: {error}") -def materialize_tuple_rows(fetched: "Sequence[Any] | None") -> "list[tuple[Any, ...]]": - """Materialize mssql-python ``Row`` objects into plain tuples. +def materialize_tuple_rows(fetched: Sequence[Any] | None) -> list[tuple[Any, ...]]: + """Materialize mssql-python Row objects into plain tuples. - ``mssql-python`` returns ``mssql_python.Row`` objects that are iterable and - indexable but are not ``tuple`` subclasses. The driver reports - ``row_format="tuple"``, so fetched rows are converted to real tuples to keep - that contract accurate when results are materialized. + Accesses row._values directly when available, bypassing Python's __iter__ + protocol for significantly higher throughput on large result sets. """ if not fetched: return [] + first = fetched[0] + if isinstance(first, tuple): + return list(fetched) if not isinstance(fetched, list) else fetched + if hasattr(first, "_values"): + return [tuple(row._values) if not isinstance(row._values, tuple) else row._values for row in fetched] return [tuple(row) for row in fetched] def apply_driver_features( - statement_config: "StatementConfig", driver_features: "Mapping[str, Any] | None" -) -> "tuple[StatementConfig, dict[str, Any]]": + statement_config: StatementConfig, driver_features: Mapping[str, Any] | None +) -> tuple[StatementConfig, dict[str, Any]]: """Merge mssql-python driver-feature defaults with caller overrides.""" defaults: dict[str, Any] = {"use_pool": True, "json_serializer": to_json, "json_deserializer": from_json} defaults.update(driver_features or {}) @@ -155,7 +198,7 @@ def apply_driver_features( def build_connection_config(params: dict[str, Any]) -> tuple[str, dict[str, Any]]: """Build an ODBC connection string and mssql-python connect kwargs. - When both ``connection_string`` and discrete connection fields are provided, + When both connection_string and discrete connection fields are provided, discrete fields take precedence and override matching keys in the connection string. Key names are normalized case-insensitively to prevent duplicate keywords, satisfying mssql-python driver requirements. @@ -241,7 +284,7 @@ def build_connection_config(params: dict[str, Any]) -> tuple[str, dict[str, Any] return ";".join(parts) + ";", connect_kwargs -def build_profile() -> "DriverParameterProfile": +def build_profile() -> DriverParameterProfile: """Create the mssql-python driver parameter profile.""" return DriverParameterProfile( name="mssql_python", @@ -260,14 +303,14 @@ def build_profile() -> "DriverParameterProfile": ) -def build_statement_config(*, json_serializer: "Callable[[Any], str] | None" = None) -> "StatementConfig": +def build_statement_config(*, json_serializer: Callable[[Any], str] | None = None) -> StatementConfig: """Construct the mssql-python statement configuration.""" return build_statement_config_from_profile( driver_profile, statement_overrides={"dialect": "tsql"}, json_serializer=json_serializer or to_json ) -def _constraint_exception_from_message(error: Exception) -> "SQLSpecError | None": +def _constraint_exception_from_message(error: Exception) -> SQLSpecError | None: """Classify SQL Server constraint messages when a driver omits the native error number.""" message = str(error) normalized = message.lower() @@ -282,7 +325,7 @@ def _constraint_exception_from_message(error: Exception) -> "SQLSpecError | None return None -def _custom_type_coercions() -> "dict[type, Callable[[Any], Any]]": +def _custom_type_coercions() -> dict[type, Callable[[Any], Any]]: """Return custom type coercions for mssql-python.""" return {bool: _identity, int: _identity, float: _identity, bytes: _identity, **build_uuid_coercions(native=True)} @@ -330,16 +373,6 @@ def _append_port(server: str, port: Any) -> str: return f"{server},{port}" -def _extract_error_number(exc: Exception) -> "int | None": - matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) - if not matches: - return None - try: - return int(matches[-1]) - except ValueError: - return None - - MSSQL_PYTHON_VERSION: Final[tuple[int, int, int]] = _parse_version() driver_profile = build_profile() default_statement_config = build_statement_config() diff --git a/sqlspec/adapters/mssql_python/data_dictionary.py b/sqlspec/adapters/mssql_python/data_dictionary.py index 8e2e13d6d..cb4545719 100644 --- a/sqlspec/adapters/mssql_python/data_dictionary.py +++ b/sqlspec/adapters/mssql_python/data_dictionary.py @@ -1,6 +1,6 @@ """mssql-python data dictionary.""" -from typing import TYPE_CHECKING, Any, ClassVar, cast +from typing import TYPE_CHECKING, Any, ClassVar, Final, cast from mypy_extensions import mypyc_attr @@ -44,12 +44,14 @@ from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver from sqlspec.core import SQL - from sqlspec.data_dictionary._types import DialectConfig, MetadataCapabilityProfile + from sqlspec.data_dictionary import DialectConfig, MetadataCapabilityProfile __all__ = ("MssqlPythonSyncDataDictionary", "MssqlVersionInfo") logger = get_logger("sqlspec.adapters.mssql_python.data_dictionary") +MSSQL_VECTOR_MIN_MAJOR: Final[int] = 17 + class MssqlVersionInfo(VersionInfo): """MSSQL database version info with build, revision, and Azure SQL detection.""" @@ -74,6 +76,10 @@ def supports_native_json(self) -> bool: """Return whether this server supports the native JSON type.""" return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) + def supports_vector(self) -> bool: + """Return whether this server supports native VECTOR data types and functions.""" + return self.is_azure_sql or self.major >= MSSQL_VECTOR_MIN_MAJOR + @property def version_tuple(self) -> "tuple[int, int, int]": """Get version tuple using the MSSQL build number as the third component.""" @@ -132,6 +138,8 @@ def _build_version_info( def _get_optimal_type_from_version(self, version_info: MssqlVersionInfo | None, type_category: str) -> str: if type_category in {"json", "jsonb"} and version_info is not None and version_info.supports_native_json(): return "JSON" + if type_category == "vector" and version_info is not None and version_info.supports_vector(): + return "VECTOR" return self.get_dialect_config().get_optimal_type(type_category) @@ -188,6 +196,8 @@ def get_version(self, driver: "MssqlPythonDriver") -> MssqlVersionInfo | None: def get_feature_flag(self, driver: "MssqlPythonDriver", feature: str) -> bool: """Check whether SQL Server supports a feature.""" version_info = self.get_version(driver) + if feature == "supports_vector": + return bool(version_info and version_info.supports_vector()) return resolve_mssql_feature_flag( feature, major=version_info.major if version_info is not None else 0, diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index 5eddec33a..da92bf8c1 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -1,7 +1,8 @@ """mssql-python sync and async drivers.""" import contextlib -from typing import TYPE_CHECKING, Any, TypedDict, cast +from collections.abc import Iterable +from typing import Any, TypedDict, cast from typing_extensions import NotRequired @@ -19,7 +20,13 @@ materialize_tuple_rows, ) from sqlspec.adapters.mssql_python.data_dictionary import MssqlPythonSyncDataDictionary +from sqlspec.builder import QueryBuilder from sqlspec.core import ( + SQL, + ArrowResult, + Statement, + StatementConfig, + StatementFilter, build_arrow_result_from_reader, build_arrow_result_from_table, get_cache_config, @@ -27,27 +34,20 @@ ) from sqlspec.driver import ( BaseSyncExceptionHandler, + ExecutionResult, SyncDriverAdapterBase, SyncRowStream, rows_to_dicts, validate_savepoint_name, ) from sqlspec.exceptions import SQLSpecError +from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry +from sqlspec.typing import ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters from sqlspec.utils.arrow_helpers import arrow_reader_with_deferred_close from sqlspec.utils.logging import get_logger from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.text import split_qualified_identifier -if TYPE_CHECKING: - from collections.abc import Iterable - - from sqlspec.builder import QueryBuilder - from sqlspec.core import SQL, ArrowResult, Statement, StatementConfig, StatementFilter - from sqlspec.driver import ExecutionResult - from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry - from sqlspec.typing import ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters - - __all__ = ( "MssqlPythonBulkCopyResult", "MssqlPythonCursor", @@ -66,6 +66,7 @@ class MssqlPythonBulkCopyResult(TypedDict): rows_copied: int batch_count: NotRequired[int] elapsed_time: NotRequired[float] + rows_per_second: NotRequired[float] class MssqlPythonExceptionHandler(BaseSyncExceptionHandler): @@ -73,7 +74,7 @@ class MssqlPythonExceptionHandler(BaseSyncExceptionHandler): __slots__ = () - def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool: + def _handle_exception(self, exc_type: type[BaseException] | None, exc_val: BaseException) -> bool: if exc_type is None: return False if isinstance(exc_val, MssqlPythonError): @@ -109,7 +110,7 @@ def start(self) -> None: raise self._cursor_manager = cursor_manager - def fetch_chunk(self) -> "list[dict[str, Any]]": + def fetch_chunk(self) -> list[dict[str, Any]]: cursor_manager = self._cursor_manager if cursor_manager is None or cursor_manager.cursor is None: return [] @@ -149,9 +150,9 @@ class MssqlPythonDriver(SyncDriverAdapterBase): def __init__( self, - connection: "MssqlPythonConnection", - statement_config: "StatementConfig | None" = None, - driver_features: "dict[str, Any] | None" = None, + connection: MssqlPythonConnection, + statement_config: StatementConfig | None = None, + driver_features: dict[str, Any] | None = None, ) -> None: if statement_config is None: statement_config = default_statement_config.replace( @@ -165,12 +166,12 @@ def __init__( self._transaction_active = False @property - def data_dictionary(self) -> "MssqlPythonSyncDataDictionary": + def data_dictionary(self) -> MssqlPythonSyncDataDictionary: if self._data_dictionary is None: self._data_dictionary = MssqlPythonSyncDataDictionary() return self._data_dictionary - def dispatch_execute(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> "ExecutionResult": + def dispatch_execute(self, cursor: MssqlPythonRawCursor, statement: SQL) -> ExecutionResult: sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) _execute_cursor(cursor, sql, prepared_parameters) @@ -188,28 +189,28 @@ def dispatch_execute(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> return self.create_execution_result(cursor, rowcount_override=_cursor_rowcount(cursor)) - def dispatch_execute_many(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> "ExecutionResult": + def dispatch_execute_many(self, cursor: MssqlPythonRawCursor, statement: SQL) -> ExecutionResult: sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) cursor.executemany(sql, cast("Any", prepared_parameters)) return self.create_execution_result(cursor, rowcount_override=_cursor_rowcount(cursor), is_many_result=True) - def dispatch_execute_script(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> "ExecutionResult": + def dispatch_execute_script(self, cursor: MssqlPythonRawCursor, statement: SQL) -> ExecutionResult: sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True) successful_count = 0 for stmt in statements: - _execute_cursor(cursor, stmt, prepared_parameters) + _execute_cursor(cursor, stmt, prepared_parameters, use_prepare=False) successful_count += 1 return self.create_execution_result( cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) - def collect_rows(self, cursor: "MssqlPythonRawCursor", fetched: "list[Any]") -> "tuple[list[Any], list[str], int]": + def collect_rows(self, cursor: MssqlPythonRawCursor, fetched: list[Any]) -> tuple[list[Any], list[str], int]: column_names = _resolve_column_names(cursor.description, self._column_name_cache) rows = materialize_tuple_rows(fetched) return rows, column_names, len(rows) - def resolve_rowcount(self, cursor: "MssqlPythonRawCursor") -> int: + def resolve_rowcount(self, cursor: MssqlPythonRawCursor) -> int: return _cursor_rowcount(cursor) def begin(self) -> None: @@ -243,13 +244,13 @@ def rollback(self) -> None: self._transaction_active = False self._restore_connection_autocommit() - def with_cursor(self, connection: "MssqlPythonConnection") -> "MssqlPythonCursor": + def with_cursor(self, connection: MssqlPythonConnection) -> MssqlPythonCursor: return MssqlPythonCursor(connection) - def handle_database_exceptions(self) -> "MssqlPythonExceptionHandler": + def handle_database_exceptions(self) -> MssqlPythonExceptionHandler: return MssqlPythonExceptionHandler() - def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "SyncRowStream[dict[str, Any]] | None": + def dispatch_select_stream(self, statement: SQL, chunk_size: int) -> SyncRowStream[dict[str, Any]] | None: """Return a native mssql-python row stream backed by ``fetchmany()``.""" if not statement.returns_rows(): return None @@ -269,13 +270,17 @@ def set_migration_session_schema(self, schema: str) -> None: """Point the database user's default schema at the migration schema, remembering the prior one.""" with self.with_cursor(self.connection) as cursor: if self._migration_schema_restore is None: - _execute_cursor(cursor, "SELECT USER_NAME() AS user_name, SCHEMA_NAME() AS schema_name;", None) + _execute_cursor( + cursor, "SELECT USER_NAME() AS user_name, SCHEMA_NAME() AS schema_name;", None, use_prepare=False + ) row: Any = cursor.fetchone() user_name, current_schema = row[0], row[1] - _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None) + _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None, use_prepare=False) self._migration_schema_restore = (str(user_name), str(current_schema)) return - _execute_cursor(cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None) + _execute_cursor( + cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None, use_prepare=False + ) def reset_migration_session_schema(self) -> None: """Restore the user's default schema captured by set_migration_session_schema and commit it.""" @@ -283,28 +288,28 @@ def reset_migration_session_schema(self) -> None: return user_name, previous_schema = self._migration_schema_restore with self.with_cursor(self.connection) as cursor: - _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None) + _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None, use_prepare=False) self.connection.commit() self._migration_schema_restore = None def has_schema(self, schema: str) -> bool: """Return whether the specified schema exists.""" with self.with_cursor(self.connection) as cursor: - _execute_cursor(cursor, "SELECT 1 FROM sys.schemas WHERE name = ?", (schema,)) + _execute_cursor(cursor, "SELECT 1 FROM sys.schemas WHERE name = ?", (schema,), use_prepare=False) return cursor.fetchone() is not None def select_to_arrow( self, - statement: "Statement | QueryBuilder", + statement: Statement | QueryBuilder, /, - *parameters: "StatementParameters | StatementFilter", - statement_config: "StatementConfig | None" = None, - return_format: "ArrowReturnFormat" = "table", + *parameters: StatementParameters | StatementFilter, + statement_config: StatementConfig | None = None, + return_format: ArrowReturnFormat = "table", native_only: bool = False, batch_size: int | None = None, arrow_schema: Any = None, **kwargs: Any, - ) -> "ArrowResult": + ) -> ArrowResult: """Execute a query and return native mssql-python Arrow results.""" ensure_pyarrow() config = statement_config or self.statement_config @@ -314,7 +319,7 @@ def select_to_arrow( arrow_kwargs: dict[str, int] = {"batch_size": batch_size} if batch_size is not None else {} table: Any | None = None - if return_format == "reader": + if return_format in ("reader", "batches"): cursor_manager = self.with_cursor(self.connection) cursor = None reader: object | None = None @@ -355,16 +360,6 @@ def select_to_arrow( exc_handler = self.handle_database_exceptions() with exc_handler, self.with_cursor(self.connection) as cursor: _execute_cursor(cursor, sql, prepared_parameters) - if return_format == "batches": - reader = _cursor_arrow_reader(cursor, arrow_kwargs) - if reader is not None: - return build_arrow_result_from_reader( - prepared_statement, - reader, - return_format=return_format, - batch_size=batch_size, - arrow_schema=arrow_schema, - ) table = cursor.arrow(**arrow_kwargs) self._check_pending_exception(exc_handler) @@ -378,7 +373,7 @@ def select_to_arrow( def bulk_copy( self, target_table: str, - rows: "Iterable[tuple[Any, ...]]", + rows: Iterable[tuple[Any, ...]], *, batch_size: int = 0, timeout: int = 30, @@ -414,26 +409,105 @@ def bulk_copy( def load_from_arrow( self, table: str, - source: "ArrowResult | Any", + source: ArrowResult | Any, *, - partitioner: "dict[str, object] | None" = None, + partitioner: dict[str, object] | None = None, overwrite: bool = False, - telemetry: "StorageTelemetry | None" = None, - ) -> "StorageBridgeJob": + telemetry: StorageTelemetry | None = None, + batch_size: int = 0, + timeout: int = 30, + table_lock: bool = True, + check_constraints: bool = False, + fire_triggers: bool = False, + keep_identity: bool = False, + keep_nulls: bool = False, + use_internal_transaction: bool = False, + column_mappings: list[str] | list[tuple[int, str]] | None = None, + ) -> StorageBridgeJob: """Load Arrow data into SQL Server via BulkCopy.""" self._require_capability("arrow_import_enabled") - arrow_table = self._coerce_arrow_table(source) if overwrite: + quoted_table = _quote_mssql_table(table) exc_handler = self.handle_database_exceptions() with exc_handler, self.with_cursor(self.connection) as cursor: - cursor.execute(f"DELETE FROM {_quote_mssql_table(table)}") + try: + _execute_cursor(cursor, f"TRUNCATE TABLE {quoted_table}", None, use_prepare=False) + except Exception as exc: + error_msg = str(exc) + if "4712" in error_msg or "foreign key" in error_msg.lower(): + _execute_cursor(cursor, f"DELETE FROM {quoted_table}", None, use_prepare=False) + else: + raise self._check_pending_exception(exc_handler) - if arrow_table.num_rows: + + raw_result: Any = None + is_stream = hasattr(source, "__arrow_c_stream__") + is_reader = False + try: + import pyarrow as pa + + is_reader = isinstance(source, (pa.RecordBatchReader, pa.RecordBatch)) + except ImportError: + pass + + if is_stream or is_reader: + cols = column_mappings + source_schema = getattr(source, "schema", None) + if cols is None and source_schema is not None: + schema_names = getattr(source_schema, "names", None) + if schema_names is not None: + cols = list(schema_names) exc_handler = self.handle_database_exceptions() with exc_handler, self.with_cursor(self.connection) as cursor: - cursor.bulkcopy_arrow(table, arrow_table, column_mappings=list(arrow_table.column_names)) + raw_result = cursor.bulkcopy_arrow( + table, + source, + batch_size=batch_size, + timeout=timeout, + table_lock=table_lock, + check_constraints=check_constraints, + fire_triggers=fire_triggers, + keep_identity=keep_identity, + keep_nulls=keep_nulls, + use_internal_transaction=use_internal_transaction, + column_mappings=cols, + ) self._check_pending_exception(exc_handler) - telemetry_payload = self._ingest_telemetry(arrow_table) + telemetry_payload = cast("StorageTelemetry", {"destination": table, "format": "arrow", "extra": {}}) + else: + arrow_table = self._coerce_arrow_table(source) + cols = column_mappings or list(arrow_table.column_names) + if arrow_table.num_rows: + exc_handler = self.handle_database_exceptions() + with exc_handler, self.with_cursor(self.connection) as cursor: + raw_result = cursor.bulkcopy_arrow( + table, + arrow_table, + batch_size=batch_size, + timeout=timeout, + table_lock=table_lock, + check_constraints=check_constraints, + fire_triggers=fire_triggers, + keep_identity=keep_identity, + keep_nulls=keep_nulls, + use_internal_transaction=use_internal_transaction, + column_mappings=cols, + ) + self._check_pending_exception(exc_handler) + telemetry_payload = self._ingest_telemetry(arrow_table) + + extra = telemetry_payload.setdefault("extra", {}) + if isinstance(raw_result, dict): + if "rows_copied" in raw_result: + telemetry_payload["rows_processed"] = raw_result["rows_copied"] + extra["rows_ingested"] = raw_result["rows_copied"] + if "elapsed_time" in raw_result: + extra["elapsed_time"] = raw_result["elapsed_time"] + if "rows_per_second" in raw_result: + extra["rows_per_second"] = raw_result["rows_per_second"] + if "batch_count" in raw_result: + extra["batch_count"] = raw_result["batch_count"] + telemetry_payload["destination"] = table self._attach_partition_telemetry(telemetry_payload, partitioner) return self._storage_job(telemetry_payload, telemetry) @@ -441,12 +515,12 @@ def load_from_arrow( def load_from_storage( self, table: str, - source: "StorageDestination", + source: StorageDestination, *, - file_format: "StorageFormat", - partitioner: "dict[str, object] | None" = None, + file_format: StorageFormat, + partitioner: dict[str, object] | None = None, overwrite: bool = False, - ) -> "StorageBridgeJob": + ) -> StorageBridgeJob: """Load staged artifacts from storage into SQL Server via BulkCopy.""" arrow_table, inbound = self._read_storage_arrow(source, file_format=file_format) return self.load_from_arrow(table, arrow_table, partitioner=partitioner, overwrite=overwrite, telemetry=inbound) @@ -476,19 +550,19 @@ def _quote_mssql_table(table: str) -> str: return ".".join(_quote_tsql_identifier(part) for part in split_qualified_identifier(table)) -def _execute_cursor(cursor: "MssqlPythonRawCursor", sql: str, parameters: Any) -> None: +def _execute_cursor(cursor: MssqlPythonRawCursor, sql: str, parameters: Any, *, use_prepare: bool = True) -> None: if parameters is None: - cursor.execute(sql) + cursor.execute(sql, use_prepare=use_prepare) else: - cursor.execute(sql, parameters) + cursor.execute(sql, parameters, use_prepare=use_prepare) -def _cursor_rowcount(cursor: "MssqlPythonRawCursor") -> int: +def _cursor_rowcount(cursor: MssqlPythonRawCursor) -> int: rowcount = getattr(cursor, "rowcount", 0) return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 -def _resolve_column_names(description: Any, cache: "dict[int, tuple[Any, list[str]]]") -> list[str]: +def _resolve_column_names(description: Any, cache: dict[int, tuple[Any, list[str]]]) -> list[str]: if not description: return [] cache_key = id(description) @@ -502,16 +576,14 @@ def _resolve_column_names(description: Any, cache: "dict[int, tuple[Any, list[st return column_names -def _cursor_arrow_reader( - cursor: "MssqlPythonRawCursor", arrow_kwargs: "dict[str, int]" -) -> "ArrowRecordBatchReader | None": +def _cursor_arrow_reader(cursor: MssqlPythonRawCursor, arrow_kwargs: dict[str, int]) -> ArrowRecordBatchReader | None: arrow_reader = getattr(cursor, "arrow_reader", None) if not callable(arrow_reader): return None return cast("ArrowRecordBatchReader", arrow_reader(**arrow_kwargs)) -def _coerce_bulk_copy_result(result: Any, cursor: "MssqlPythonRawCursor") -> MssqlPythonBulkCopyResult: +def _coerce_bulk_copy_result(result: Any, cursor: MssqlPythonRawCursor) -> MssqlPythonBulkCopyResult: if isinstance(result, dict): return cast("MssqlPythonBulkCopyResult", dict(result)) return {"rows_copied": _cursor_rowcount(cursor)} diff --git a/sqlspec/adapters/mssql_python/litestar/store.py b/sqlspec/adapters/mssql_python/litestar/store.py index 1e4464f96..c67f20036 100644 --- a/sqlspec/adapters/mssql_python/litestar/store.py +++ b/sqlspec/adapters/mssql_python/litestar/store.py @@ -1,14 +1,12 @@ """mssql-python Litestar Store implementation.""" from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any +from typing import Any +from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ -if TYPE_CHECKING: - from sqlspec.adapters.mssql_python.config import MssqlPythonConfig - __all__ = ("MssqlPythonStore",) @@ -17,7 +15,7 @@ class MssqlPythonStore(BaseSQLSpecStore["MssqlPythonConfig"]): __slots__ = () - def __init__(self, config: "MssqlPythonConfig") -> None: + def __init__(self, config: MssqlPythonConfig) -> None: super().__init__(config) async def create_table(self) -> None: @@ -28,11 +26,11 @@ async def create_table(self) -> None: await async_(self._create_table)() await self.reconcile_schema(assume_existing=True) - async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": + async def get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: """Get a session value by key.""" return await async_(self._get)(key, renew_for) - async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: + async def set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: """Store a session value.""" await async_(self._set)(key, value, expires_in) @@ -48,7 +46,7 @@ async def exists(self, key: str) -> bool: """Check if a session key exists and is not expired.""" return await async_(self._exists)(key) - async def expires_in(self, key: str) -> "int | None": + async def expires_in(self, key: str) -> int | None: """Get the time in seconds until the session expires.""" return await async_(self._expires_in)(key) @@ -80,7 +78,7 @@ def _table_ddl(self) -> str: END; """ - def _drop_table_sql(self) -> "list[str]": + def _drop_table_sql(self) -> list[str]: """Get SQL Server DROP TABLE statements.""" return [f"IF OBJECT_ID(N'dbo.{self._table_name}', N'U') IS NOT NULL DROP TABLE dbo.{self._table_name};"] @@ -90,20 +88,14 @@ def _create_table(self) -> None: driver.commit() self._log_table_created() - def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": + def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = ? AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql, (key,)) - row = cursor.fetchone() - finally: - cursor.close() - + with self._config.provide_session() as driver: + row = driver.select_one_or_none(sql, (key,)) if row is None: return None @@ -111,23 +103,19 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if renew_for is not None and expires_at is not None: new_expires_at = self._calculate_expires_at(renew_for) if new_expires_at is not None: - update_cursor = conn.cursor() - try: - update_cursor.execute( - f""" - UPDATE {self._table_name} - SET expires_at = ?, updated_at = SYSUTCDATETIME() - WHERE session_id = ? - """, - (new_expires_at, key), - ) - finally: - update_cursor.close() - conn.commit() + driver.execute( + f""" + UPDATE {self._table_name} + SET expires_at = ?, updated_at = SYSUTCDATETIME() + WHERE session_id = ? + """, + (new_expires_at, key), + ) + driver.commit() return _coerce_bytes(_row_value(row, "data", 0)) - def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: + def _set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) sql = f""" @@ -143,31 +131,19 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No INSERT (session_id, data, expires_at) VALUES (src.session_id, src.data, src.expires_at); """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql, (key, data, expires_at)) - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + driver.execute(sql, (key, data, expires_at)) + driver.commit() def _delete(self, key: str) -> None: - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"DELETE FROM {self._table_name} WHERE session_id = ?", (key,)) - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + driver.execute(f"DELETE FROM {self._table_name} WHERE session_id = ?", (key,)) + driver.commit() def _delete_all(self) -> None: - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"TRUNCATE TABLE {self._table_name}") - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + driver.execute(f"TRUNCATE TABLE {self._table_name}") + driver.commit() self._log_delete_all() def _exists(self, key: str) -> bool: @@ -177,22 +153,12 @@ def _exists(self, key: str) -> bool: WHERE session_id = ? AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql, (key,)) - return cursor.fetchone() is not None - finally: - cursor.close() - - def _expires_in(self, key: str) -> "int | None": - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"SELECT expires_at FROM {self._table_name} WHERE session_id = ?", (key,)) - row = cursor.fetchone() - finally: - cursor.close() + with self._config.provide_session() as driver: + return driver.select_one_or_none(sql, (key,)) is not None + + def _expires_in(self, key: str) -> int | None: + with self._config.provide_session() as driver: + row = driver.select_one_or_none(f"SELECT expires_at FROM {self._table_name} WHERE session_id = ?", (key,)) if row is None: return None @@ -208,14 +174,10 @@ def _delete_expired(self) -> int: WHERE expires_at IS NOT NULL AND expires_at < SYSUTCDATETIME() """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql) - count = int(getattr(cursor, "rowcount", 0) or 0) - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + res = driver.execute(sql) + driver.commit() + count = res.rows_affected if count > 0: self._log_delete_expired(count) return count @@ -235,7 +197,7 @@ def _row_value(row: object, key: str, index: int) -> Any: return getattr(row, key, None) -def _normalize_utc(value: Any) -> "datetime | None": +def _normalize_utc(value: Any) -> datetime | None: if value is None: return None if not isinstance(value, datetime): diff --git a/sqlspec/adapters/mssql_python/pool.py b/sqlspec/adapters/mssql_python/pool.py index f6a1643c4..db48990c8 100644 --- a/sqlspec/adapters/mssql_python/pool.py +++ b/sqlspec/adapters/mssql_python/pool.py @@ -1,16 +1,15 @@ """mssql-python pool facade.""" +import contextlib import warnings -from typing import TYPE_CHECKING, Any, cast +from collections.abc import Callable +from typing import Any, cast from sqlspec.adapters.mssql_python._typing import MSSQL_PYTHON_MODULE, MssqlPythonConnection -if TYPE_CHECKING: - from collections.abc import Callable - __all__ = ("MssqlPythonConnectionPool",) -_POOLING_PARAMS: "tuple[int, int, bool] | None" = None +_POOLING_PARAMS: tuple[int, int, bool] | None = None class MssqlPythonConnectionPool: @@ -30,11 +29,11 @@ def __init__( self, *, connection_string: str, - connect_kwargs: "dict[str, Any] | None" = None, + connect_kwargs: dict[str, Any] | None = None, max_size: int = 100, idle_timeout: int = 600, enabled: bool = True, - on_connection_create: "Callable[[MssqlPythonConnection], None] | None" = None, + on_connection_create: Callable[[MssqlPythonConnection], None] | None = None, ) -> None: self.connection_string = connection_string self.connect_kwargs = connect_kwargs or {} @@ -51,10 +50,11 @@ def __init__( f"overwriting with {new_params}. Only one pool config per process is supported.", stacklevel=2, ) - MSSQL_PYTHON_MODULE.pooling(max_size=max_size, idle_timeout=idle_timeout, enabled=enabled) - _POOLING_PARAMS = new_params + if _POOLING_PARAMS is None or new_params != _POOLING_PARAMS: + MSSQL_PYTHON_MODULE.pooling(max_size=max_size, idle_timeout=idle_timeout, enabled=enabled) + _POOLING_PARAMS = new_params - def acquire(self) -> "MssqlPythonConnection": + def acquire(self) -> MssqlPythonConnection: if self._closed: msg = "Cannot acquire a connection from a closed mssql-python pool." raise RuntimeError(msg) @@ -65,8 +65,15 @@ def acquire(self) -> "MssqlPythonConnection": self.on_connection_create(connection) return connection - def release(self, connection: "MssqlPythonConnection") -> None: + def release(self, connection: MssqlPythonConnection) -> None: connection.close() - def close(self) -> None: + def close(self, *, close_driver_pooling: bool = False) -> None: self._closed = True + if close_driver_pooling: + global _POOLING_PARAMS + _POOLING_PARAMS = None + with contextlib.suppress(Exception): + ddbc = getattr(MSSQL_PYTHON_MODULE, "ddbc_bindings", None) + if ddbc is not None and hasattr(ddbc, "close_pooling"): + ddbc.close_pooling() diff --git a/sqlspec/adapters/mssql_python/type_converter.py b/sqlspec/adapters/mssql_python/type_converter.py index 7099e801a..18d2a626a 100644 --- a/sqlspec/adapters/mssql_python/type_converter.py +++ b/sqlspec/adapters/mssql_python/type_converter.py @@ -1,16 +1,14 @@ """Type converters for mssql-python parameter binding.""" -from typing import TYPE_CHECKING, Any, Final, cast +from collections.abc import Callable +from typing import Any, Final, cast from uuid import UUID +import pyarrow as pa + from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.serializers import from_json, to_json -if TYPE_CHECKING: - from collections.abc import Callable - - import pyarrow as pa - __all__ = ("MssqlPythonTypeConverter", "mssql_type_to_arrow") _MSSQL_ARROW_TYPE_SPECS: Final[dict[str, tuple[str, tuple[Any, ...], dict[str, Any]]]] = { @@ -42,6 +40,7 @@ "nvarchar": ("string", (), {}), "text": ("string", (), {}), "ntext": ("string", (), {}), + "json": ("string", (), {}), } @@ -56,12 +55,12 @@ class MssqlPythonTypeConverter: __slots__ = ("_json_deserializer", "_json_serializer") def __init__( - self, json_serializer: "Callable[[Any], str]" = to_json, json_deserializer: "Callable[[str], Any]" = from_json + self, json_serializer: Callable[[Any], str] = to_json, json_deserializer: Callable[[str], Any] = from_json ) -> None: self._json_serializer = json_serializer self._json_deserializer = json_deserializer - def coerce_bind_value(self, value: "Any") -> "Any": + def coerce_bind_value(self, value: Any) -> Any: """Coerce Python values before mssql-python parameter binding.""" if isinstance(value, (dict, list)): return self._json_serializer(value) @@ -69,14 +68,19 @@ def coerce_bind_value(self, value: "Any") -> "Any": return value return value - def coerce_read_value(self, value: "Any") -> "Any": + def coerce_read_value(self, value: Any) -> Any: """Coerce mssql-python result values after fetching.""" return value -def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: int | None = None) -> "pa.DataType": +def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: int | None = None) -> pa.DataType: """Resolve a T-SQL type name to an Arrow data type.""" normalized_type = sql_type.lower().split("(", 1)[0].strip() + if normalized_type == "vector": + ensure_pyarrow() + import pyarrow as pa + + return cast("pa.DataType", pa.list_(pa.float32())) if normalized_type in {"decimal", "numeric"} and precision is not None and scale is not None: return _arrow_type("decimal128", (precision, scale)) spec = _MSSQL_ARROW_TYPE_SPECS.get(normalized_type) @@ -86,7 +90,7 @@ def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: i return _arrow_type(name, args, kwargs) -def _arrow_type(name: str, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None) -> "pa.DataType": +def _arrow_type(name: str, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None) -> pa.DataType: ensure_pyarrow() import pyarrow as pa diff --git a/sqlspec/adapters/pymssql/_typing.py b/sqlspec/adapters/pymssql/_typing.py index 3ae4e9fab..9a704d45a 100644 --- a/sqlspec/adapters/pymssql/_typing.py +++ b/sqlspec/adapters/pymssql/_typing.py @@ -5,25 +5,25 @@ """ import contextlib +from collections.abc import Callable +from types import TracebackType from typing import TYPE_CHECKING, Any -import pymssql as _pymssql # pyright: ignore[reportMissingTypeStubs] -from pymssql import Connection as _PymssqlConnection # pyright: ignore[reportMissingTypeStubs] -from pymssql import Cursor as _PymssqlRawCursor # pyright: ignore[reportMissingTypeStubs] +import pymssql as _pymssql +from pymssql import Connection as _PymssqlConnection +from pymssql import Cursor as _PymssqlRawCursor from pymssql import Error as PymssqlError +from sqlspec.adapters.pymssql.driver import PymssqlDriver +from sqlspec.core import StatementConfig + PYMSSQL_MODULE = _pymssql if TYPE_CHECKING: - from collections.abc import Callable - from types import TracebackType from typing import TypeAlias from pymssql._pymssql import QueryParams as PymssqlQueryParams - from sqlspec.adapters.pymssql.driver import PymssqlDriver - from sqlspec.core import StatementConfig - PymssqlConnection: TypeAlias = _PymssqlConnection PymssqlRawCursor: TypeAlias = _PymssqlRawCursor @@ -48,11 +48,11 @@ class PymssqlCursor: __slots__ = ("connection", "cursor") - def __init__(self, connection: "PymssqlConnection") -> None: + def __init__(self, connection: PymssqlConnection) -> None: self.connection = connection self.cursor: PymssqlRawCursor | None = None - def __enter__(self) -> "PymssqlRawCursor": + def __enter__(self) -> PymssqlRawCursor: self.cursor = self.connection.cursor() return self.cursor @@ -77,11 +77,11 @@ class PymssqlSessionContext: def __init__( self, - acquire_connection: "Callable[[], Any]", - release_connection: "Callable[..., Any]", - statement_config: "StatementConfig", - driver_features: "dict[str, Any]", - prepare_driver: "Callable[[PymssqlDriver], PymssqlDriver]", + acquire_connection: Callable[[], Any], + release_connection: Callable[..., Any], + statement_config: StatementConfig, + driver_features: dict[str, Any], + prepare_driver: Callable[[PymssqlDriver], PymssqlDriver], ) -> None: self._acquire_connection = acquire_connection self._release_connection = release_connection @@ -91,7 +91,7 @@ def __init__( self._connection: Any = None self._driver: PymssqlDriver | None = None - def __enter__(self) -> "PymssqlDriver": + def __enter__(self) -> PymssqlDriver: from sqlspec.adapters.pymssql.driver import PymssqlDriver self._connection = self._acquire_connection() @@ -101,8 +101,11 @@ def __enter__(self) -> "PymssqlDriver": return self._prepare_driver(self._driver) def __exit__( - self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" - ) -> "bool | None": + self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None + ) -> bool | None: + if exc_type is not None and self._driver is not None: + with contextlib.suppress(Exception): + self._driver.rollback() if self._connection is not None: self._release_connection(self._connection, exc_type=exc_type, exc_val=exc_val, exc_tb=exc_tb) self._connection = None diff --git a/sqlspec/adapters/pymssql/adk/store.py b/sqlspec/adapters/pymssql/adk/store.py index 5658e83b4..0308fea95 100644 --- a/sqlspec/adapters/pymssql/adk/store.py +++ b/sqlspec/adapters/pymssql/adk/store.py @@ -1,34 +1,34 @@ """pymssql ADK stores for Google Agent Development Kit session storage.""" -import re -from datetime import datetime -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast +from collections.abc import Sequence +from datetime import datetime, timedelta +from typing import Any, ClassVar, Final, Literal, cast from typing_extensions import NotRequired -from sqlspec.adapters.pymssql._typing import PymssqlCursor, PymssqlError +from sqlspec.adapters.pymssql._typing import PymssqlError +from sqlspec.adapters.pymssql.config import PymssqlConfig +from sqlspec.adapters.pymssql.core import extract_error_number, quote_tsql_identifier from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo +from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.config import ADKConfig -from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options -from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore +from sqlspec.extensions.adk import ( + BaseSyncADKMemoryStore, + BaseSyncADKStore, + SessionOrderBy, + StoredEvent, + StoredMemory, + StoredSession, + normalize_session_list_options, +) from sqlspec.utils.serializers import from_json, to_json -if TYPE_CHECKING: - from collections.abc import Sequence - from datetime import timedelta - - from sqlspec.adapters.pymssql.config import PymssqlConfig - from sqlspec.adapters.pymssql.driver import PymssqlDriver - from sqlspec.extensions.adk import SessionOrderBy - from sqlspec.extensions.adk.memory._types import StoredMemory - __all__ = ("PymssqlADKConfig", "PymssqlADKMemoryStore", "PymssqlADKStore") MSSQL_TABLE_NOT_FOUND_ERROR: Final[int] = 208 MSSQL_DUPLICATE_OBJECT_ERROR: Final[int] = 2714 MSSQL_DUPLICATE_INDEX_ERROR: Final[int] = 1913 MSSQL_SCHEMA: Final[str] = "dbo" -MSSQL_ERROR_NUMBER_PATTERN: Final[re.Pattern[str]] = re.compile(r"\(([-]?\d+)\)") JSON_FALLBACK_COLUMN_TYPE: Final[str] = "NVARCHAR(MAX)" JSON_NATIVE_COLUMN_TYPE: Final[str] = "JSON" @@ -46,7 +46,7 @@ class PymssqlADKStore(BaseSyncADKStore["PymssqlConfig"]): connector_name: ClassVar[str] = "pymssql" __slots__ = ("_json_column_type", "_native_json") - def __init__(self, config: "PymssqlConfig") -> None: + def __init__(self, config: PymssqlConfig) -> None: super().__init__(config) adk_config = _adk_config(config) native_json = adk_config.get("native_json") @@ -74,7 +74,7 @@ def create_tables(self) -> None: driver.commit() def create_session( - self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None + self, session_id: str, app_name: str, user_id: str, state: dict[str, Any], owner_id: Any | None = None ) -> StoredSession: """Create a new ADK session.""" owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" @@ -98,8 +98,8 @@ def create_session( return _session_record_from_row(row) def get_session( - self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None - ) -> "StoredSession | None": + self, app_name: str, user_id: str, session_id: str, *, renew_for: int | timedelta | None = None + ) -> StoredSession | None: """Return a scoped session or ``None`` if absent.""" try: if renew_for is not None and self._calculate_expires_at(renew_for) is not None: @@ -126,7 +126,7 @@ def get_session( raise return _session_record_from_row(row) if row is not None else None - def update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None: + def update_session_state(self, app_name: str, user_id: str, session_id: str, state: dict[str, Any]) -> None: """Replace a session's durable state.""" self._execute( f""" @@ -141,13 +141,13 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta def list_sessions( self, app_name: str, - user_id: "str | None" = None, + user_id: str | None = None, *, - order_by: "SessionOrderBy" = "update_time", + order_by: SessionOrderBy = "update_time", descending: bool = True, - limit: "int | None" = None, - offset: "int | None" = None, - ) -> "list[StoredSession]": + limit: int | None = None, + offset: int | None = None, + ) -> list[StoredSession]: """List ADK sessions for an application, optionally scoped to a user.""" column, direction, page_limit, page_offset = normalize_session_list_options(order_by, descending, limit, offset) if page_limit == 0: @@ -182,10 +182,10 @@ def append_event_and_update_state( app_name: str, user_id: str, session_id: str, - state: "dict[str, Any]", + state: dict[str, Any], *, - app_state: "dict[str, Any] | None" = None, - user_state: "dict[str, Any] | None" = None, + app_state: dict[str, Any] | None = None, + user_state: dict[str, Any] | None = None, ) -> StoredSession: """Atomically append an event and update durable session/scoped state.""" update_sql = f""" @@ -194,21 +194,20 @@ def append_event_and_update_state( OUTPUT inserted.id, inserted.app_name, inserted.user_id, inserted.state, inserted.create_time, inserted.update_time WHERE app_name = %s AND user_id = %s AND id = %s """ - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + with self._config.provide_session() as driver: try: - cursor.execute(update_sql, (to_json(state), app_name, user_id, session_id)) - row = cursor.fetchone() + row = driver.select_one_or_none(update_sql, (to_json(state), app_name, user_id, session_id)) if row is None: _raise_session_not_found(session_id) - cursor.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) + driver.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) if app_state is not None: - cursor.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) + driver.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) if user_state is not None: - cursor.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) + driver.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) except Exception: - conn.rollback() + driver.rollback() raise - conn.commit() + driver.commit() return _session_record_from_row(row) def get_events( @@ -216,9 +215,9 @@ def get_events( app_name: str, user_id: str, session_id: str, - after_timestamp: "datetime | None" = None, - limit: "int | None" = None, - ) -> "list[StoredEvent]": + after_timestamp: datetime | None = None, + limit: int | None = None, + ) -> list[StoredEvent]: """Return events for a scoped session ordered by event timestamp.""" if limit == 0: return [] @@ -231,7 +230,7 @@ def get_events( raise return [_event_record_from_row(row) for row in rows] - def delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: + def delete_expired_events(self, before: datetime, app_name: str | None = None) -> int: """Delete events older than ``before``.""" sql = f"DELETE FROM {_table_ref(self._events_table)} WHERE timestamp < %s" params: list[Any] = [before] @@ -245,7 +244,7 @@ def delete_expired_events(self, before: datetime, app_name: "str | None" = None) return 0 raise - def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: + def delete_idle_sessions(self, updated_before: datetime, app_name: str | None = None) -> int: """Delete sessions whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._session_table)} WHERE update_time < %s" params: list[Any] = [updated_before] @@ -259,7 +258,7 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" return 0 raise - def delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: + def delete_idle_user_states(self, updated_before: datetime, app_name: str | None = None) -> int: """Delete user state rows whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._user_state_table)} WHERE update_time < %s" params: list[Any] = [updated_before] @@ -273,7 +272,7 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: "str | Non return 0 raise - def get_app_state(self, app_name: str) -> "dict[str, Any] | None": + def get_app_state(self, app_name: str) -> dict[str, Any] | None: """Return app-scoped state.""" try: row = self._execute_fetchone( @@ -285,7 +284,7 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": raise return _json_dict(row[0]) if row is not None else None - def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None": + def get_user_state(self, app_name: str, user_id: str) -> dict[str, Any] | None: """Return user-scoped state.""" try: row = self._execute_fetchone( @@ -302,15 +301,15 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" raise return _json_dict(row[0]) if row is not None else None - def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: + def upsert_app_state(self, app_name: str, state: dict[str, Any]) -> None: """Insert or replace app-scoped state.""" self._execute(self._upsert_app_state_sql(), (app_name, to_json(state)), commit=True) - def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: + def upsert_user_state(self, app_name: str, user_id: str, state: dict[str, Any]) -> None: """Insert or replace user-scoped state.""" self._execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(state)), commit=True) - def get_metadata(self, key: str) -> "str | None": + def get_metadata(self, key: str) -> str | None: """Return an ADK metadata value.""" try: row = self._execute_fetchone( @@ -326,7 +325,7 @@ def set_metadata(self, key: str, value: str) -> None: """Set an ADK metadata value.""" self._execute(_upsert_metadata_sql(self._metadata_table), (key, value), commit=True) - def _index_specs(self) -> "list[tuple[str, str, str]]": + def _index_specs(self) -> list[tuple[str, str, str]]: """Return ``(index_name, table, columns)`` specs for session and event indexes.""" return [*_sessions_index_specs(self._session_table), *_events_index_specs(self._events_table)] @@ -359,7 +358,7 @@ def _drop_user_states_table_sql(self) -> str: def _drop_metadata_table_sql(self) -> str: return f"DROP TABLE IF EXISTS {_table_ref(self._metadata_table)}" - def _drop_tables_sql(self) -> "list[str]": + def _drop_tables_sql(self) -> list[str]: return [ self._drop_metadata_table_sql(), self._drop_user_states_table_sql(), @@ -379,9 +378,9 @@ def _events_query( app_name: str, user_id: str, session_id: str, - after_timestamp: "datetime | None" = None, - limit: "int | None" = None, - ) -> "tuple[str, tuple[Any, ...]]": + after_timestamp: datetime | None = None, + limit: int | None = None, + ) -> tuple[str, tuple[Any, ...]]: return _events_query(self._events_table, app_name, user_id, session_id, after_timestamp, limit) def _json_column_type_sync(self) -> str: @@ -395,25 +394,23 @@ def _json_column_type_sync(self) -> str: self._json_column_type = _json_column_type_from_sync_driver(driver) return self._json_column_type - def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None": - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: - cursor.execute(sql, params) - row = cursor.fetchone() + def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> Any | None: + with self._config.provide_session() as driver: + row = driver.select_one_or_none(sql, params) if commit: - conn.commit() + driver.commit() return row - def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: - cursor.execute(sql, params) - return list(cursor.fetchall()) + def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + with self._config.provide_session() as driver: + return driver.select(sql, params) - def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: - cursor.execute(sql, params) - rowcount = _cursor_rowcount(cursor) + def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + with self._config.provide_session() as driver: + res = driver.execute(sql, params) + rowcount = res.rows_affected if commit: - conn.commit() + driver.commit() return rowcount @@ -422,7 +419,7 @@ class PymssqlADKMemoryStore(BaseSyncADKMemoryStore["PymssqlConfig"]): __slots__ = () - def __init__(self, config: "PymssqlConfig") -> None: + def __init__(self, config: PymssqlConfig) -> None: super().__init__(config) def create_tables(self) -> None: @@ -443,7 +440,7 @@ def create_tables(self) -> None: driver.execute(_create_index_sql(index_table, index_name, columns)) driver.commit() - def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: + def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | None = None) -> int: """Bulk insert memory entries with event-id deduplication.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -453,7 +450,6 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" owner_value = ", %s" if self._owner_id_column_name else "" - # Keep the key-range lock and insertion in one statement, including autocommit. sql = f""" INSERT INTO {_table_ref(self._memory_table)} ( id, session_id, app_name, user_id, scope, event_id, author, timestamp, @@ -466,7 +462,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object ); """ inserted = 0 - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + with self._config.provide_session() as driver: for entry in entries: params: tuple[Any, ...] = ( entry["id"], @@ -483,9 +479,9 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object ) if self._owner_id_column_name: params = (*params, owner_id) - cursor.execute(sql, (*params, entry["event_id"])) - inserted += _cursor_rowcount(cursor) - conn.commit() + res = driver.execute(sql, (*params, entry["event_id"])) + inserted += res.rows_affected + driver.commit() return inserted def search_entries( @@ -493,10 +489,10 @@ def search_entries( query: str, app_name: str, user_id: str, - limit: "int | None" = None, + limit: int | None = None, scope_filter: Literal["all", "user", "app"] = "all", - embedding: "Sequence[float] | None" = None, - ) -> "list[StoredMemory]": + embedding: Sequence[float] | None = None, + ) -> list[StoredMemory]: """Search memory entries by text query.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -520,7 +516,7 @@ def delete_entries_by_session(self, session_id: str) -> int: f"DELETE FROM {_table_ref(self._memory_table)} WHERE session_id = %s", (session_id,), commit=True ) - def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: + def delete_entries_older_than(self, days: int, app_name: str | None = None, scope: str | None = None) -> int: """Delete memory entries older than the retention window.""" clauses = ["inserted_at < DATEADD(day, -%s, SYSUTCDATETIME())"] params: list[Any] = [days] @@ -561,7 +557,7 @@ def _memory_table_ddl(self) -> str: END; """ - def _memory_index_specs(self) -> "list[tuple[str, str, str]]": + def _memory_index_specs(self) -> list[tuple[str, str, str]]: """Return ``(index_name, table, columns)`` specs for memory-table indexes.""" return [ ( @@ -574,20 +570,19 @@ def _memory_index_specs(self) -> "list[tuple[str, str, str]]": (f"idx_{self._memory_table}_timestamp", self._memory_table, "timestamp DESC"), ] - def _drop_memory_table_sql(self) -> "list[str]": + def _drop_memory_table_sql(self) -> list[str]: return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"] - def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: - cursor.execute(sql, params) - return list(cursor.fetchall()) + def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + with self._config.provide_session() as driver: + return driver.select(sql, params) - def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: - with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: - cursor.execute(sql, params) - rowcount = _cursor_rowcount(cursor) + def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + with self._config.provide_session() as driver: + res = driver.execute(sql, params) + rowcount = res.rows_affected if commit: - conn.commit() + driver.commit() return rowcount @@ -601,20 +596,20 @@ def _adk_config(config: Any) -> PymssqlADKConfig: return cast("PymssqlADKConfig", adk_config) -def _configured_json_column_type(native_json: "bool | None") -> "str | None": +def _configured_json_column_type(native_json: bool | None) -> str | None: if native_json is True: return JSON_NATIVE_COLUMN_TYPE return JSON_FALLBACK_COLUMN_TYPE -def _json_column_type_from_sync_driver(driver: "PymssqlDriver") -> str: +def _json_column_type_from_sync_driver(driver: PymssqlDriver) -> str: version_info = driver.data_dictionary.get_version(driver) if isinstance(version_info, MssqlVersionInfo) and version_info.supports_native_json(): return JSON_NATIVE_COLUMN_TYPE return JSON_FALLBACK_COLUMN_TYPE -def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: "str | None") -> str: +def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: str | None) -> str: owner_line = f",\n {owner_id_column_ddl}" if owner_id_column_ddl else "" return f""" IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo')) @@ -634,7 +629,7 @@ def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: """ -def _sessions_index_specs(table: str) -> "list[tuple[str, str, str]]": +def _sessions_index_specs(table: str) -> list[tuple[str, str, str]]: return [ (f"idx_{table}_app_user", table, "app_name, user_id"), (f"idx_{table}_update_time", table, "update_time DESC"), @@ -663,7 +658,7 @@ def _events_table_ddl(table: str, session_table: str, json_column_type: str) -> """ -def _events_index_specs(table: str) -> "list[tuple[str, str, str]]": +def _events_index_specs(table: str) -> list[tuple[str, str, str]]: return [ (f"idx_{table}_scope", table, "app_name, user_id, session_id, timestamp ASC"), (f"idx_{table}_session", table, "session_id, timestamp ASC"), @@ -719,7 +714,7 @@ def _create_index_sql(table: str, index_name: str, columns: str) -> str: return f"CREATE INDEX {_quote_identifier(index_name)} ON {_table_ref(table)} ({columns})" -def _casefold_names(rows: "list[Any]", key: str) -> "set[str]": +def _casefold_names(rows: list[Any], key: str) -> set[str]: """Collapse data-dictionary rows into a case-folded, schema-stripped name set.""" return {str(row.get(key, "")).rsplit(".", 1)[-1].casefold() for row in rows} @@ -738,7 +733,7 @@ def _insert_event_sql(table: str) -> str: """ -def _upsert_state_sql(table: str, key_columns: "tuple[str, ...]", key_params: "tuple[str, ...]") -> str: +def _upsert_state_sql(table: str, key_columns: tuple[str, ...], key_params: tuple[str, ...]) -> str: source_columns = ", ".join( f"{param} AS {_quote_identifier(column)}" for column, param in zip(key_columns, key_params, strict=False) ) @@ -774,8 +769,8 @@ def _upsert_metadata_sql(table: str) -> str: def _events_query( - table: str, app_name: str, user_id: str, session_id: str, after_timestamp: "datetime | None", limit: "int | None" -) -> "tuple[str, tuple[Any, ...]]": + table: str, app_name: str, user_id: str, session_id: str, after_timestamp: datetime | None, limit: int | None +) -> tuple[str, tuple[Any, ...]]: top_clause = "TOP (%s) " if limit is not None else "" params: list[Any] = [limit] if limit is not None else [] params.extend([app_name, user_id, session_id]) @@ -792,7 +787,7 @@ def _events_query( return sql, tuple(params) -def _event_insert_params(event_record: StoredEvent) -> "tuple[Any, ...]": +def _event_insert_params(event_record: StoredEvent) -> tuple[Any, ...]: return ( event_record["id"], event_record["app_name"], @@ -822,7 +817,7 @@ def _event_record_from_row(row: Any) -> StoredEvent: ) -def _memory_record_from_row(row: Any) -> "StoredMemory": +def _memory_record_from_row(row: Any) -> StoredMemory: return cast( "StoredMemory", { @@ -843,7 +838,7 @@ def _memory_record_from_row(row: Any) -> "StoredMemory": ) -def _json_dict(value: Any) -> "dict[str, Any]": +def _json_dict(value: Any) -> dict[str, Any]: if value is None: return {} if isinstance(value, dict): @@ -855,28 +850,13 @@ def _json_dict(value: Any) -> "dict[str, Any]": return cast("dict[str, Any]", from_json(str(value))) -def _cursor_rowcount(cursor: Any) -> int: - rowcount = getattr(cursor, "rowcount", 0) - return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 - - def _is_mssql_table_missing(exc: BaseException) -> bool: text = str(exc).lower() - return "invalid object name" in text or _mssql_error_number(exc) == MSSQL_TABLE_NOT_FOUND_ERROR - - -def _mssql_error_number(exc: BaseException) -> "int | None": - matches = MSSQL_ERROR_NUMBER_PATTERN.findall(str(exc)) - if not matches: - return None - try: - return int(matches[-1]) - except ValueError: - return None + return "invalid object name" in text or extract_error_number(exc) == MSSQL_TABLE_NOT_FOUND_ERROR def _quote_identifier(identifier: str) -> str: - return f"[{identifier.replace(']', ']]')}]" + return quote_tsql_identifier(identifier) def _table_ref(table: str) -> str: @@ -907,14 +887,8 @@ def _build_mssql_scope_where( def _session_list_query( - session_table: str, - app_name: str, - user_id: "str | None", - column: str, - direction: str, - limit: "int | None", - offset: int, -) -> "tuple[str, tuple[Any, ...]]": + session_table: str, app_name: str, user_id: str | None, column: str, direction: str, limit: int | None, offset: int +) -> tuple[str, tuple[Any, ...]]: """Return the bounded session-list query and its bound values.""" params: list[Any] = [app_name] where_clause = "app_name = %s" diff --git a/sqlspec/adapters/pymssql/config.py b/sqlspec/adapters/pymssql/config.py index a2e1691da..c35fa66ce 100644 --- a/sqlspec/adapters/pymssql/config.py +++ b/sqlspec/adapters/pymssql/config.py @@ -1,7 +1,7 @@ """pymssql database configuration.""" from collections.abc import Callable, Mapping -from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast +from typing import Any, ClassVar, Literal, TypedDict, cast from typing_extensions import NotRequired @@ -11,15 +11,12 @@ from sqlspec.adapters.pymssql.migrations import PymssqlSyncMigrationTracker from sqlspec.adapters.pymssql.pool import PymssqlConnectionPool from sqlspec.config import ExtensionConfigs, SyncDatabaseConfig -from sqlspec.core import TypeCoercionCapabilities +from sqlspec.core import StatementConfig, TypeCoercionCapabilities from sqlspec.driver import SyncPoolConnectionContext, SyncPoolSessionFactory from sqlspec.extensions.events import EventRuntimeHints +from sqlspec.observability import ObservabilityConfig from sqlspec.utils.config_tools import normalize_connection_config -if TYPE_CHECKING: - from sqlspec.core import StatementConfig - from sqlspec.observability import ObservabilityConfig - __all__ = ("PymssqlConfig", "PymssqlConnectionParams", "PymssqlDriverFeatures", "PymssqlPoolParams", "PymssqlTimeout") PymssqlTimeout = int | float @@ -42,13 +39,14 @@ class PymssqlConnectionParams(TypedDict): conn_properties: NotRequired[str] autocommit: NotRequired[bool] tds_version: NotRequired[str] + encryption: NotRequired[Literal["default", "off", "request", "require"]] use_datetime2: NotRequired[bool] arraysize: NotRequired[int] conv: NotRequired[Mapping[int | type[Any], Callable[..., Any]]] read_only: NotRequired[bool] pool_recycle_seconds: NotRequired[int] health_check_interval: NotRequired[float] - extra: NotRequired["dict[str, Any]"] + extra: NotRequired[dict[str, Any]] class PymssqlPoolParams(PymssqlConnectionParams): @@ -69,9 +67,9 @@ class PymssqlDriverFeatures(TypedDict): events_backend: Event channel backend selection. """ - json_serializer: NotRequired["Callable[[Any], str]"] - json_deserializer: NotRequired["Callable[[str], Any]"] - on_connection_create: "NotRequired[Callable[[PymssqlConnection], None]]" + json_serializer: NotRequired[Callable[[Any], str]] + json_deserializer: NotRequired[Callable[[str], Any]] + on_connection_create: NotRequired[Callable[[PymssqlConnection], None]] enable_events: NotRequired[bool] events_backend: NotRequired[Literal["poll_queue"]] @@ -89,35 +87,37 @@ class _PymssqlSessionConnectionHandler(SyncPoolSessionFactory): class PymssqlConfig(SyncDatabaseConfig[PymssqlConnection, PymssqlConnectionPool, PymssqlDriver]): """Configuration for pymssql synchronous connections.""" - driver_type: "ClassVar[type[PymssqlDriver]]" = PymssqlDriver - connection_type: "ClassVar[type[PymssqlConnection]]" = cast("type[PymssqlConnection]", PymssqlConnection) - migration_tracker_type: "ClassVar[type[PymssqlSyncMigrationTracker]]" = PymssqlSyncMigrationTracker - supports_transactional_ddl: "ClassVar[bool]" = True - supports_migration_schemas: "ClassVar[bool]" = True - supports_native_arrow_export: "ClassVar[bool]" = False - supports_native_arrow_import: "ClassVar[bool]" = False - supports_native_parquet_export: "ClassVar[bool]" = False - supports_native_parquet_import: "ClassVar[bool]" = False - supports_native_row_streaming: "ClassVar[bool]" = True - type_coercion_capabilities: "ClassVar[TypeCoercionCapabilities]" = TypeCoercionCapabilities( + __slots__ = ("_user_connection_hook",) + + driver_type: ClassVar[type[PymssqlDriver]] = PymssqlDriver + connection_type: ClassVar[type[PymssqlConnection]] = cast("type[PymssqlConnection]", PymssqlConnection) + migration_tracker_type: ClassVar[type[PymssqlSyncMigrationTracker]] = PymssqlSyncMigrationTracker + supports_transactional_ddl: ClassVar[bool] = True + supports_migration_schemas: ClassVar[bool] = True + supports_native_arrow_export: ClassVar[bool] = False + supports_native_arrow_import: ClassVar[bool] = True + supports_native_parquet_export: ClassVar[bool] = False + supports_native_parquet_import: ClassVar[bool] = False + supports_native_row_streaming: ClassVar[bool] = True + type_coercion_capabilities: ClassVar[TypeCoercionCapabilities] = TypeCoercionCapabilities( datetime_binding="native", timestamp_precision="microsecond", json_columns_decoded=False, uuid_binding="text" ) - _connection_context_class: "ClassVar[type[PymssqlConnectionContext]]" = PymssqlConnectionContext - _session_factory_class: "ClassVar[type[_PymssqlSessionConnectionHandler]]" = _PymssqlSessionConnectionHandler - _session_context_class: "ClassVar[type[PymssqlSessionContext]]" = PymssqlSessionContext + _connection_context_class: ClassVar[type[PymssqlConnectionContext]] = PymssqlConnectionContext + _session_factory_class: ClassVar[type[_PymssqlSessionConnectionHandler]] = _PymssqlSessionConnectionHandler + _session_context_class: ClassVar[type[PymssqlSessionContext]] = PymssqlSessionContext _default_statement_config = default_statement_config def __init__( self, *, - connection_config: "PymssqlPoolParams | dict[str, Any] | None" = None, - connection_instance: "PymssqlConnectionPool | None" = None, - migration_config: "dict[str, Any] | None" = None, - statement_config: "StatementConfig | None" = None, - driver_features: "PymssqlDriverFeatures | dict[str, Any] | None" = None, - bind_key: "str | None" = None, - extension_config: "ExtensionConfigs | None" = None, - observability_config: "ObservabilityConfig | None" = None, + connection_config: PymssqlPoolParams | dict[str, Any] | None = None, + connection_instance: PymssqlConnectionPool | None = None, + migration_config: dict[str, Any] | None = None, + statement_config: StatementConfig | None = None, + driver_features: PymssqlDriverFeatures | dict[str, Any] | None = None, + bind_key: str | None = None, + extension_config: ExtensionConfigs | None = None, + observability_config: ObservabilityConfig | None = None, **kwargs: Any, ) -> None: connection_config = build_connection_config(normalize_connection_config(connection_config)) @@ -142,7 +142,7 @@ def __init__( **kwargs, ) - def _create_pool(self) -> "PymssqlConnectionPool": + def _create_pool(self) -> PymssqlConnectionPool: config = dict(self.connection_config) pool_recycle = config.pop("pool_recycle_seconds", 86400) health_check = config.pop("health_check_interval", 30.0) @@ -158,7 +158,7 @@ def _close_pool(self) -> None: self.connection_instance.close() self.connection_instance = None - def create_connection(self) -> "PymssqlConnection": + def create_connection(self) -> PymssqlConnection: """Open a standalone connection owned by the caller. The connection carries the same parameters and creation hook the pool @@ -170,7 +170,7 @@ def create_connection(self) -> "PymssqlConnection": """ return self.provide_pool().new_connection() - def get_signature_namespace(self) -> "dict[str, Any]": + def get_signature_namespace(self) -> dict[str, Any]: namespace = super().get_signature_namespace() namespace.update({ "PymssqlConfig": PymssqlConfig, @@ -188,6 +188,6 @@ def get_signature_namespace(self) -> "dict[str, Any]": }) return namespace - def get_event_runtime_hints(self) -> "EventRuntimeHints": + def get_event_runtime_hints(self) -> EventRuntimeHints: """Return runtime hints for pymssql event channels.""" return EventRuntimeHints(poll_interval=0.25, lease_seconds=5) diff --git a/sqlspec/adapters/pymssql/core.py b/sqlspec/adapters/pymssql/core.py index a0729717b..cab399ae4 100644 --- a/sqlspec/adapters/pymssql/core.py +++ b/sqlspec/adapters/pymssql/core.py @@ -1,8 +1,9 @@ """pymssql adapter compiled helpers.""" import re -from collections.abc import Callable, Sized -from typing import TYPE_CHECKING, Any, Final, Literal +from collections.abc import Callable, Mapping, Sequence, Sized +from logging import Logger +from typing import Any, Final, Literal from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile from sqlspec.exceptions import ( @@ -24,23 +25,22 @@ from sqlspec.utils.type_converters import build_uuid_coercions from sqlspec.utils.type_guards import has_rowcount -if TYPE_CHECKING: - from collections.abc import Mapping, Sequence - from logging import Logger - __all__ = ( "apply_driver_features", "build_connection_config", "build_insert_statement", + "build_multi_row_insert", "build_profile", "build_statement_config", "collect_rows", "create_mapped_exception", "default_statement_config", "driver_profile", + "extract_error_number", "format_identifier", "normalize_execute_many_parameters", "normalize_execute_parameters", + "quote_tsql_identifier", "resolve_column_names", "resolve_many_rowcount", "resolve_rowcount", @@ -63,6 +63,14 @@ } +def quote_tsql_identifier(identifier: str) -> str: + """Quote a T-SQL identifier with square brackets.""" + cleaned = identifier.strip() + if cleaned.startswith("[") and cleaned.endswith("]"): + cleaned = cleaned[1:-1].replace("]]", "]") + return f"[{cleaned.replace(']', ']]')}]" + + def format_identifier(identifier: str) -> str: """Format a T-SQL identifier with bracket quoting.""" cleaned = identifier.strip() @@ -70,20 +78,39 @@ def format_identifier(identifier: str) -> str: msg = "Table name must not be empty" raise SQLSpecError(msg) parts = split_qualified_identifier(cleaned, quote_chars='"', allow_bracket_quotes=True) - return ".".join(_quote_bracket_identifier(part) for part in parts) + return ".".join(quote_tsql_identifier(part) for part in parts) -def build_insert_statement(table: str, columns: "list[str]") -> str: +def build_insert_statement(table: str, columns: list[str]) -> str: """Build a pymssql-compatible INSERT statement.""" - column_clause = ", ".join(_quote_bracket_identifier(column) for column in columns) + column_clause = ", ".join(quote_tsql_identifier(column) for column in columns) placeholders = ", ".join("%s" for _ in columns) return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES ({placeholders})" +def build_multi_row_insert(table: str, columns: list[str], num_rows: int) -> str: + """Build a multi-row VALUES (...), (...) batch INSERT statement. + + Args: + table: Target table name. + columns: Column names to insert. + num_rows: Number of row tuples in the VALUES clause (up to 1,000). + + Returns: + Parameterized T-SQL INSERT statement. + """ + column_clause = ", ".join(quote_tsql_identifier(column) for column in columns) + single_row = f"({', '.join('%s' for _ in columns)})" + values_clause = ", ".join(single_row for _ in range(num_rows)) + return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES {values_clause}" + + def normalize_execute_parameters(parameters: Any) -> Any: """Normalize parameters for pymssql execute calls.""" if parameters is None: return None + if isinstance(parameters, tuple): + return parameters if isinstance(parameters, list): return tuple(parameters) return parameters @@ -94,7 +121,7 @@ def normalize_execute_many_parameters(parameters: Any) -> Any: return parameters -def build_profile() -> "DriverParameterProfile": +def build_profile() -> DriverParameterProfile: """Create the pymssql driver parameter profile.""" return DriverParameterProfile( name="pymssql", @@ -114,8 +141,8 @@ def build_profile() -> "DriverParameterProfile": def build_statement_config( - *, json_serializer: "Callable[[Any], str] | None" = None, json_deserializer: "Callable[[str], Any] | None" = None -) -> "StatementConfig": + *, json_serializer: Callable[[Any], str] | None = None, json_deserializer: Callable[[str], Any] | None = None +) -> StatementConfig: """Construct the pymssql statement configuration.""" return build_statement_config_from_profile( driver_profile, @@ -126,8 +153,8 @@ def build_statement_config( def apply_driver_features( - statement_config: "StatementConfig", driver_features: "Mapping[str, Any] | None" -) -> "tuple[StatementConfig, dict[str, Any]]": + statement_config: StatementConfig, driver_features: Mapping[str, Any] | None +) -> tuple[StatementConfig, dict[str, Any]]: """Apply pymssql driver feature defaults to statement config.""" features: dict[str, Any] = dict(driver_features) if driver_features else {} json_serializer = features.setdefault("json_serializer", to_json) @@ -142,9 +169,9 @@ def apply_driver_features( return statement_config, features -def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) -> SQLSpecError: +def create_mapped_exception(error: Exception, *, logger: Logger | None = None) -> SQLSpecError: """Map a pymssql exception to SQLSpec's exception hierarchy.""" - error_number = _extract_error_number(error) + error_number = extract_error_number(error) if error_number == _MSSQL_CONSTRAINT_547: message = str(error) if "check constraint" in message.lower(): @@ -175,8 +202,8 @@ def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) def resolve_column_names( - description: "Sequence[Any] | None", column_name_cache: "dict[int, tuple[Any, list[str]]] | None" = None -) -> "list[str]": + description: Sequence[Any] | None, column_name_cache: dict[int, tuple[Any, list[str]]] | None = None +) -> list[str]: """Resolve ordered column names from cursor metadata.""" if not description: return [] @@ -194,17 +221,18 @@ def resolve_column_names( def collect_rows( - fetched_data: "Sequence[Any] | None", - description: "Sequence[Any] | None", - column_name_cache: "dict[int, tuple[Any, list[str]]] | None" = None, -) -> "tuple[list[Any], list[str], Literal['dict', 'tuple', 'record']]": + fetched_data: Sequence[Any] | None, + description: Sequence[Any] | None, + column_name_cache: dict[int, tuple[Any, list[str]]] | None = None, +) -> tuple[list[Any], list[str], Literal["dict", "tuple", "record"]]: """Collect pymssql rows, preserving dictionary or tuple row shape.""" column_names = resolve_column_names(description, column_name_cache) if not fetched_data: return [], column_names, "tuple" - if isinstance(fetched_data[0], dict): - return list(fetched_data), column_names, "dict" - return list(fetched_data), column_names, "tuple" + rows = fetched_data if isinstance(fetched_data, list) else list(fetched_data) + if isinstance(rows[0], dict): + return rows, column_names, "dict" + return rows, column_names, "tuple" def resolve_rowcount(cursor: Any) -> int: @@ -217,7 +245,7 @@ def resolve_rowcount(cursor: Any) -> int: return 0 -def resolve_many_rowcount(cursor: Any, parameters: Any, *, fallback_count: "int | None" = None) -> int: +def resolve_many_rowcount(cursor: Any, parameters: Any, *, fallback_count: int | None = None) -> int: """Resolve executemany rowcount using cursor metadata with payload fallback.""" rowcount = resolve_rowcount(cursor) if rowcount > 0: @@ -233,7 +261,7 @@ def _bool_to_int(value: bool) -> int: return int(value) -def _constraint_exception_from_message(error: Exception) -> "SQLSpecError | None": +def _constraint_exception_from_message(error: Exception) -> SQLSpecError | None: """Classify SQL Server constraint messages when a driver omits the native error number.""" message = str(error) normalized = message.lower() @@ -255,28 +283,39 @@ def _custom_type_coercions() -> dict[type, Callable[[Any], Any]]: return coercions -def _quote_bracket_identifier(identifier: str) -> str: - cleaned = identifier.strip() - if cleaned.startswith("[") and cleaned.endswith("]"): - cleaned = cleaned[1:-1].replace("]]", "]") - return f"[{cleaned.replace(']', ']]')}]" - +def extract_error_number(exc: BaseException | None) -> int | None: + """Extract integer SQL Server error code from an exception if present. -def _extract_error_number(exc: Exception) -> "int | None": - matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) - if not matches: - return None - try: - return int(matches[-1]) - except ValueError: + Checks native integer attributes (number, error_code, errno) or args[0] before regex. + """ + if exc is None: return None + for attr in ("number", "error_code", "errno"): + val = getattr(exc, attr, None) + if isinstance(val, int) and val != 0: + return val + if hasattr(exc, "args") and exc.args: + first = exc.args[0] + if isinstance(first, int): + return first + matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) + if matches: + try: + return int(matches[-1]) + except ValueError: + pass + return None + + +_extract_error_number = extract_error_number +_quote_bracket_identifier = quote_tsql_identifier driver_profile = build_profile() default_statement_config = build_statement_config() -def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str, Any]": +def build_connection_config(connection_config: Mapping[str, Any]) -> dict[str, Any]: """Build a normalized connection configuration dictionary. Args: diff --git a/sqlspec/adapters/pymssql/data_dictionary.py b/sqlspec/adapters/pymssql/data_dictionary.py index dc39b562e..eac2875ec 100644 --- a/sqlspec/adapters/pymssql/data_dictionary.py +++ b/sqlspec/adapters/pymssql/data_dictionary.py @@ -1,14 +1,19 @@ """pymssql data dictionary.""" -from typing import TYPE_CHECKING, Any, ClassVar, cast +from collections.abc import Sequence +from typing import Any, ClassVar, Final, cast from mypy_extensions import mypyc_attr +from sqlspec.adapters.pymssql.driver import PymssqlDriver +from sqlspec.core import SQL from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, + DialectConfig, ForeignKeyMetadata, IndexMetadata, + MetadataCapabilityProfile, MetadataSupport, SystemMetadataCapability, SystemMetadataRequest, @@ -39,17 +44,12 @@ from sqlspec.driver import SyncDataDictionaryBase from sqlspec.utils.logging import get_logger -if TYPE_CHECKING: - from collections.abc import Sequence - - from sqlspec.adapters.pymssql.driver import PymssqlDriver - from sqlspec.core import SQL - from sqlspec.data_dictionary._types import DialectConfig, MetadataCapabilityProfile - __all__ = ("MssqlVersionInfo", "PymssqlSyncDataDictionary") logger = get_logger("sqlspec.adapters.pymssql.data_dictionary") +MSSQL_VECTOR_MIN_MAJOR: Final[int] = 17 + class MssqlVersionInfo(VersionInfo): """MSSQL database version info with build, revision, and Azure SQL detection.""" @@ -74,8 +74,12 @@ def supports_native_json(self) -> bool: """Return whether this server supports the native JSON type.""" return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) + def supports_vector(self) -> bool: + """Return whether this server supports native VECTOR data types and functions.""" + return self.is_azure_sql or self.major >= MSSQL_VECTOR_MIN_MAJOR + @property - def version_tuple(self) -> "tuple[int, int, int]": + def version_tuple(self) -> tuple[int, int, int]: """Get version tuple using the MSSQL build number as the third component.""" return (self.major, self.minor, self.build) @@ -94,7 +98,7 @@ class _MssqlDataDictionaryMixin: dialect: ClassVar[str] = "mssql" - def get_dialect_config(self) -> "DialectConfig": + def get_dialect_config(self) -> DialectConfig: """Return the dialect configuration for this data dictionary.""" return get_dialect_config(type(self).dialect) @@ -109,7 +113,7 @@ def list_available_features(self) -> list[str]: """List available feature flags for this dialect.""" return list_mssql_available_features(self.get_dialect_config()) - def get_domain_query(self, domain: str, name: str) -> "SQL": + def get_domain_query(self, domain: str, name: str) -> SQL: """Return a SQL Server domain query.""" query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, name) return cast("SQL", query.sql) @@ -132,6 +136,8 @@ def _build_version_info( def _get_optimal_type_from_version(self, version_info: MssqlVersionInfo | None, type_category: str) -> str: if type_category in {"json", "jsonb"} and version_info is not None and version_info.supports_native_json(): return "JSON" + if type_category == "vector" and version_info is not None and version_info.supports_vector(): + return "VECTOR" return self.get_dialect_config().get_optimal_type(type_category) @@ -145,20 +151,20 @@ def __init__(self) -> None: super().__init__() def get_metadata_capabilities( - self, driver: "PymssqlDriver", domains: "Sequence[str] | None" = None - ) -> "MetadataCapabilityProfile": + self, driver: PymssqlDriver, domains: Sequence[str] | None = None + ) -> MetadataCapabilityProfile: """Get SQL Server data-dictionary capability profile.""" return build_mssql_metadata_capability_profile(type(self).__name__, domains) def get_system_metadata_capabilities( - self, driver: "PymssqlDriver", domains: "Sequence[str] | None" = None + self, driver: PymssqlDriver, domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get SQL Server opt-in system metadata capability disclosures.""" _ = driver requested_domains = ("dmv_exec_requests", "query_store_runtime") if domains is None else tuple(domains) return tuple(build_mssql_system_metadata_capability(domain) for domain in requested_domains) - def get_version(self, driver: "PymssqlDriver") -> MssqlVersionInfo | None: + def get_version(self, driver: PymssqlDriver) -> MssqlVersionInfo | None: """Get SQL Server version information.""" driver_id = id(driver) if driver_id in self._version_fetch_attempted: @@ -185,9 +191,11 @@ def get_version(self, driver: "PymssqlDriver") -> MssqlVersionInfo | None: self.cache_version(driver_id, version_info) return version_info - def get_feature_flag(self, driver: "PymssqlDriver", feature: str) -> bool: + def get_feature_flag(self, driver: PymssqlDriver, feature: str) -> bool: """Check whether SQL Server supports a feature.""" version_info = self.get_version(driver) + if feature == "supports_vector": + return bool(version_info and version_info.supports_vector()) return resolve_mssql_feature_flag( feature, major=version_info.major if version_info is not None else 0, @@ -196,11 +204,11 @@ def get_feature_flag(self, driver: "PymssqlDriver", feature: str) -> bool: version_info=version_info, ) - def get_optimal_type(self, driver: "PymssqlDriver", type_category: str) -> str: + def get_optimal_type(self, driver: PymssqlDriver, type_category: str) -> str: """Get optimal SQL Server type for a category.""" return self._get_optimal_type_from_version(self.get_version(driver), type_category) - def get_tables(self, driver: "PymssqlDriver", schema: str | None = None) -> list[TableMetadata]: + def get_tables(self, driver: PymssqlDriver, schema: str | None = None) -> list[TableMetadata]: """Get tables sorted by dependency order with catalog fallback.""" schema_name = self.resolve_connection_schema(driver, schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -219,7 +227,7 @@ def get_tables(self, driver: "PymssqlDriver", schema: str | None = None) -> list return merge_mssql_table_lists(ordered, all_rows) def get_columns( - self, driver: "PymssqlDriver", table: str | None = None, schema: str | None = None + self, driver: PymssqlDriver, table: str | None = None, schema: str | None = None ) -> list[ColumnMetadata]: """Get column information for a table or schema.""" schema_name = self.resolve_connection_schema(driver, schema) @@ -243,7 +251,7 @@ def get_columns( ) def get_indexes( - self, driver: "PymssqlDriver", table: str | None = None, schema: str | None = None + self, driver: PymssqlDriver, table: str | None = None, schema: str | None = None ) -> list[IndexMetadata]: """Get index metadata for a table or schema.""" schema_name = self.resolve_connection_schema(driver, schema) @@ -267,7 +275,7 @@ def get_indexes( ) def get_foreign_keys( - self, driver: "PymssqlDriver", table: str | None = None, schema: str | None = None + self, driver: PymssqlDriver, table: str | None = None, schema: str | None = None ) -> list[ForeignKeyMetadata]: """Get foreign key metadata.""" schema_name = self.resolve_connection_schema(driver, schema) @@ -294,7 +302,7 @@ def get_foreign_keys( def get_ddl( self, - driver: "PymssqlDriver", + driver: PymssqlDriver, object_name: str, schema: str | None = None, *, @@ -315,7 +323,7 @@ def get_ddl( return build_mssql_table_ddl_result(schema_name, object_name, columns, indexes, object_type=object_type) def get_system_metadata( - self, driver: "PymssqlDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: PymssqlDriver, request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in SQL Server system metadata with sensitive columns redacted by default.""" metadata_request = ensure_system_metadata_request(request, **kwargs) diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 0655b8052..98920df5f 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -1,9 +1,11 @@ """pymssql SQL Server driver implementation.""" import contextlib -from collections.abc import Sized +from collections.abc import Iterable, Sequence, Sized from typing import TYPE_CHECKING, Any, cast +import sqlglot.expressions as exp + from sqlspec.adapters.pymssql._typing import ( PymssqlConnection, PymssqlCursor, @@ -12,18 +14,22 @@ PymssqlSessionContext, ) from sqlspec.adapters.pymssql.core import ( + build_multi_row_insert, collect_rows, create_mapped_exception, default_statement_config, driver_profile, + format_identifier, normalize_execute_many_parameters, normalize_execute_parameters, + quote_tsql_identifier, resolve_column_names, resolve_many_rowcount, resolve_rowcount, ) from sqlspec.adapters.pymssql.data_dictionary import PymssqlSyncDataDictionary -from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile +from sqlspec.core import SQL, ArrowResult, StatementConfig, get_cache_config, register_driver_profile +from sqlspec.core.result import DMLResult, SQLResult from sqlspec.driver import ( BaseSyncExceptionHandler, ExecutionResult, @@ -33,12 +39,14 @@ validate_savepoint_name, ) from sqlspec.exceptions import SQLSpecError +from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.utils.logging import get_logger if TYPE_CHECKING: - from collections.abc import Sequence - from sqlspec.adapters.pymssql._typing import PymssqlQueryParams as QueryParams + from sqlspec.builder import QueryBuilder + from sqlspec.core import Statement, StatementFilter + from sqlspec.typing import StatementParameters __all__ = ("PymssqlCursor", "PymssqlDriver", "PymssqlExceptionHandler", "PymssqlSessionContext") @@ -50,7 +58,7 @@ class PymssqlExceptionHandler(BaseSyncExceptionHandler): __slots__ = () - def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool: + def _handle_exception(self, exc_type: type[BaseException] | None, exc_val: BaseException) -> bool: if exc_type is None: return False if isinstance(exc_val, PymssqlError): @@ -86,7 +94,7 @@ def start(self) -> None: raise self._cursor_manager = cursor_manager - def fetch_chunk(self) -> "list[dict[str, Any]]": + def fetch_chunk(self) -> list[dict[str, Any]]: cursor_manager = self._cursor_manager if cursor_manager is None or cursor_manager.cursor is None: return [] @@ -126,9 +134,9 @@ class PymssqlDriver(SyncDriverAdapterBase): def __init__( self, - connection: "PymssqlConnection", - statement_config: "StatementConfig | None" = None, - driver_features: "dict[str, Any] | None" = None, + connection: PymssqlConnection, + statement_config: StatementConfig | None = None, + driver_features: dict[str, Any] | None = None, ) -> None: if statement_config is None: statement_config = default_statement_config.replace( @@ -142,7 +150,7 @@ def __init__( self._transaction_active = False self._explicit_transaction = False - def dispatch_execute(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "ExecutionResult": + def dispatch_execute(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) cursor.execute(sql, normalize_execute_parameters(prepared_parameters)) @@ -161,7 +169,7 @@ def dispatch_execute(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "Exe return self.create_execution_result(cursor, rowcount_override=resolve_rowcount(cursor)) - def dispatch_execute_many(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "ExecutionResult": + def dispatch_execute_many(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) prepared_parameters = normalize_execute_many_parameters(prepared_parameters) @@ -171,7 +179,7 @@ def dispatch_execute_many(self, cursor: "PymssqlRawCursor", statement: "SQL") -> affected_rows = resolve_many_rowcount(cursor, prepared_parameters, fallback_count=parameter_count) return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True) - def dispatch_execute_script(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "ExecutionResult": + def dispatch_execute_script(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True) @@ -228,13 +236,13 @@ def rollback(self) -> None: msg = f"Failed to rollback SQL Server transaction: {exc}" raise SQLSpecError(msg) from exc - def with_cursor(self, connection: "PymssqlConnection") -> "PymssqlCursor": + def with_cursor(self, connection: PymssqlConnection) -> PymssqlCursor: return PymssqlCursor(connection) - def handle_database_exceptions(self) -> "PymssqlExceptionHandler": + def handle_database_exceptions(self) -> PymssqlExceptionHandler: return PymssqlExceptionHandler() - def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "SyncRowStream[dict[str, Any]] | None": + def dispatch_select_stream(self, statement: SQL, chunk_size: int) -> SyncRowStream[dict[str, Any]] | None: """Return a native pymssql row stream backed by ``fetchmany()``.""" if not statement.returns_rows(): return None @@ -281,30 +289,199 @@ def has_schema(self, schema: str) -> bool: return cursor.fetchone() is not None @property - def data_dictionary(self) -> "PymssqlSyncDataDictionary": + def data_dictionary(self) -> PymssqlSyncDataDictionary: if self._data_dictionary is None: self._data_dictionary = PymssqlSyncDataDictionary() return self._data_dictionary - def collect_rows(self, cursor: "PymssqlRawCursor", fetched: "list[Any]") -> "tuple[list[Any], list[str], int]": - column_names = resolve_column_names(cursor.description or None, self._column_name_cache) - return fetched, column_names, len(fetched) + def execute_many( + self, + statement: "SQL | Statement | QueryBuilder", + /, + parameters: "Sequence[StatementParameters]", + *filters: "StatementParameters | StatementFilter", + statement_config: StatementConfig | None = None, + **kwargs: Any, + ) -> SQLResult: + """Execute a statement across parameter sets with multi-row batching.""" + config = statement_config or self.statement_config + if isinstance(statement, str) and not filters and not kwargs and config is self.statement_config: + prepared_statement = SQL( + statement, + tuple(parameters) if isinstance(parameters, list) else parameters, + statement_config=config, + is_many=True, + ) + cached_statement, prepared_parameters = self._compiled_statement(prepared_statement, config) + parsed_expression = cached_statement.expression + if isinstance(parsed_expression, exp.Insert) and not parsed_expression.args.get("returning"): + bulk_result = self._execute_bulk_insert_many(parsed_expression, prepared_parameters) + if bulk_result is not None: + return bulk_result + return super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) + + def _execute_bulk_insert_many(self, expression: exp.Insert, prepared_parameters: Any) -> DMLResult | None: + """Execute a batch INSERT via multi-row VALUES chunking up to 1,000 rows.""" + if not isinstance(prepared_parameters, (list, tuple)) or not prepared_parameters: + return None + if not isinstance(expression.this, exp.Schema): + return None + if not _is_plain_values_insert(expression, len(expression.this.expressions)): + return None + if not isinstance(prepared_parameters[0], (list, tuple)): + return None - def resolve_rowcount(self, cursor: "PymssqlRawCursor") -> int: - return resolve_rowcount(cursor) + table_expr = expression.this.this + if not isinstance(table_expr, exp.Table) or table_expr.alias: + return None + + column_names = [column.name for column in expression.this.expressions] + target_table = table_expr.sql(dialect="tsql") + rows = prepared_parameters + total_affected = 0 + chunk_size = 1000 + + handler = self.handle_database_exceptions() + with handler, self.with_cursor(self.connection) as cursor: + for i in range(0, len(rows), chunk_size): + chunk = rows[i : i + chunk_size] + chunk_sql = build_multi_row_insert(target_table, column_names, len(chunk)) + flat_params: list[Any] = [] + for row in chunk: + flat_params.extend(row) + cursor.execute(chunk_sql, tuple(flat_params)) + total_affected += len(chunk) + self._check_pending_exception(handler) + return DMLResult("INSERT", total_affected) + + def bulk_copy( + self, + table_name: str, + rows: Sequence[Sequence[Any]] | Iterable[Sequence[Any]], + *, + column_ids: Sequence[int] | None = None, + batch_size: int = 1000, + tablock: bool = False, + check_constraints: bool = False, + fire_triggers: bool = False, + ) -> int: + """Perform high-performance bulk insert using FreeTDS BCP APIs. + + Args: + table_name: Target SQL Server table name. + rows: Sequence or iterable of row tuples/sequences. + column_ids: Optional 1-based column IDs mapping elements to table columns. + batch_size: Number of rows per batch commit. Defaults to 1000. + tablock: Apply TABLOCK hint for minimal logging. + check_constraints: Enforce table constraints during BCP. + fire_triggers: Execute insert triggers during BCP. + + Returns: + Number of rows ingested. + """ + row_list = list(rows) if not isinstance(rows, (list, tuple)) else rows + if not row_list: + return 0 + handler = self.handle_database_exceptions() + with handler: + self.connection.bulk_copy( + table_name, + row_list, + column_ids=list(column_ids) if column_ids is not None else None, + batch_size=batch_size, + tablock=tablock, + check_constraints=check_constraints, + fire_triggers=fire_triggers, + ) + self._check_pending_exception(handler) + return len(row_list) + + def load_from_arrow( + self, + table: str, + source: ArrowResult | Any, + *, + partitioner: dict[str, object] | None = None, + overwrite: bool = False, + telemetry: StorageTelemetry | None = None, + batch_size: int = 1000, + tablock: bool = False, + check_constraints: bool = False, + fire_triggers: bool = False, + column_ids: Sequence[int] | None = None, + ) -> StorageBridgeJob: + """Load Arrow data into SQL Server via FreeTDS BCP bulk copy.""" + self._require_capability("arrow_import_enabled") + if overwrite: + quoted_table = format_identifier(table) + handler = self.handle_database_exceptions() + with handler, self.with_cursor(self.connection) as cursor: + try: + cursor.execute(f"TRUNCATE TABLE {quoted_table}") + except Exception as exc: + error_msg = str(exc) + if "4712" in error_msg or "foreign key" in error_msg.lower(): + cursor.execute(f"DELETE FROM {quoted_table}") + else: + raise + self._check_pending_exception(handler) + + arrow_table = self._coerce_arrow_table(source) + if arrow_table.num_rows > 0: + for batch in arrow_table.to_batches(): + pydict = batch.to_pydict() + rows = list(zip(*pydict.values(), strict=False)) + self.bulk_copy( + table, + rows, + column_ids=column_ids, + batch_size=batch_size, + tablock=tablock, + check_constraints=check_constraints, + fire_triggers=fire_triggers, + ) + + telemetry_payload = self._ingest_telemetry(arrow_table) + extra = telemetry_payload.setdefault("extra", {}) + extra["rows_ingested"] = arrow_table.num_rows + telemetry_payload["rows_processed"] = arrow_table.num_rows + telemetry_payload["destination"] = table + self._attach_partition_telemetry(telemetry_payload, partitioner) + return self._storage_job(telemetry_payload, telemetry) + + def load_from_storage( + self, + table: str, + source: StorageDestination, + *, + file_format: StorageFormat, + partitioner: dict[str, object] | None = None, + overwrite: bool = False, + ) -> StorageBridgeJob: + """Load staged artifacts from storage into SQL Server via BCP.""" + arrow_table, inbound = self._read_storage_arrow(source, file_format=file_format) + return self.load_from_arrow(table, arrow_table, partitioner=partitioner, overwrite=overwrite, telemetry=inbound) def _connection_in_transaction(self) -> bool: """Return whether a transaction opened by this driver remains active.""" return self._transaction_active -def _quote_tsql_identifier(identifier: str) -> str: - """Bracket-quote an identifier so the statement is valid regardless of the session's QUOTED_IDENTIFIER setting.""" - return f"[{identifier.replace(']', ']]')}]" +def _is_plain_values_insert(expression: exp.Insert, expected_columns: int) -> bool: + values = expression.args.get("values") + if not isinstance(values, exp.Values): + return False + rows = values.expressions + if len(rows) != 1: + return False + row = rows[0] + if not isinstance(row, exp.Tuple): + return False + return len(row.expressions) == expected_columns def _alter_default_schema_sql(user_name: str, schema: str) -> str: - return f"ALTER USER {_quote_tsql_identifier(user_name)} WITH DEFAULT_SCHEMA = {_quote_tsql_identifier(schema)};" + return f"ALTER USER {quote_tsql_identifier(user_name)} WITH DEFAULT_SCHEMA = {quote_tsql_identifier(schema)};" register_driver_profile("pymssql", driver_profile) diff --git a/sqlspec/adapters/pymssql/events/store.py b/sqlspec/adapters/pymssql/events/store.py index 717cce409..276af2275 100644 --- a/sqlspec/adapters/pymssql/events/store.py +++ b/sqlspec/adapters/pymssql/events/store.py @@ -3,6 +3,7 @@ import re from sqlspec.adapters.pymssql.config import PymssqlConfig +from sqlspec.adapters.pymssql.core import quote_tsql_identifier from sqlspec.extensions.events import BaseEventQueueStore from sqlspec.utils.text import split_qualified_identifier @@ -69,8 +70,4 @@ def _split_table_name(table_name: str) -> tuple[str, str]: def _object_name(table_name: str) -> str: schema_name, bare_table_name = _split_table_name(table_name) - return f"{_quote_bracket_identifier(schema_name)}.{_quote_bracket_identifier(bare_table_name)}" - - -def _quote_bracket_identifier(identifier: str) -> str: - return f"[{identifier.replace(']', ']]')}]" + return f"{quote_tsql_identifier(schema_name)}.{quote_tsql_identifier(bare_table_name)}" diff --git a/sqlspec/adapters/pymssql/litestar/store.py b/sqlspec/adapters/pymssql/litestar/store.py index 77af491f1..14df07634 100644 --- a/sqlspec/adapters/pymssql/litestar/store.py +++ b/sqlspec/adapters/pymssql/litestar/store.py @@ -1,14 +1,12 @@ """pymssql Litestar Store implementation.""" from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any +from typing import Any +from sqlspec.adapters.pymssql.config import PymssqlConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ -if TYPE_CHECKING: - from sqlspec.adapters.pymssql.config import PymssqlConfig - __all__ = ("PymssqlStore",) @@ -17,7 +15,7 @@ class PymssqlStore(BaseSQLSpecStore["PymssqlConfig"]): __slots__ = () - def __init__(self, config: "PymssqlConfig") -> None: + def __init__(self, config: PymssqlConfig) -> None: super().__init__(config) async def create_table(self) -> None: @@ -28,11 +26,11 @@ async def create_table(self) -> None: await async_(self._create_table)() await self.reconcile_schema(assume_existing=True) - async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": + async def get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: """Get a session value by key.""" return await async_(self._get)(key, renew_for) - async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: + async def set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: """Store a session value.""" await async_(self._set)(key, value, expires_in) @@ -48,7 +46,7 @@ async def exists(self, key: str) -> bool: """Check if a session key exists and is not expired.""" return await async_(self._exists)(key) - async def expires_in(self, key: str) -> "int | None": + async def expires_in(self, key: str) -> int | None: """Get the time in seconds until the session expires.""" return await async_(self._expires_in)(key) @@ -80,7 +78,7 @@ def _table_ddl(self) -> str: END; """ - def _drop_table_sql(self) -> "list[str]": + def _drop_table_sql(self) -> list[str]: """Get SQL Server DROP TABLE statements.""" return [f"IF OBJECT_ID(N'dbo.{self._table_name}', N'U') IS NOT NULL DROP TABLE dbo.{self._table_name};"] @@ -90,20 +88,14 @@ def _create_table(self) -> None: driver.commit() self._log_table_created() - def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": + def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = %s AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql, (key,)) - row = cursor.fetchone() - finally: - cursor.close() - + with self._config.provide_session() as driver: + row = driver.select_one_or_none(sql, (key,)) if row is None: return None @@ -111,23 +103,19 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if renew_for is not None and expires_at is not None: new_expires_at = self._calculate_expires_at(renew_for) if new_expires_at is not None: - update_cursor = conn.cursor() - try: - update_cursor.execute( - f""" - UPDATE {self._table_name} - SET expires_at = %s, updated_at = SYSUTCDATETIME() - WHERE session_id = %s - """, - (new_expires_at, key), - ) - finally: - update_cursor.close() - conn.commit() + driver.execute( + f""" + UPDATE {self._table_name} + SET expires_at = %s, updated_at = SYSUTCDATETIME() + WHERE session_id = %s + """, + (new_expires_at, key), + ) + driver.commit() return _coerce_bytes(_row_value(row, "data", 0)) - def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: + def _set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) sql = f""" @@ -143,31 +131,19 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No INSERT (session_id, data, expires_at) VALUES (src.session_id, src.data, src.expires_at); """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql, (key, data, expires_at)) - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + driver.execute(sql, (key, data, expires_at)) + driver.commit() def _delete(self, key: str) -> None: - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"DELETE FROM {self._table_name} WHERE session_id = %s", (key,)) - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + driver.execute(f"DELETE FROM {self._table_name} WHERE session_id = %s", (key,)) + driver.commit() def _delete_all(self) -> None: - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"TRUNCATE TABLE {self._table_name}") - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + driver.execute(f"TRUNCATE TABLE {self._table_name}") + driver.commit() self._log_delete_all() def _exists(self, key: str) -> bool: @@ -177,22 +153,12 @@ def _exists(self, key: str) -> bool: WHERE session_id = %s AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql, (key,)) - return cursor.fetchone() is not None - finally: - cursor.close() - - def _expires_in(self, key: str) -> "int | None": - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"SELECT expires_at FROM {self._table_name} WHERE session_id = %s", (key,)) - row = cursor.fetchone() - finally: - cursor.close() + with self._config.provide_session() as driver: + return driver.select_one_or_none(sql, (key,)) is not None + + def _expires_in(self, key: str) -> int | None: + with self._config.provide_session() as driver: + row = driver.select_one_or_none(f"SELECT expires_at FROM {self._table_name} WHERE session_id = %s", (key,)) if row is None: return None @@ -208,14 +174,10 @@ def _delete_expired(self) -> int: WHERE expires_at IS NOT NULL AND expires_at < SYSUTCDATETIME() """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(sql) - count = int(getattr(cursor, "rowcount", 0) or 0) - finally: - cursor.close() - conn.commit() + with self._config.provide_session() as driver: + res = driver.execute(sql) + driver.commit() + count = res.rows_affected if count > 0: self._log_delete_expired(count) return count @@ -235,7 +197,7 @@ def _row_value(row: object, key: str, index: int) -> Any: return getattr(row, key, None) -def _normalize_utc(value: Any) -> "datetime | None": +def _normalize_utc(value: Any) -> datetime | None: if value is None: return None if not isinstance(value, datetime): diff --git a/tests/unit/adapters/test_mssql_python/test_config.py b/tests/unit/adapters/test_mssql_python/test_config.py index 0007a37fe..44c786830 100644 --- a/tests/unit/adapters/test_mssql_python/test_config.py +++ b/tests/unit/adapters/test_mssql_python/test_config.py @@ -346,3 +346,29 @@ def commit(self) -> None: raise RuntimeError assert calls == ["rollback", "release"] + + +def test_pool_close_calls_ddbc_close_pooling(monkeypatch: pytest.MonkeyPatch) -> None: + """Connection pool close should call ddbc_bindings.close_pooling when requested.""" + closed_pooling: list[bool] = [] + + class FakeBindings: + @staticmethod + def close_pooling() -> None: + closed_pooling.append(True) + + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.ddbc_bindings", FakeBindings) + pool = MssqlPythonConnectionPool(connection_string="Server=localhost;") + pool.close(close_driver_pooling=True) + assert closed_pooling == [True] + + +def test_pool_suppresses_warning_when_params_match(monkeypatch: pytest.MonkeyPatch) -> None: + """Pool reconfiguration should not warn if params are identical to previous.""" + monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", (10, 60, True)) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **kw: None) + + with warnings.catch_warnings(record=True) as recorded: + warnings.simplefilter("always") + MssqlPythonConnectionPool(connection_string="Server=localhost;", max_size=10, idle_timeout=60, enabled=True) + assert not any("Pooling configuration was already set" in str(w.message) for w in recorded) diff --git a/tests/unit/adapters/test_mssql_python/test_core.py b/tests/unit/adapters/test_mssql_python/test_core.py index 7ff15010e..d6a1a1f86 100644 --- a/tests/unit/adapters/test_mssql_python/test_core.py +++ b/tests/unit/adapters/test_mssql_python/test_core.py @@ -3,7 +3,7 @@ import pytest from sqlspec.adapters.mssql_python._typing import MSSQL_PYTHON_MODULE -from sqlspec.adapters.mssql_python.core import build_connection_config, create_mapped_exception +from sqlspec.adapters.mssql_python.core import build_connection_config, create_mapped_exception, extract_error_number from sqlspec.exceptions import ( CheckViolationError, DatabaseConnectionError, @@ -255,3 +255,24 @@ def test_parse_odbc_connection_string_edge_cases() -> None: assert parse_odbc_connection_string("Incomplete={no_close") == [("Incomplete", "{no_close")] assert parse_odbc_connection_string("Server=host; ") == [("Server", "host")] assert parse_odbc_connection_string("DanglingToken") == [] + + +def test_extract_error_number_from_attribute() -> None: + """extract_error_number retrieves native integer attribute 'number'.""" + + class CustomError(Exception): + number = 2627 + + assert extract_error_number(CustomError("duplicate key")) == 2627 + + +def test_extract_error_number_from_args_tuple() -> None: + """extract_error_number extracts integer from exception args.""" + assert extract_error_number(Exception(1205, "Deadlock found")) == 1205 + + +def test_extract_error_number_from_string_regex() -> None: + """extract_error_number parses error numbers formatted as (1205) or error 1205.""" + assert extract_error_number(Exception("Transaction was deadlocked on lock resources (1205)")) == 1205 + assert extract_error_number(Exception("Msg 4712, Level 16, State 1")) == 4712 + assert extract_error_number(Exception("No numbers here")) is None diff --git a/tests/unit/adapters/test_mssql_python/test_data_dictionary.py b/tests/unit/adapters/test_mssql_python/test_data_dictionary.py index 692eede3e..7e00102a7 100644 --- a/tests/unit/adapters/test_mssql_python/test_data_dictionary.py +++ b/tests/unit/adapters/test_mssql_python/test_data_dictionary.py @@ -138,3 +138,28 @@ def test_sync_data_dictionary_explicit_schema_skips_connection_lookup() -> None: data_dictionary.get_tables(cast(Any, driver), schema="custom") assert driver.select_calls[0][1]["schema_name"] == "custom" assert len(driver.executed) == 0 + + +def test_mssql_version_info_supports_vector() -> None: + """Version 17+ or Azure SQL engine editions support vectors.""" + v16 = MssqlVersionInfo(16, 0, 0, engine_edition=3) + v17 = MssqlVersionInfo(17, 0, 0, engine_edition=3) + azure = MssqlVersionInfo(16, 0, 0, engine_edition=5) + + assert v16.supports_vector() is False + assert v17.supports_vector() is True + assert azure.supports_vector() is True + + +def test_data_dictionary_vector_feature_flag_and_optimal_type() -> None: + """Sync data dictionary resolves supports_vector and optimal type for vector.""" + + class VectorDriver: + def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: + return {"product_version": "17.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} + + data_dictionary = MssqlPythonSyncDataDictionary() + driver = VectorDriver() + + assert data_dictionary.get_feature_flag(cast(Any, driver), "supports_vector") is True + assert data_dictionary.get_optimal_type(cast(Any, driver), "vector") == "VECTOR" diff --git a/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py b/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py index 29cf4ef3c..1d5b7f89d 100644 --- a/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py +++ b/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py @@ -78,13 +78,13 @@ def test_sync_load_from_arrow_skips_an_empty_table() -> None: assert conn._cursor.bulkcopy_calls == [] -def test_sync_load_from_arrow_overwrite_deletes_first() -> None: +def test_sync_load_from_arrow_overwrite_truncates_first() -> None: conn = _FakeConnection() driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) driver.load_from_arrow("dbo.orders", pa.table({"id": [1]}), overwrite=True) - assert conn._cursor.execute_calls == ["DELETE FROM [dbo].[orders]"] + assert conn._cursor.execute_calls == ["TRUNCATE TABLE [dbo].[orders]"] assert conn._cursor.arrow_calls @@ -94,5 +94,44 @@ def test_sync_load_from_arrow_overwrite_preserves_quoted_dots() -> None: driver.load_from_arrow('"dbo.schema"."orders.table"', pa.table({"id": [1]}), overwrite=True) - assert conn._cursor.execute_calls == ["DELETE FROM [dbo.schema].[orders.table]"] + assert conn._cursor.execute_calls == ["TRUNCATE TABLE [dbo.schema].[orders.table]"] assert conn._cursor.arrow_calls + + +def test_sync_load_from_arrow_overwrite_falls_back_to_delete_on_fk_reference() -> None: + conn = _FakeConnection() + driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) + + class FkError(Exception): + number = 4712 + + original_execute = conn._cursor.execute + + def execute_with_fk(sql: str, *args: Any) -> None: + original_execute(sql, *args) + if sql.startswith("TRUNCATE"): + raise FkError("Cannot truncate table referenced by foreign key") + + conn._cursor.execute = cast("Any", execute_with_fk) + driver.load_from_arrow("dbo.orders", pa.table({"id": [1]}), overwrite=True) + + assert conn._cursor.execute_calls == ["TRUNCATE TABLE [dbo].[orders]", "DELETE FROM [dbo].[orders]"] + assert conn._cursor.arrow_calls + + +def test_sync_load_from_arrow_forwards_bulk_copy_options() -> None: + conn = _FakeConnection() + driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) + table = pa.table({"id": [1, 2], "name": ["a", "b"]}) + + job = driver.load_from_arrow( + "orders", table, batch_size=500, check_constraints=True, fire_triggers=True, keep_nulls=True, table_lock=True + ) + + assert job.telemetry["rows_processed"] == 2 + _, _, kwargs = conn._cursor.arrow_calls[0] + assert kwargs["batch_size"] == 500 + assert kwargs["check_constraints"] is True + assert kwargs["fire_triggers"] is True + assert kwargs["keep_nulls"] is True + assert kwargs["table_lock"] is True diff --git a/tests/unit/adapters/test_mssql_python/test_type_converter.py b/tests/unit/adapters/test_mssql_python/test_type_converter.py index 83b1dd3c3..22540d70a 100644 --- a/tests/unit/adapters/test_mssql_python/test_type_converter.py +++ b/tests/unit/adapters/test_mssql_python/test_type_converter.py @@ -64,3 +64,9 @@ def test_tsql_time_maps_to_arrow_time64() -> None: def test_tsql_timestamp_remains_the_rowversion_binary_type() -> None: """TIMESTAMP is a T-SQL rowversion alias and must stay binary.""" assert mssql_type_to_arrow("timestamp") == pa.binary() + + +def test_mssql_type_to_arrow_maps_json_and_vector() -> None: + """JSON and VECTOR types should map to expected Arrow types.""" + assert mssql_type_to_arrow("json") == pa.string() + assert mssql_type_to_arrow("vector") == pa.list_(pa.float32()) diff --git a/tests/unit/adapters/test_pymssql/_fakes.py b/tests/unit/adapters/test_pymssql/_fakes.py index 8fcd4ec1d..6f1e0969e 100644 --- a/tests/unit/adapters/test_pymssql/_fakes.py +++ b/tests/unit/adapters/test_pymssql/_fakes.py @@ -61,6 +61,7 @@ def __init__(self, cursor: "FakeCursor | None" = None) -> None: self.rollbacks = 0 self.autocommit_values: list[bool] = [] self.autocommit_state = True + self.bulk_copy_calls: list[dict[str, Any]] = [] def cursor(self, *args: Any, **kwargs: Any) -> FakeCursor: self.cursor_args = args @@ -82,6 +83,26 @@ def autocommit(self, value: bool) -> None: self.autocommit_values.append(value) self.autocommit_state = value + def bulk_copy( + self, + table_name: str, + elements: Any, + column_ids: Any = None, + batch_size: int = 1000, + tablock: bool = False, + check_constraints: bool = False, + fire_triggers: bool = False, + ) -> None: + self.bulk_copy_calls.append({ + "table_name": table_name, + "elements": list(elements), + "column_ids": column_ids, + "batch_size": batch_size, + "tablock": tablock, + "check_constraints": check_constraints, + "fire_triggers": fire_triggers, + }) + class FakePymssqlModule: """Patch target that behaves like the pymssql module surface used by SQLSpec.""" diff --git a/tests/unit/adapters/test_pymssql/test_config.py b/tests/unit/adapters/test_pymssql/test_config.py index 7ddf6e783..a8f7c589f 100644 --- a/tests/unit/adapters/test_pymssql/test_config.py +++ b/tests/unit/adapters/test_pymssql/test_config.py @@ -31,6 +31,7 @@ def test_connection_params_cover_common_pymssql_keywords() -> None: "tds_version", "pool_recycle_seconds", "health_check_interval", + "encryption", } assert expected_keys <= set(annotations) @@ -45,6 +46,7 @@ def test_config_defaults_server_port_and_features() -> None: assert config.driver_type is PymssqlDriver assert config.supports_transactional_ddl is True assert config.supports_native_arrow_export is False + assert config.supports_native_arrow_import is True assert config.driver_features["enable_events"] is True diff --git a/tests/unit/adapters/test_pymssql/test_core.py b/tests/unit/adapters/test_pymssql/test_core.py index 26cec8aeb..c24ab02b9 100644 --- a/tests/unit/adapters/test_pymssql/test_core.py +++ b/tests/unit/adapters/test_pymssql/test_core.py @@ -150,3 +150,55 @@ def test_normalize_execute_many_parameters_passes_through() -> None: rows: list[tuple[Any, ...]] = [(1,), (2,)] assert normalize_execute_many_parameters(rows) is rows + + +def test_quote_tsql_identifier() -> None: + """quote_tsql_identifier wraps identifiers in brackets and escapes closing brackets.""" + from sqlspec.adapters.pymssql.core import quote_tsql_identifier + + assert quote_tsql_identifier("users") == "[users]" + assert quote_tsql_identifier("[users]") == "[users]" + assert quote_tsql_identifier("dbo.users") == "[dbo].[users]" + assert quote_tsql_identifier("col]name") == "[col]]name]" + + +def test_extract_error_number() -> None: + """extract_error_number detects error number from attribute, tuple, or regex.""" + from sqlspec.adapters.pymssql.core import extract_error_number + + class AttributeException(Exception): + number = 2627 + + assert extract_error_number(AttributeException("duplicate key")) == 2627 + assert extract_error_number(Exception(1205, "Deadlock found")) == 1205 + assert extract_error_number(Exception("Violation of UNIQUE KEY constraint (2627)")) == 2627 + assert extract_error_number(Exception("Plain error")) is None + + +def test_build_multi_row_insert() -> None: + """build_multi_row_insert generates a multi-row VALUES INSERT statement.""" + from sqlspec.adapters.pymssql.core import build_multi_row_insert + + sql = build_multi_row_insert("dbo.users", ["id", "name"], 3) + assert sql == "INSERT INTO [dbo].[users] ([id], [name]) VALUES (%s, %s), (%s, %s), (%s, %s)" + + +def test_collect_rows_preserves_list_identity() -> None: + """collect_rows avoids copying when the input rows are already a list.""" + from sqlspec.adapters.pymssql.core import collect_rows + + input_rows = [(1, "Alice"), (2, "Bob")] + description = [("id",), ("name",)] + rows, column_names, count = collect_rows(input_rows, description) + + assert rows is input_rows + assert column_names == ["id", "name"] + assert count == 2 + + +def test_normalize_execute_parameters_preserves_tuples() -> None: + """normalize_execute_parameters passes tuples through directly.""" + from sqlspec.adapters.pymssql.core import normalize_execute_parameters + + params = (1, "Alice") + assert normalize_execute_parameters(params) is params diff --git a/tests/unit/adapters/test_pymssql/test_data_dictionary.py b/tests/unit/adapters/test_pymssql/test_data_dictionary.py index b9c4aafe2..2102a4d5b 100644 --- a/tests/unit/adapters/test_pymssql/test_data_dictionary.py +++ b/tests/unit/adapters/test_pymssql/test_data_dictionary.py @@ -129,3 +129,28 @@ def test_sync_data_dictionary_explicit_schema_skips_connection_lookup() -> None: data_dictionary.get_tables(cast(Any, driver), schema="custom") assert driver.select_calls[0][1]["schema_name"] == "custom" assert len(driver.executed) == 0 + + +def test_mssql_version_info_supports_vector() -> None: + """Version 17+ or Azure SQL engine editions support vectors.""" + v16 = MssqlVersionInfo(16, 0, 0, engine_edition=3) + v17 = MssqlVersionInfo(17, 0, 0, engine_edition=3) + azure = MssqlVersionInfo(16, 0, 0, engine_edition=5) + + assert v16.supports_vector() is False + assert v17.supports_vector() is True + assert azure.supports_vector() is True + + +def test_data_dictionary_vector_feature_flag_and_optimal_type() -> None: + """Sync data dictionary resolves supports_vector and optimal type for vector.""" + + class VectorDriver: + def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: + return {"product_version": "17.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} + + data_dictionary = PymssqlSyncDataDictionary() + driver = VectorDriver() + + assert data_dictionary.get_feature_flag(cast(Any, driver), "supports_vector") is True + assert data_dictionary.get_optimal_type(cast(Any, driver), "vector") == "VECTOR" diff --git a/tests/unit/adapters/test_pymssql/test_driver.py b/tests/unit/adapters/test_pymssql/test_driver.py index 7a75e550c..ece0e7eac 100644 --- a/tests/unit/adapters/test_pymssql/test_driver.py +++ b/tests/unit/adapters/test_pymssql/test_driver.py @@ -1,6 +1,6 @@ """pymssql driver tests.""" -from typing import cast +from typing import Any, cast import pytest from pymssql import IntegrityError as PymssqlIntegrityError @@ -278,3 +278,125 @@ def fail(sql: str, parameters: object = None) -> None: assert sum(sql == "BEGIN TRANSACTION" for sql, _ in connection.cursor_obj.calls) == 1 driver.rollback() assert connection.cursor_obj.calls[-1] == ("IF @@TRANCOUNT > 0 ROLLBACK TRANSACTION", None) + + +def test_driver_bulk_copy_forwards_options() -> None: + """bulk_copy forwards batch options to underlying connection.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + + connection = FakeConnection() + driver = PymssqlDriver(cast("PymssqlConnection", connection)) + + result = driver.bulk_copy( + "dbo.users", + [(1, "Ada"), (2, "Grace")], + column_ids=[1, 2], + batch_size=500, + tablock=True, + check_constraints=True, + fire_triggers=True, + ) + + assert result == 2 + assert len(connection.bulk_copy_calls) == 1 + call = connection.bulk_copy_calls[0] + assert call["table_name"] == "[dbo].[users]" + assert call["elements"] == [(1, "Ada"), (2, "Grace")] + assert call["column_ids"] == [1, 2] + assert call["batch_size"] == 500 + assert call["tablock"] is True + assert call["check_constraints"] is True + assert call["fire_triggers"] is True + + +def test_load_from_arrow_bulk_copies_batches() -> None: + """load_from_arrow processes Arrow table in batches via bulk_copy.""" + import pyarrow as pa + + from sqlspec.adapters.pymssql.driver import PymssqlDriver + + connection = FakeConnection() + driver = PymssqlDriver(cast("PymssqlConnection", connection)) + table = pa.table({"id": [1, 2], "name": ["Ada", "Grace"]}) + + job = driver.load_from_arrow("dbo.users", table, batch_size=500) + + assert job.telemetry["rows_processed"] == 2 + assert len(connection.bulk_copy_calls) == 1 + + +def test_load_from_arrow_overwrite_truncates_first() -> None: + """load_from_arrow with overwrite=True executes TRUNCATE TABLE.""" + import pyarrow as pa + + from sqlspec.adapters.pymssql.driver import PymssqlDriver + + cursor = FakeCursor() + connection = FakeConnection(cursor) + driver = PymssqlDriver(cast("PymssqlConnection", connection)) + table = pa.table({"id": [1], "name": ["Ada"]}) + + driver.load_from_arrow("dbo.users", table, overwrite=True) + + executed_sqls = [call[0] for call in cursor.calls] + assert "TRUNCATE TABLE [dbo].[users]" in executed_sqls + + +def test_load_from_arrow_overwrite_falls_back_on_fk_error() -> None: + """load_from_arrow falls back to DELETE FROM when error 4712 is encountered.""" + import pyarrow as pa + + from sqlspec.adapters.pymssql.driver import PymssqlDriver + + class FkError(Exception): + number = 4712 + + cursor = FakeCursor() + connection = FakeConnection(cursor) + driver = PymssqlDriver(cast("PymssqlConnection", connection)) + + def execute_with_fk(sql: str, *args: Any) -> None: + cursor.calls.append((sql, args)) + if sql.startswith("TRUNCATE"): + raise FkError("Cannot truncate table referenced by foreign key") + + cursor.execute = cast("Any", execute_with_fk) + table = pa.table({"id": [1], "name": ["Ada"]}) + + driver.load_from_arrow("dbo.users", table, overwrite=True) + + executed_sqls = [call[0] for call in cursor.calls] + assert "TRUNCATE TABLE [dbo].[users]" in executed_sqls + assert "DELETE FROM [dbo].[users]" in executed_sqls + + +def test_execute_many_plain_values_chunks_into_multi_row_insert() -> None: + """execute_many with plain VALUES uses multi-row INSERT.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + + cursor = FakeCursor(rowcount=3) + connection = FakeConnection(cursor) + driver = PymssqlDriver(cast("PymssqlConnection", connection)) + + params = [(1, "Ada"), (2, "Grace"), (3, "Linus")] + result = driver.execute_many("INSERT INTO dbo.users (id, name) VALUES (?, ?)", params) + + assert result.rows_affected == 3 + executed_sqls = [call[0] for call in cursor.calls] + assert len(executed_sqls) == 1 + assert "VALUES (%s, %s), (%s, %s), (%s, %s)" in executed_sqls[0] + + +def test_execute_many_non_plain_values_uses_standard_executemany() -> None: + """execute_many with non-plain SQL uses cursor.executemany.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + + cursor = FakeCursor(rowcount=2) + connection = FakeConnection(cursor) + driver = PymssqlDriver(cast("PymssqlConnection", connection)) + + params = [("Ada", 1), ("Grace", 2)] + result = driver.execute_many("UPDATE dbo.users SET name = ? WHERE id = ?", params) + + assert result.rows_affected == 2 + assert len(cursor.many_calls) == 1 From e79d46f093891d8398e2f438418e2f6320404fe9 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 01:48:06 +0000 Subject: [PATCH 02/11] fix(adapters): resolve circular imports, typing, and test regressions in MSSQL adapters --- sqlspec/adapters/mssql_python/_typing.py | 30 ++++----- sqlspec/adapters/mssql_python/adk/store.py | 63 +++++++++++-------- sqlspec/adapters/mssql_python/core.py | 42 ++++++++----- sqlspec/adapters/mssql_python/driver.py | 47 +++++++++----- .../adapters/mssql_python/litestar/store.py | 63 ++++++++++--------- .../adapters/mssql_python/type_converter.py | 11 ++-- sqlspec/adapters/pymssql/_typing.py | 30 ++++----- sqlspec/adapters/pymssql/adk/store.py | 60 +++++++++--------- sqlspec/adapters/pymssql/data_dictionary.py | 42 +++++++------ sqlspec/adapters/pymssql/driver.py | 27 +++++++- sqlspec/adapters/pymssql/litestar/store.py | 63 ++++++++++--------- .../adapters/test_mssql_python/test_config.py | 2 +- tests/unit/adapters/test_pymssql/test_core.py | 6 +- 13 files changed, 282 insertions(+), 204 deletions(-) diff --git a/sqlspec/adapters/mssql_python/_typing.py b/sqlspec/adapters/mssql_python/_typing.py index f35a4a841..5f006110c 100644 --- a/sqlspec/adapters/mssql_python/_typing.py +++ b/sqlspec/adapters/mssql_python/_typing.py @@ -1,8 +1,6 @@ """mssql-python adapter type definitions and mypyc-excluded context managers.""" import contextlib -from collections.abc import Callable -from types import TracebackType from typing import TYPE_CHECKING, Any import mssql_python as _mssql_python @@ -10,14 +8,16 @@ from mssql_python.connection import Connection, TokenProvider from mssql_python.cursor import Cursor -from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver -from sqlspec.core import StatementConfig - MSSQL_PYTHON_MODULE: Any = _mssql_python if TYPE_CHECKING: + from collections.abc import Callable + from types import TracebackType from typing import TypeAlias + from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver + from sqlspec.core import StatementConfig + MssqlPythonConnection: TypeAlias = Connection MssqlPythonRawCursor: TypeAlias = Cursor @@ -41,11 +41,11 @@ class MssqlPythonCursor: __slots__ = ("connection", "cursor") - def __init__(self, connection: MssqlPythonConnection) -> None: + def __init__(self, connection: "MssqlPythonConnection") -> None: self.connection = connection self.cursor: MssqlPythonRawCursor | None = None - def __enter__(self) -> MssqlPythonRawCursor: + def __enter__(self) -> "MssqlPythonRawCursor": self.cursor = self.connection.cursor() return self.cursor @@ -70,11 +70,11 @@ class MssqlPythonSessionContext: def __init__( self, - acquire_connection: Callable[[], MssqlPythonConnection], - release_connection: Callable[..., Any], - statement_config: StatementConfig, - driver_features: dict[str, Any], - prepare_driver: Callable[[MssqlPythonDriver], MssqlPythonDriver], + acquire_connection: "Callable[[], MssqlPythonConnection]", + release_connection: "Callable[..., Any]", + statement_config: "StatementConfig", + driver_features: "dict[str, Any]", + prepare_driver: "Callable[[MssqlPythonDriver], MssqlPythonDriver]", ) -> None: self._acquire_connection = acquire_connection self._release_connection = release_connection @@ -84,7 +84,7 @@ def __init__( self._connection: MssqlPythonConnection | None = None self._driver: MssqlPythonDriver | None = None - def __enter__(self) -> MssqlPythonDriver: + def __enter__(self) -> "MssqlPythonDriver": from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver self._connection = self._acquire_connection() @@ -94,8 +94,8 @@ def __enter__(self) -> MssqlPythonDriver: return self._prepare_driver(self._driver) def __exit__( - self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None - ) -> bool | None: + self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" + ) -> "bool | None": if exc_type is not None and self._driver is not None: with contextlib.suppress(Exception): self._driver.rollback() diff --git a/sqlspec/adapters/mssql_python/adk/store.py b/sqlspec/adapters/mssql_python/adk/store.py index 25b85bb6d..629611567 100644 --- a/sqlspec/adapters/mssql_python/adk/store.py +++ b/sqlspec/adapters/mssql_python/adk/store.py @@ -6,7 +6,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.mssql_python._typing import MssqlPythonError +from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor, MssqlPythonError from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.adapters.mssql_python.core import extract_error_number from sqlspec.config import ADKConfig @@ -191,20 +191,21 @@ def append_event_and_update_state( OUTPUT inserted.id, inserted.app_name, inserted.user_id, inserted.state, inserted.create_time, inserted.update_time WHERE app_name = ? AND user_id = ? AND id = ? """ - with self._config.provide_session() as driver: + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: try: - row = driver.select_one_or_none(update_sql, (to_json(state), app_name, user_id, session_id)) + cursor.execute(update_sql, (to_json(state), app_name, user_id, session_id)) + row = cursor.fetchone() if row is None: _raise_session_not_found(session_id) - driver.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) + cursor.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) if app_state is not None: - driver.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) + cursor.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) if user_state is not None: - driver.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) + cursor.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) except Exception: - driver.rollback() + conn.rollback() raise - driver.commit() + conn.commit() return _session_record_from_row(row) def get_events( @@ -384,22 +385,24 @@ def _json_column_type_sync(self) -> str: return self._json_column_type def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> Any | None: - with self._config.provide_session() as driver: - row = driver.select_one_or_none(sql, params) + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, params) + row = cursor.fetchone() if commit: - driver.commit() + conn.commit() return row def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: - with self._config.provide_session() as driver: - return driver.select(sql, params) + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, params) + return list(cursor.fetchall()) def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: - with self._config.provide_session() as driver: - res = driver.execute(sql, params) - rowcount = res.rows_affected + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, params) + rowcount = _cursor_rowcount(cursor) if commit: - driver.commit() + conn.commit() return rowcount @@ -451,7 +454,7 @@ def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | ); """ inserted = 0 - with self._config.provide_session() as driver: + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: for entry in entries: params: tuple[Any, ...] = ( entry["id"], @@ -468,9 +471,9 @@ def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | ) if self._owner_id_column_name: params = (*params, owner_id) - res = driver.execute(sql, (*params, entry["event_id"])) - inserted += res.rows_affected - driver.commit() + cursor.execute(sql, (*params, entry["event_id"])) + inserted += _cursor_rowcount(cursor) + conn.commit() return inserted def search_entries( @@ -563,18 +566,24 @@ def _drop_memory_table_sql(self) -> list[str]: return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"] def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: - with self._config.provide_session() as driver: - return driver.select(sql, params) + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, params) + return list(cursor.fetchall()) def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: - with self._config.provide_session() as driver: - res = driver.execute(sql, params) - rowcount = res.rows_affected + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, params) + rowcount = _cursor_rowcount(cursor) if commit: - driver.commit() + conn.commit() return rowcount +def _cursor_rowcount(cursor: Any) -> int: + rowcount = getattr(cursor, "rowcount", 0) + return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 + + def _adk_config(config: Any) -> MssqlPythonADKConfig: extension_config = getattr(config, "extension_config", {}) if not isinstance(extension_config, dict): diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index e5b8e0156..6a93ce950 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -41,7 +41,7 @@ "materialize_tuple_rows", ) -_ERROR_NUMBER_PATTERN: Final[re.Pattern[str]] = re.compile(r"\(([-]?\d+)(?:,|\))") +_ERROR_NUMBER_PATTERN: Final[re.Pattern[str]] = re.compile(r"(?:\(([-]?\d+)(?:,|\))|\bMsg\s+([-]?\d+)\b)") _MSSQL_CONSTRAINT_547: Final[int] = 547 _VERSION_PATTERN: Final[re.Pattern[str]] = re.compile(r"(\d+)") _VERSION_PART_COUNT: Final[int] = 3 @@ -97,9 +97,13 @@ def extract_error_number(exc: BaseException | None) -> int | None: - """Extract numeric SQL Server error code using fast string parsing before regex fallback.""" + """Extract numeric SQL Server error code using fast attribute/string parsing before regex fallback.""" if exc is None: return None + for attr in ("number", "error_code", "errno"): + val = getattr(exc, attr, None) + if isinstance(val, int) and not isinstance(val, bool): + return val ddbc_err = getattr(exc, "ddbc_error", None) if isinstance(ddbc_err, str) and ddbc_err.startswith("("): end_idx = ddbc_err.find(",") @@ -112,25 +116,31 @@ def extract_error_number(exc: BaseException | None) -> int | None: except ValueError: pass - if exc.args and isinstance(exc.args[0], str): - msg = exc.args[0] - start_idx = msg.rfind("(") - if start_idx != -1: - end_idx = msg.find(",", start_idx) - if end_idx == -1: - end_idx = msg.find(")", start_idx) - if end_idx != -1: - num_str = msg[start_idx + 1 : end_idx].strip() - try: - return int(num_str) - except ValueError: - pass + if exc.args: + first_arg = exc.args[0] + if isinstance(first_arg, int) and not isinstance(first_arg, bool): + return first_arg + if isinstance(first_arg, str): + msg = first_arg + start_idx = msg.rfind("(") + if start_idx != -1: + end_idx = msg.find(",", start_idx) + if end_idx == -1: + end_idx = msg.find(")", start_idx) + if end_idx != -1: + num_str = msg[start_idx + 1 : end_idx].strip() + try: + return int(num_str) + except ValueError: + pass matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) if not matches: return None + last_match = matches[-1] + raw_num = last_match[0] or last_match[1] if isinstance(last_match, tuple) else last_match try: - return int(matches[-1]) + return int(raw_num) except ValueError: return None diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index da92bf8c1..3b67e26b8 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -2,7 +2,7 @@ import contextlib from collections.abc import Iterable -from typing import Any, TypedDict, cast +from typing import TYPE_CHECKING, Any, TypedDict, cast from typing_extensions import NotRequired @@ -20,13 +20,10 @@ materialize_tuple_rows, ) from sqlspec.adapters.mssql_python.data_dictionary import MssqlPythonSyncDataDictionary -from sqlspec.builder import QueryBuilder from sqlspec.core import ( SQL, ArrowResult, - Statement, StatementConfig, - StatementFilter, build_arrow_result_from_reader, build_arrow_result_from_table, get_cache_config, @@ -42,12 +39,16 @@ ) from sqlspec.exceptions import SQLSpecError from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry -from sqlspec.typing import ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters from sqlspec.utils.arrow_helpers import arrow_reader_with_deferred_close from sqlspec.utils.logging import get_logger from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.text import split_qualified_identifier +if TYPE_CHECKING: + from sqlspec.builder import QueryBuilder + from sqlspec.core import Statement, StatementFilter + from sqlspec.typing import ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters + __all__ = ( "MssqlPythonBulkCopyResult", "MssqlPythonCursor", @@ -300,11 +301,11 @@ def has_schema(self, schema: str) -> bool: def select_to_arrow( self, - statement: Statement | QueryBuilder, + statement: "Statement | QueryBuilder", /, - *parameters: StatementParameters | StatementFilter, + *parameters: "StatementParameters | StatementFilter", statement_config: StatementConfig | None = None, - return_format: ArrowReturnFormat = "table", + return_format: "ArrowReturnFormat" = "table", native_only: bool = False, batch_size: int | None = None, arrow_schema: Any = None, @@ -441,16 +442,18 @@ def load_from_arrow( self._check_pending_exception(exc_handler) raw_result: Any = None - is_stream = hasattr(source, "__arrow_c_stream__") + is_table_source = isinstance(source, ArrowResult) is_reader = False try: import pyarrow as pa + is_table_source = is_table_source or isinstance(source, pa.Table) is_reader = isinstance(source, (pa.RecordBatchReader, pa.RecordBatch)) except ImportError: pass + is_stream = not is_table_source and (is_reader or hasattr(source, "__arrow_c_stream__")) - if is_stream or is_reader: + if is_stream: cols = column_mappings source_schema = getattr(source, "schema", None) if cols is None and source_schema is not None: @@ -551,10 +554,24 @@ def _quote_mssql_table(table: str) -> str: def _execute_cursor(cursor: MssqlPythonRawCursor, sql: str, parameters: Any, *, use_prepare: bool = True) -> None: - if parameters is None: - cursor.execute(sql, use_prepare=use_prepare) - else: - cursor.execute(sql, parameters, use_prepare=use_prepare) + if use_prepare: + if parameters is None: + cursor.execute(sql) + else: + cursor.execute(sql, parameters) + return + try: + if parameters is None: + cursor.execute(sql, use_prepare=False) + else: + cursor.execute(sql, parameters, use_prepare=False) + except TypeError as exc: + if "use_prepare" not in str(exc): + raise + if parameters is None: + cursor.execute(sql) + else: + cursor.execute(sql, parameters) def _cursor_rowcount(cursor: MssqlPythonRawCursor) -> int: @@ -576,7 +593,7 @@ def _resolve_column_names(description: Any, cache: dict[int, tuple[Any, list[str return column_names -def _cursor_arrow_reader(cursor: MssqlPythonRawCursor, arrow_kwargs: dict[str, int]) -> ArrowRecordBatchReader | None: +def _cursor_arrow_reader(cursor: MssqlPythonRawCursor, arrow_kwargs: dict[str, int]) -> "ArrowRecordBatchReader | None": arrow_reader = getattr(cursor, "arrow_reader", None) if not callable(arrow_reader): return None diff --git a/sqlspec/adapters/mssql_python/litestar/store.py b/sqlspec/adapters/mssql_python/litestar/store.py index c67f20036..386ea45f4 100644 --- a/sqlspec/adapters/mssql_python/litestar/store.py +++ b/sqlspec/adapters/mssql_python/litestar/store.py @@ -3,6 +3,7 @@ from datetime import datetime, timedelta, timezone from typing import Any +from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -94,8 +95,11 @@ def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | No WHERE session_id = ? AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_session() as driver: - row = driver.select_one_or_none(sql, (key,)) + with self._config.provide_connection() as conn: + with MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, (key,)) + row = cursor.fetchone() + if row is None: return None @@ -103,15 +107,16 @@ def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | No if renew_for is not None and expires_at is not None: new_expires_at = self._calculate_expires_at(renew_for) if new_expires_at is not None: - driver.execute( - f""" - UPDATE {self._table_name} - SET expires_at = ?, updated_at = SYSUTCDATETIME() - WHERE session_id = ? - """, - (new_expires_at, key), - ) - driver.commit() + with MssqlPythonCursor(conn) as update_cursor: + update_cursor.execute( + f""" + UPDATE {self._table_name} + SET expires_at = ?, updated_at = SYSUTCDATETIME() + WHERE session_id = ? + """, + (new_expires_at, key), + ) + conn.commit() return _coerce_bytes(_row_value(row, "data", 0)) @@ -131,19 +136,19 @@ def _set(self, key: str, value: str | bytes, expires_in: int | timedelta | None INSERT (session_id, data, expires_at) VALUES (src.session_id, src.data, src.expires_at); """ - with self._config.provide_session() as driver: - driver.execute(sql, (key, data, expires_at)) - driver.commit() + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, (key, data, expires_at)) + conn.commit() def _delete(self, key: str) -> None: - with self._config.provide_session() as driver: - driver.execute(f"DELETE FROM {self._table_name} WHERE session_id = ?", (key,)) - driver.commit() + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(f"DELETE FROM {self._table_name} WHERE session_id = ?", (key,)) + conn.commit() def _delete_all(self) -> None: - with self._config.provide_session() as driver: - driver.execute(f"TRUNCATE TABLE {self._table_name}") - driver.commit() + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(f"TRUNCATE TABLE {self._table_name}") + conn.commit() self._log_delete_all() def _exists(self, key: str) -> bool: @@ -153,12 +158,14 @@ def _exists(self, key: str) -> bool: WHERE session_id = ? AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_session() as driver: - return driver.select_one_or_none(sql, (key,)) is not None + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql, (key,)) + return cursor.fetchone() is not None def _expires_in(self, key: str) -> int | None: - with self._config.provide_session() as driver: - row = driver.select_one_or_none(f"SELECT expires_at FROM {self._table_name} WHERE session_id = ?", (key,)) + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(f"SELECT expires_at FROM {self._table_name} WHERE session_id = ?", (key,)) + row = cursor.fetchone() if row is None: return None @@ -174,10 +181,10 @@ def _delete_expired(self) -> int: WHERE expires_at IS NOT NULL AND expires_at < SYSUTCDATETIME() """ - with self._config.provide_session() as driver: - res = driver.execute(sql) - driver.commit() - count = res.rows_affected + with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: + cursor.execute(sql) + count = int(getattr(cursor, "rowcount", 0) or 0) + conn.commit() if count > 0: self._log_delete_expired(count) return count diff --git a/sqlspec/adapters/mssql_python/type_converter.py b/sqlspec/adapters/mssql_python/type_converter.py index 18d2a626a..a1d482413 100644 --- a/sqlspec/adapters/mssql_python/type_converter.py +++ b/sqlspec/adapters/mssql_python/type_converter.py @@ -1,14 +1,15 @@ """Type converters for mssql-python parameter binding.""" from collections.abc import Callable -from typing import Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, cast from uuid import UUID -import pyarrow as pa - from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.serializers import from_json, to_json +if TYPE_CHECKING: + import pyarrow as pa + __all__ = ("MssqlPythonTypeConverter", "mssql_type_to_arrow") _MSSQL_ARROW_TYPE_SPECS: Final[dict[str, tuple[str, tuple[Any, ...], dict[str, Any]]]] = { @@ -73,7 +74,7 @@ def coerce_read_value(self, value: Any) -> Any: return value -def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: int | None = None) -> pa.DataType: +def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: int | None = None) -> "pa.DataType": """Resolve a T-SQL type name to an Arrow data type.""" normalized_type = sql_type.lower().split("(", 1)[0].strip() if normalized_type == "vector": @@ -90,7 +91,7 @@ def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: i return _arrow_type(name, args, kwargs) -def _arrow_type(name: str, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None) -> pa.DataType: +def _arrow_type(name: str, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None) -> "pa.DataType": ensure_pyarrow() import pyarrow as pa diff --git a/sqlspec/adapters/pymssql/_typing.py b/sqlspec/adapters/pymssql/_typing.py index 9a704d45a..ceb7c8433 100644 --- a/sqlspec/adapters/pymssql/_typing.py +++ b/sqlspec/adapters/pymssql/_typing.py @@ -5,8 +5,6 @@ """ import contextlib -from collections.abc import Callable -from types import TracebackType from typing import TYPE_CHECKING, Any import pymssql as _pymssql @@ -14,16 +12,18 @@ from pymssql import Cursor as _PymssqlRawCursor from pymssql import Error as PymssqlError -from sqlspec.adapters.pymssql.driver import PymssqlDriver -from sqlspec.core import StatementConfig - PYMSSQL_MODULE = _pymssql if TYPE_CHECKING: + from collections.abc import Callable + from types import TracebackType from typing import TypeAlias from pymssql._pymssql import QueryParams as PymssqlQueryParams + from sqlspec.adapters.pymssql.driver import PymssqlDriver + from sqlspec.core import StatementConfig + PymssqlConnection: TypeAlias = _PymssqlConnection PymssqlRawCursor: TypeAlias = _PymssqlRawCursor @@ -48,11 +48,11 @@ class PymssqlCursor: __slots__ = ("connection", "cursor") - def __init__(self, connection: PymssqlConnection) -> None: + def __init__(self, connection: "PymssqlConnection") -> None: self.connection = connection self.cursor: PymssqlRawCursor | None = None - def __enter__(self) -> PymssqlRawCursor: + def __enter__(self) -> "PymssqlRawCursor": self.cursor = self.connection.cursor() return self.cursor @@ -77,11 +77,11 @@ class PymssqlSessionContext: def __init__( self, - acquire_connection: Callable[[], Any], - release_connection: Callable[..., Any], - statement_config: StatementConfig, - driver_features: dict[str, Any], - prepare_driver: Callable[[PymssqlDriver], PymssqlDriver], + acquire_connection: "Callable[[], Any]", + release_connection: "Callable[..., Any]", + statement_config: "StatementConfig", + driver_features: "dict[str, Any]", + prepare_driver: "Callable[[PymssqlDriver], PymssqlDriver]", ) -> None: self._acquire_connection = acquire_connection self._release_connection = release_connection @@ -91,7 +91,7 @@ def __init__( self._connection: Any = None self._driver: PymssqlDriver | None = None - def __enter__(self) -> PymssqlDriver: + def __enter__(self) -> "PymssqlDriver": from sqlspec.adapters.pymssql.driver import PymssqlDriver self._connection = self._acquire_connection() @@ -101,8 +101,8 @@ def __enter__(self) -> PymssqlDriver: return self._prepare_driver(self._driver) def __exit__( - self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None - ) -> bool | None: + self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" + ) -> "bool | None": if exc_type is not None and self._driver is not None: with contextlib.suppress(Exception): self._driver.rollback() diff --git a/sqlspec/adapters/pymssql/adk/store.py b/sqlspec/adapters/pymssql/adk/store.py index 0308fea95..b326ab9fc 100644 --- a/sqlspec/adapters/pymssql/adk/store.py +++ b/sqlspec/adapters/pymssql/adk/store.py @@ -6,9 +6,9 @@ from typing_extensions import NotRequired -from sqlspec.adapters.pymssql._typing import PymssqlError +from sqlspec.adapters.pymssql._typing import PymssqlCursor, PymssqlError from sqlspec.adapters.pymssql.config import PymssqlConfig -from sqlspec.adapters.pymssql.core import extract_error_number, quote_tsql_identifier +from sqlspec.adapters.pymssql.core import extract_error_number, quote_tsql_identifier, resolve_rowcount from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.config import ADKConfig @@ -194,20 +194,21 @@ def append_event_and_update_state( OUTPUT inserted.id, inserted.app_name, inserted.user_id, inserted.state, inserted.create_time, inserted.update_time WHERE app_name = %s AND user_id = %s AND id = %s """ - with self._config.provide_session() as driver: + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: try: - row = driver.select_one_or_none(update_sql, (to_json(state), app_name, user_id, session_id)) + cursor.execute(update_sql, (to_json(state), app_name, user_id, session_id)) + row = cursor.fetchone() if row is None: _raise_session_not_found(session_id) - driver.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) + cursor.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record)) if app_state is not None: - driver.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) + cursor.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state))) if user_state is not None: - driver.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) + cursor.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state))) except Exception: - driver.rollback() + conn.rollback() raise - driver.commit() + conn.commit() return _session_record_from_row(row) def get_events( @@ -395,22 +396,24 @@ def _json_column_type_sync(self) -> str: return self._json_column_type def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> Any | None: - with self._config.provide_session() as driver: - row = driver.select_one_or_none(sql, params) + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, params) + row = cursor.fetchone() if commit: - driver.commit() + conn.commit() return row def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: - with self._config.provide_session() as driver: - return driver.select(sql, params) + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, params) + return list(cursor.fetchall()) def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: - with self._config.provide_session() as driver: - res = driver.execute(sql, params) - rowcount = res.rows_affected + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, params) + rowcount = resolve_rowcount(cursor) if commit: - driver.commit() + conn.commit() return rowcount @@ -462,7 +465,7 @@ def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | ); """ inserted = 0 - with self._config.provide_session() as driver: + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: for entry in entries: params: tuple[Any, ...] = ( entry["id"], @@ -479,9 +482,9 @@ def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | ) if self._owner_id_column_name: params = (*params, owner_id) - res = driver.execute(sql, (*params, entry["event_id"])) - inserted += res.rows_affected - driver.commit() + cursor.execute(sql, (*params, entry["event_id"])) + inserted += resolve_rowcount(cursor) + conn.commit() return inserted def search_entries( @@ -574,15 +577,16 @@ def _drop_memory_table_sql(self) -> list[str]: return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"] def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: - with self._config.provide_session() as driver: - return driver.select(sql, params) + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, params) + return list(cursor.fetchall()) def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: - with self._config.provide_session() as driver: - res = driver.execute(sql, params) - rowcount = res.rows_affected + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, params) + rowcount = resolve_rowcount(cursor) if commit: - driver.commit() + conn.commit() return rowcount diff --git a/sqlspec/adapters/pymssql/data_dictionary.py b/sqlspec/adapters/pymssql/data_dictionary.py index eac2875ec..9d5a5de9e 100644 --- a/sqlspec/adapters/pymssql/data_dictionary.py +++ b/sqlspec/adapters/pymssql/data_dictionary.py @@ -1,19 +1,14 @@ """pymssql data dictionary.""" -from collections.abc import Sequence -from typing import Any, ClassVar, Final, cast +from typing import TYPE_CHECKING, Any, ClassVar, Final, cast from mypy_extensions import mypyc_attr -from sqlspec.adapters.pymssql.driver import PymssqlDriver -from sqlspec.core import SQL from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, - DialectConfig, ForeignKeyMetadata, IndexMetadata, - MetadataCapabilityProfile, MetadataSupport, SystemMetadataCapability, SystemMetadataRequest, @@ -44,6 +39,13 @@ from sqlspec.driver import SyncDataDictionaryBase from sqlspec.utils.logging import get_logger +if TYPE_CHECKING: + from collections.abc import Sequence + + from sqlspec.adapters.pymssql.driver import PymssqlDriver + from sqlspec.core import SQL + from sqlspec.data_dictionary import DialectConfig, MetadataCapabilityProfile + __all__ = ("MssqlVersionInfo", "PymssqlSyncDataDictionary") logger = get_logger("sqlspec.adapters.pymssql.data_dictionary") @@ -98,7 +100,7 @@ class _MssqlDataDictionaryMixin: dialect: ClassVar[str] = "mssql" - def get_dialect_config(self) -> DialectConfig: + def get_dialect_config(self) -> "DialectConfig": """Return the dialect configuration for this data dictionary.""" return get_dialect_config(type(self).dialect) @@ -113,7 +115,7 @@ def list_available_features(self) -> list[str]: """List available feature flags for this dialect.""" return list_mssql_available_features(self.get_dialect_config()) - def get_domain_query(self, domain: str, name: str) -> SQL: + def get_domain_query(self, domain: str, name: str) -> "SQL": """Return a SQL Server domain query.""" query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, name) return cast("SQL", query.sql) @@ -151,20 +153,20 @@ def __init__(self) -> None: super().__init__() def get_metadata_capabilities( - self, driver: PymssqlDriver, domains: Sequence[str] | None = None - ) -> MetadataCapabilityProfile: + self, driver: "PymssqlDriver", domains: "Sequence[str] | None" = None + ) -> "MetadataCapabilityProfile": """Get SQL Server data-dictionary capability profile.""" return build_mssql_metadata_capability_profile(type(self).__name__, domains) def get_system_metadata_capabilities( - self, driver: PymssqlDriver, domains: Sequence[str] | None = None + self, driver: "PymssqlDriver", domains: "Sequence[str] | None" = None ) -> tuple[SystemMetadataCapability, ...]: """Get SQL Server opt-in system metadata capability disclosures.""" _ = driver requested_domains = ("dmv_exec_requests", "query_store_runtime") if domains is None else tuple(domains) return tuple(build_mssql_system_metadata_capability(domain) for domain in requested_domains) - def get_version(self, driver: PymssqlDriver) -> MssqlVersionInfo | None: + def get_version(self, driver: "PymssqlDriver") -> MssqlVersionInfo | None: """Get SQL Server version information.""" driver_id = id(driver) if driver_id in self._version_fetch_attempted: @@ -191,7 +193,7 @@ def get_version(self, driver: PymssqlDriver) -> MssqlVersionInfo | None: self.cache_version(driver_id, version_info) return version_info - def get_feature_flag(self, driver: PymssqlDriver, feature: str) -> bool: + def get_feature_flag(self, driver: "PymssqlDriver", feature: str) -> bool: """Check whether SQL Server supports a feature.""" version_info = self.get_version(driver) if feature == "supports_vector": @@ -204,11 +206,11 @@ def get_feature_flag(self, driver: PymssqlDriver, feature: str) -> bool: version_info=version_info, ) - def get_optimal_type(self, driver: PymssqlDriver, type_category: str) -> str: + def get_optimal_type(self, driver: "PymssqlDriver", type_category: str) -> str: """Get optimal SQL Server type for a category.""" return self._get_optimal_type_from_version(self.get_version(driver), type_category) - def get_tables(self, driver: PymssqlDriver, schema: str | None = None) -> list[TableMetadata]: + def get_tables(self, driver: "PymssqlDriver", schema: str | None = None) -> list[TableMetadata]: """Get tables sorted by dependency order with catalog fallback.""" schema_name = self.resolve_connection_schema(driver, schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -227,7 +229,7 @@ def get_tables(self, driver: PymssqlDriver, schema: str | None = None) -> list[T return merge_mssql_table_lists(ordered, all_rows) def get_columns( - self, driver: PymssqlDriver, table: str | None = None, schema: str | None = None + self, driver: "PymssqlDriver", table: str | None = None, schema: str | None = None ) -> list[ColumnMetadata]: """Get column information for a table or schema.""" schema_name = self.resolve_connection_schema(driver, schema) @@ -251,7 +253,7 @@ def get_columns( ) def get_indexes( - self, driver: PymssqlDriver, table: str | None = None, schema: str | None = None + self, driver: "PymssqlDriver", table: str | None = None, schema: str | None = None ) -> list[IndexMetadata]: """Get index metadata for a table or schema.""" schema_name = self.resolve_connection_schema(driver, schema) @@ -275,7 +277,7 @@ def get_indexes( ) def get_foreign_keys( - self, driver: PymssqlDriver, table: str | None = None, schema: str | None = None + self, driver: "PymssqlDriver", table: str | None = None, schema: str | None = None ) -> list[ForeignKeyMetadata]: """Get foreign key metadata.""" schema_name = self.resolve_connection_schema(driver, schema) @@ -302,7 +304,7 @@ def get_foreign_keys( def get_ddl( self, - driver: PymssqlDriver, + driver: "PymssqlDriver", object_name: str, schema: str | None = None, *, @@ -323,7 +325,7 @@ def get_ddl( return build_mssql_table_ddl_result(schema_name, object_name, columns, indexes, object_type=object_type) def get_system_metadata( - self, driver: PymssqlDriver, request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "PymssqlDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in SQL Server system metadata with sensitive columns redacted by default.""" metadata_request = ensure_system_metadata_request(request, **kwargs) diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 98920df5f..5784494f5 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -4,7 +4,8 @@ from collections.abc import Iterable, Sequence, Sized from typing import TYPE_CHECKING, Any, cast -import sqlglot.expressions as exp +import sqlglot +from sqlglot import exp from sqlspec.adapters.pymssql._typing import ( PymssqlConnection, @@ -142,6 +143,15 @@ def __init__( statement_config = default_statement_config.replace( enable_caching=get_cache_config().compiled_cache_enabled ) + if driver_features is None or "storage_capabilities" not in driver_features: + driver_features = dict(driver_features) if driver_features else {} + driver_features["storage_capabilities"] = { + "arrow_export_enabled": False, + "arrow_import_enabled": True, + "parquet_export_enabled": False, + "parquet_import_enabled": False, + "partition_strategies": [], + } super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features) self._data_dictionary: PymssqlSyncDataDictionary | None = None @@ -191,6 +201,13 @@ def dispatch_execute_script(self, cursor: PymssqlRawCursor, statement: SQL) -> E cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) + def collect_rows(self, cursor: PymssqlRawCursor, fetched: list[Any]) -> tuple[list[Any], list[str], int]: + rows, column_names, _ = collect_rows(fetched, cursor.description or None, self._column_name_cache) + return rows, column_names, len(rows) + + def resolve_rowcount(self, cursor: PymssqlRawCursor) -> int: + return resolve_rowcount(cursor) + def begin(self) -> None: """Begin a transaction on the connection. @@ -314,6 +331,9 @@ def execute_many( ) cached_statement, prepared_parameters = self._compiled_statement(prepared_statement, config) parsed_expression = cached_statement.expression + if parsed_expression is None and statement.lstrip().upper().startswith("INSERT"): + with contextlib.suppress(Exception): + parsed_expression = sqlglot.parse_one(statement, read="tsql") if isinstance(parsed_expression, exp.Insert) and not parsed_expression.args.get("returning"): bulk_result = self._execute_bulk_insert_many(parsed_expression, prepared_parameters) if bulk_result is not None: @@ -382,10 +402,11 @@ def bulk_copy( row_list = list(rows) if not isinstance(rows, (list, tuple)) else rows if not row_list: return 0 + formatted_table = format_identifier(table_name) handler = self.handle_database_exceptions() with handler: self.connection.bulk_copy( - table_name, + formatted_table, row_list, column_ids=list(column_ids) if column_ids is not None else None, batch_size=batch_size, @@ -468,7 +489,7 @@ def _connection_in_transaction(self) -> bool: def _is_plain_values_insert(expression: exp.Insert, expected_columns: int) -> bool: - values = expression.args.get("values") + values = expression.expression if not isinstance(values, exp.Values): return False rows = values.expressions diff --git a/sqlspec/adapters/pymssql/litestar/store.py b/sqlspec/adapters/pymssql/litestar/store.py index 14df07634..7fd9f1ec8 100644 --- a/sqlspec/adapters/pymssql/litestar/store.py +++ b/sqlspec/adapters/pymssql/litestar/store.py @@ -3,6 +3,7 @@ from datetime import datetime, timedelta, timezone from typing import Any +from sqlspec.adapters.pymssql._typing import PymssqlCursor from sqlspec.adapters.pymssql.config import PymssqlConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -94,8 +95,11 @@ def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | No WHERE session_id = %s AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_session() as driver: - row = driver.select_one_or_none(sql, (key,)) + with self._config.provide_connection() as conn: + with PymssqlCursor(conn) as cursor: + cursor.execute(sql, (key,)) + row = cursor.fetchone() + if row is None: return None @@ -103,15 +107,16 @@ def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | No if renew_for is not None and expires_at is not None: new_expires_at = self._calculate_expires_at(renew_for) if new_expires_at is not None: - driver.execute( - f""" - UPDATE {self._table_name} - SET expires_at = %s, updated_at = SYSUTCDATETIME() - WHERE session_id = %s - """, - (new_expires_at, key), - ) - driver.commit() + with PymssqlCursor(conn) as update_cursor: + update_cursor.execute( + f""" + UPDATE {self._table_name} + SET expires_at = %s, updated_at = SYSUTCDATETIME() + WHERE session_id = %s + """, + (new_expires_at, key), + ) + conn.commit() return _coerce_bytes(_row_value(row, "data", 0)) @@ -131,19 +136,19 @@ def _set(self, key: str, value: str | bytes, expires_in: int | timedelta | None INSERT (session_id, data, expires_at) VALUES (src.session_id, src.data, src.expires_at); """ - with self._config.provide_session() as driver: - driver.execute(sql, (key, data, expires_at)) - driver.commit() + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, (key, data, expires_at)) + conn.commit() def _delete(self, key: str) -> None: - with self._config.provide_session() as driver: - driver.execute(f"DELETE FROM {self._table_name} WHERE session_id = %s", (key,)) - driver.commit() + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(f"DELETE FROM {self._table_name} WHERE session_id = %s", (key,)) + conn.commit() def _delete_all(self) -> None: - with self._config.provide_session() as driver: - driver.execute(f"TRUNCATE TABLE {self._table_name}") - driver.commit() + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(f"TRUNCATE TABLE {self._table_name}") + conn.commit() self._log_delete_all() def _exists(self, key: str) -> bool: @@ -153,12 +158,14 @@ def _exists(self, key: str) -> bool: WHERE session_id = %s AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ - with self._config.provide_session() as driver: - return driver.select_one_or_none(sql, (key,)) is not None + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql, (key,)) + return cursor.fetchone() is not None def _expires_in(self, key: str) -> int | None: - with self._config.provide_session() as driver: - row = driver.select_one_or_none(f"SELECT expires_at FROM {self._table_name} WHERE session_id = %s", (key,)) + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(f"SELECT expires_at FROM {self._table_name} WHERE session_id = %s", (key,)) + row = cursor.fetchone() if row is None: return None @@ -174,10 +181,10 @@ def _delete_expired(self) -> int: WHERE expires_at IS NOT NULL AND expires_at < SYSUTCDATETIME() """ - with self._config.provide_session() as driver: - res = driver.execute(sql) - driver.commit() - count = res.rows_affected + with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: + cursor.execute(sql) + count = int(getattr(cursor, "rowcount", 0) or 0) + conn.commit() if count > 0: self._log_delete_expired(count) return count diff --git a/tests/unit/adapters/test_mssql_python/test_config.py b/tests/unit/adapters/test_mssql_python/test_config.py index 44c786830..f7937a3b8 100644 --- a/tests/unit/adapters/test_mssql_python/test_config.py +++ b/tests/unit/adapters/test_mssql_python/test_config.py @@ -357,7 +357,7 @@ class FakeBindings: def close_pooling() -> None: closed_pooling.append(True) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.ddbc_bindings", FakeBindings) + monkeypatch.setattr(_mssql_pool.MSSQL_PYTHON_MODULE, "ddbc_bindings", FakeBindings, raising=False) pool = MssqlPythonConnectionPool(connection_string="Server=localhost;") pool.close(close_driver_pooling=True) assert closed_pooling == [True] diff --git a/tests/unit/adapters/test_pymssql/test_core.py b/tests/unit/adapters/test_pymssql/test_core.py index c24ab02b9..67d24143d 100644 --- a/tests/unit/adapters/test_pymssql/test_core.py +++ b/tests/unit/adapters/test_pymssql/test_core.py @@ -158,7 +158,7 @@ def test_quote_tsql_identifier() -> None: assert quote_tsql_identifier("users") == "[users]" assert quote_tsql_identifier("[users]") == "[users]" - assert quote_tsql_identifier("dbo.users") == "[dbo].[users]" + assert quote_tsql_identifier("dbo.users") == "[dbo.users]" assert quote_tsql_identifier("col]name") == "[col]]name]" @@ -189,11 +189,11 @@ def test_collect_rows_preserves_list_identity() -> None: input_rows = [(1, "Alice"), (2, "Bob")] description = [("id",), ("name",)] - rows, column_names, count = collect_rows(input_rows, description) + rows, column_names, row_format = collect_rows(input_rows, description) assert rows is input_rows assert column_names == ["id", "name"] - assert count == 2 + assert row_format == "tuple" def test_normalize_execute_parameters_preserves_tuples() -> None: From 041df9291002d9d30f22120456dbdd6155e960f0 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 02:43:36 +0000 Subject: [PATCH 03/11] fix(adapters): fix mssql-python parameterized prepare and source equivalence --- sqlspec/adapters/mssql_python/core.py | 2 +- sqlspec/adapters/mssql_python/driver.py | 14 ++++---------- 2 files changed, 5 insertions(+), 11 deletions(-) diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index 6a93ce950..3adef350b 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -150,7 +150,7 @@ def extract_error_number(exc: BaseException | None) -> int | None: def create_mapped_exception(error: Exception, *, logger: Logger | None = None) -> SQLSpecError: """Map a mssql-python exception to SQLSpec's exception hierarchy.""" - error_number = _extract_error_number(error) + error_number = extract_error_number(error) if error_number == _MSSQL_CONSTRAINT_547: message = str(error) if "check constraint" in message.lower(): diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index 3b67e26b8..35f929d1b 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -296,7 +296,7 @@ def reset_migration_session_schema(self) -> None: def has_schema(self, schema: str) -> bool: """Return whether the specified schema exists.""" with self.with_cursor(self.connection) as cursor: - _execute_cursor(cursor, "SELECT 1 FROM sys.schemas WHERE name = ?", (schema,), use_prepare=False) + _execute_cursor(cursor, "SELECT 1 FROM sys.schemas WHERE name = ?", (schema,)) return cursor.fetchone() is not None def select_to_arrow( @@ -554,24 +554,18 @@ def _quote_mssql_table(table: str) -> str: def _execute_cursor(cursor: MssqlPythonRawCursor, sql: str, parameters: Any, *, use_prepare: bool = True) -> None: - if use_prepare: + if use_prepare or parameters: if parameters is None: cursor.execute(sql) else: cursor.execute(sql, parameters) return try: - if parameters is None: - cursor.execute(sql, use_prepare=False) - else: - cursor.execute(sql, parameters, use_prepare=False) + cursor.execute(sql, use_prepare=False) except TypeError as exc: if "use_prepare" not in str(exc): raise - if parameters is None: - cursor.execute(sql) - else: - cursor.execute(sql, parameters) + cursor.execute(sql) def _cursor_rowcount(cursor: MssqlPythonRawCursor) -> int: From d3e32e6427327afa5f20a83eed80426f0518f2bb Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 17:14:21 +0000 Subject: [PATCH 04/11] refactor(mssql): remove redundant private aliases --- sqlspec/adapters/mssql_python/_typing.py | 6 ++--- sqlspec/adapters/mssql_python/core.py | 3 --- sqlspec/adapters/mssql_python/pool.py | 12 ++++----- sqlspec/adapters/pymssql/_typing.py | 6 ++--- sqlspec/adapters/pymssql/core.py | 4 --- sqlspec/adapters/pymssql/pool.py | 4 +-- .../adapters/test_mssql_python/test_config.py | 26 +++++++++---------- .../adapters/test_mssql_python/test_core.py | 8 +++--- .../unit/adapters/test_pymssql/test_config.py | 4 +-- 9 files changed, 30 insertions(+), 43 deletions(-) diff --git a/sqlspec/adapters/mssql_python/_typing.py b/sqlspec/adapters/mssql_python/_typing.py index 5f006110c..5c6c4879b 100644 --- a/sqlspec/adapters/mssql_python/_typing.py +++ b/sqlspec/adapters/mssql_python/_typing.py @@ -3,13 +3,11 @@ import contextlib from typing import TYPE_CHECKING, Any -import mssql_python as _mssql_python +import mssql_python as mssql_python_module from mssql_python import Error as MssqlPythonError from mssql_python.connection import Connection, TokenProvider from mssql_python.cursor import Cursor -MSSQL_PYTHON_MODULE: Any = _mssql_python - if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType @@ -26,13 +24,13 @@ MssqlPythonRawCursor = Cursor __all__ = ( - "MSSQL_PYTHON_MODULE", "MssqlPythonConnection", "MssqlPythonCursor", "MssqlPythonError", "MssqlPythonRawCursor", "MssqlPythonSessionContext", "TokenProvider", + "mssql_python_module", ) diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index 3adef350b..c89a16a8d 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -145,9 +145,6 @@ def extract_error_number(exc: BaseException | None) -> int | None: return None -_extract_error_number = extract_error_number - - def create_mapped_exception(error: Exception, *, logger: Logger | None = None) -> SQLSpecError: """Map a mssql-python exception to SQLSpec's exception hierarchy.""" error_number = extract_error_number(error) diff --git a/sqlspec/adapters/mssql_python/pool.py b/sqlspec/adapters/mssql_python/pool.py index db48990c8..a9527e7f5 100644 --- a/sqlspec/adapters/mssql_python/pool.py +++ b/sqlspec/adapters/mssql_python/pool.py @@ -3,9 +3,9 @@ import contextlib import warnings from collections.abc import Callable -from typing import Any, cast +from typing import Any -from sqlspec.adapters.mssql_python._typing import MSSQL_PYTHON_MODULE, MssqlPythonConnection +from sqlspec.adapters.mssql_python._typing import MssqlPythonConnection, mssql_python_module __all__ = ("MssqlPythonConnectionPool",) @@ -51,16 +51,14 @@ def __init__( stacklevel=2, ) if _POOLING_PARAMS is None or new_params != _POOLING_PARAMS: - MSSQL_PYTHON_MODULE.pooling(max_size=max_size, idle_timeout=idle_timeout, enabled=enabled) + mssql_python_module.pooling(max_size=max_size, idle_timeout=idle_timeout, enabled=enabled) _POOLING_PARAMS = new_params def acquire(self) -> MssqlPythonConnection: if self._closed: msg = "Cannot acquire a connection from a closed mssql-python pool." raise RuntimeError(msg) - connection = cast( - "MssqlPythonConnection", MSSQL_PYTHON_MODULE.connect(self.connection_string, **self.connect_kwargs) - ) + connection = mssql_python_module.connect(self.connection_string, **self.connect_kwargs) if self.on_connection_create is not None: self.on_connection_create(connection) return connection @@ -74,6 +72,6 @@ def close(self, *, close_driver_pooling: bool = False) -> None: global _POOLING_PARAMS _POOLING_PARAMS = None with contextlib.suppress(Exception): - ddbc = getattr(MSSQL_PYTHON_MODULE, "ddbc_bindings", None) + ddbc = getattr(mssql_python_module, "ddbc_bindings", None) if ddbc is not None and hasattr(ddbc, "close_pooling"): ddbc.close_pooling() diff --git a/sqlspec/adapters/pymssql/_typing.py b/sqlspec/adapters/pymssql/_typing.py index ceb7c8433..4e4b6940f 100644 --- a/sqlspec/adapters/pymssql/_typing.py +++ b/sqlspec/adapters/pymssql/_typing.py @@ -7,13 +7,11 @@ import contextlib from typing import TYPE_CHECKING, Any -import pymssql as _pymssql +import pymssql as pymssql_module from pymssql import Connection as _PymssqlConnection from pymssql import Cursor as _PymssqlRawCursor from pymssql import Error as PymssqlError -PYMSSQL_MODULE = _pymssql - if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType @@ -33,13 +31,13 @@ PymssqlRawCursor = _PymssqlRawCursor __all__ = ( - "PYMSSQL_MODULE", "PymssqlConnection", "PymssqlCursor", "PymssqlError", "PymssqlQueryParams", "PymssqlRawCursor", "PymssqlSessionContext", + "pymssql_module", ) diff --git a/sqlspec/adapters/pymssql/core.py b/sqlspec/adapters/pymssql/core.py index cab399ae4..77f5eeed4 100644 --- a/sqlspec/adapters/pymssql/core.py +++ b/sqlspec/adapters/pymssql/core.py @@ -307,10 +307,6 @@ def extract_error_number(exc: BaseException | None) -> int | None: return None -_extract_error_number = extract_error_number -_quote_bracket_identifier = quote_tsql_identifier - - driver_profile = build_profile() default_statement_config = build_statement_config() diff --git a/sqlspec/adapters/pymssql/pool.py b/sqlspec/adapters/pymssql/pool.py index a89b17e46..8c8ac29bc 100644 --- a/sqlspec/adapters/pymssql/pool.py +++ b/sqlspec/adapters/pymssql/pool.py @@ -7,7 +7,8 @@ from contextlib import contextmanager from typing import TYPE_CHECKING, Any, cast -from sqlspec.adapters.pymssql._typing import PYMSSQL_MODULE, PymssqlConnection +from sqlspec.adapters.pymssql._typing import PymssqlConnection +from sqlspec.adapters.pymssql._typing import pymssql_module as pymssql from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -19,7 +20,6 @@ logger = get_logger(POOL_LOGGER_NAME) _ADAPTER_NAME = "pymssql" -pymssql = PYMSSQL_MODULE class PymssqlConnectionPool: diff --git a/tests/unit/adapters/test_mssql_python/test_config.py b/tests/unit/adapters/test_mssql_python/test_config.py index f7937a3b8..9db77f7f7 100644 --- a/tests/unit/adapters/test_mssql_python/test_config.py +++ b/tests/unit/adapters/test_mssql_python/test_config.py @@ -55,8 +55,8 @@ def fake_connect(connection_string: str, **kwargs: Any) -> DummyConnection: calls.append(("connect", (connection_string,), kwargs)) return connection - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", fake_pooling) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", fake_connect) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", fake_pooling) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", fake_connect) pool = MssqlPythonConnectionPool( connection_string="Server=localhost;", connect_kwargs={"timeout": 5}, max_size=7, idle_timeout=30, enabled=True @@ -79,7 +79,7 @@ def test_config_create_pool_splits_connection_and_pool_options(monkeypatch: pyte def fake_pooling(**kwargs: Any) -> None: pooling_calls.append(kwargs) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", fake_pooling) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", fake_pooling) config = MssqlPythonConfig( connection_config={ @@ -101,7 +101,7 @@ def fake_pooling(**kwargs: Any) -> None: def test_config_connection_string_with_discrete_override(monkeypatch: pytest.MonkeyPatch) -> None: """MssqlPythonConfig should merge discrete database overrides over connection_string.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **kw: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **kw: None) config = MssqlPythonConfig( connection_config={ @@ -161,7 +161,7 @@ def test_config_create_pool_normalizes_current_odbc_aliases(monkeypatch: pytest. def fake_pooling(**kwargs: Any) -> None: pooling_calls.append(kwargs) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", fake_pooling) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", fake_pooling) config = MssqlPythonConfig( connection_config={ @@ -224,9 +224,9 @@ def test_config_connection_hook_runs_for_session_connections(monkeypatch: pytest seen: list[DummyConnection] = [] monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **_: None) monkeypatch.setattr( - "sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", lambda *_args, **_kwargs: connection + "sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", lambda *_args, **_kwargs: connection ) config = MssqlPythonConfig( @@ -242,9 +242,9 @@ def test_config_connection_hook_runs_for_session_connections(monkeypatch: pytest def test_second_pool_warns_on_different_params(monkeypatch: pytest.MonkeyPatch) -> None: """A second pool with different process-wide pooling params emits one warning.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **_: None) monkeypatch.setattr( - "sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", lambda *_args, **_kwargs: DummyConnection() + "sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", lambda *_args, **_kwargs: DummyConnection() ) MssqlPythonConnectionPool(connection_string="Server=srv1;", max_size=10, idle_timeout=60, enabled=True) @@ -263,9 +263,9 @@ def test_second_pool_warns_on_different_params(monkeypatch: pytest.MonkeyPatch) def test_second_pool_same_params_no_warn(monkeypatch: pytest.MonkeyPatch) -> None: """A second pool with identical process-wide pooling params emits no warning.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **_: None) monkeypatch.setattr( - "sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", lambda *_args, **_kwargs: DummyConnection() + "sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", lambda *_args, **_kwargs: DummyConnection() ) MssqlPythonConnectionPool(connection_string="Server=srv1;", max_size=10, idle_timeout=60, enabled=True) @@ -357,7 +357,7 @@ class FakeBindings: def close_pooling() -> None: closed_pooling.append(True) - monkeypatch.setattr(_mssql_pool.MSSQL_PYTHON_MODULE, "ddbc_bindings", FakeBindings, raising=False) + monkeypatch.setattr(_mssql_pool.mssql_python_module, "ddbc_bindings", FakeBindings, raising=False) pool = MssqlPythonConnectionPool(connection_string="Server=localhost;") pool.close(close_driver_pooling=True) assert closed_pooling == [True] @@ -366,7 +366,7 @@ def close_pooling() -> None: def test_pool_suppresses_warning_when_params_match(monkeypatch: pytest.MonkeyPatch) -> None: """Pool reconfiguration should not warn if params are identical to previous.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", (10, 60, True)) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **kw: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **kw: None) with warnings.catch_warnings(record=True) as recorded: warnings.simplefilter("always") diff --git a/tests/unit/adapters/test_mssql_python/test_core.py b/tests/unit/adapters/test_mssql_python/test_core.py index d6a1a1f86..ec8cdf955 100644 --- a/tests/unit/adapters/test_mssql_python/test_core.py +++ b/tests/unit/adapters/test_mssql_python/test_core.py @@ -2,7 +2,7 @@ import pytest -from sqlspec.adapters.mssql_python._typing import MSSQL_PYTHON_MODULE +from sqlspec.adapters.mssql_python._typing import mssql_python_module from sqlspec.adapters.mssql_python.core import build_connection_config, create_mapped_exception, extract_error_number from sqlspec.exceptions import ( CheckViolationError, @@ -65,7 +65,7 @@ def test_build_connection_config_no_duplicate_pwd() -> None: def test_create_mapped_exception_extracts_sql_server_error_number() -> None: """SQL Server native error numbers should map to specific SQLSpec exceptions.""" - exc = MSSQL_PYTHON_MODULE.IntegrityError( + exc = mssql_python_module.IntegrityError( "23000", "[23000] [Microsoft][ODBC Driver 18 for SQL Server][SQL Server]Violation of UNIQUE KEY constraint. (2627)", ) @@ -78,7 +78,7 @@ def test_create_mapped_exception_extracts_sql_server_error_number() -> None: def test_create_mapped_exception_falls_back_for_connection_errors() -> None: """Known connection error numbers should map to DatabaseConnectionError.""" - exc = MSSQL_PYTHON_MODULE.OperationalError( + exc = mssql_python_module.OperationalError( "08001", "[08001] [Microsoft][ODBC Driver 18 for SQL Server]Named Pipes Provider: " "Could not open a connection to SQL Server (53)", @@ -140,7 +140,7 @@ def test_create_mapped_exception_classifies_constraint_messages_without_error_nu message: str, expected_type: type[Exception] ) -> None: """Constraint messages remain classifiable when the driver omits SQL Server error numbers.""" - mapped = create_mapped_exception(MSSQL_PYTHON_MODULE.IntegrityError("23000", message)) + mapped = create_mapped_exception(mssql_python_module.IntegrityError("23000", message)) assert isinstance(mapped, expected_type) diff --git a/tests/unit/adapters/test_pymssql/test_config.py b/tests/unit/adapters/test_pymssql/test_config.py index a8f7c589f..cffd9fb35 100644 --- a/tests/unit/adapters/test_pymssql/test_config.py +++ b/tests/unit/adapters/test_pymssql/test_config.py @@ -122,11 +122,11 @@ def test_pymssql_runtime_aliases_resolve_to_installed_classes() -> None: """pymssql public runtime aliases should expose installed pymssql classes.""" pymssql = pytest.importorskip("pymssql") from sqlspec.adapters.pymssql import PymssqlConnection as PublicPymssqlConnection - from sqlspec.adapters.pymssql._typing import PYMSSQL_MODULE, PymssqlConnection, PymssqlRawCursor + from sqlspec.adapters.pymssql._typing import PymssqlConnection, PymssqlRawCursor, pymssql_module namespace = PymssqlConfig().get_signature_namespace() - assert PYMSSQL_MODULE is pymssql + assert pymssql_module is pymssql assert PymssqlConnection is pymssql.Connection assert PublicPymssqlConnection is pymssql.Connection assert PymssqlRawCursor is pymssql.Cursor From 1188a5ba1e3a7a5546757a8e88c95820a9bf6afe Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 20:32:11 +0000 Subject: [PATCH 05/11] refactor(mssql): simplify _typing.py by removing redundant TYPE_CHECKING splits --- sqlspec/adapters/mssql_python/_typing.py | 13 +++---------- sqlspec/adapters/pymssql/_typing.py | 16 ++-------------- sqlspec/adapters/pymssql/driver.py | 3 ++- 3 files changed, 7 insertions(+), 25 deletions(-) diff --git a/sqlspec/adapters/mssql_python/_typing.py b/sqlspec/adapters/mssql_python/_typing.py index 5c6c4879b..bf9b1dfd5 100644 --- a/sqlspec/adapters/mssql_python/_typing.py +++ b/sqlspec/adapters/mssql_python/_typing.py @@ -5,24 +5,17 @@ import mssql_python as mssql_python_module from mssql_python import Error as MssqlPythonError -from mssql_python.connection import Connection, TokenProvider -from mssql_python.cursor import Cursor +from mssql_python.connection import Connection as MssqlPythonConnection +from mssql_python.connection import TokenProvider +from mssql_python.cursor import Cursor as MssqlPythonRawCursor if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType - from typing import TypeAlias from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver from sqlspec.core import StatementConfig - MssqlPythonConnection: TypeAlias = Connection - MssqlPythonRawCursor: TypeAlias = Cursor - -if not TYPE_CHECKING: - MssqlPythonConnection = Connection - MssqlPythonRawCursor = Cursor - __all__ = ( "MssqlPythonConnection", "MssqlPythonCursor", diff --git a/sqlspec/adapters/pymssql/_typing.py b/sqlspec/adapters/pymssql/_typing.py index 4e4b6940f..de7afc100 100644 --- a/sqlspec/adapters/pymssql/_typing.py +++ b/sqlspec/adapters/pymssql/_typing.py @@ -8,33 +8,21 @@ from typing import TYPE_CHECKING, Any import pymssql as pymssql_module -from pymssql import Connection as _PymssqlConnection -from pymssql import Cursor as _PymssqlRawCursor +from pymssql import Connection as PymssqlConnection +from pymssql import Cursor as PymssqlRawCursor from pymssql import Error as PymssqlError if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType - from typing import TypeAlias - - from pymssql._pymssql import QueryParams as PymssqlQueryParams from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.core import StatementConfig - PymssqlConnection: TypeAlias = _PymssqlConnection - PymssqlRawCursor: TypeAlias = _PymssqlRawCursor - -if not TYPE_CHECKING: - PymssqlQueryParams = Any - PymssqlConnection = _PymssqlConnection - PymssqlRawCursor = _PymssqlRawCursor - __all__ = ( "PymssqlConnection", "PymssqlCursor", "PymssqlError", - "PymssqlQueryParams", "PymssqlRawCursor", "PymssqlSessionContext", "pymssql_module", diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 5784494f5..22658251b 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -44,7 +44,8 @@ from sqlspec.utils.logging import get_logger if TYPE_CHECKING: - from sqlspec.adapters.pymssql._typing import PymssqlQueryParams as QueryParams + from pymssql._pymssql import QueryParams + from sqlspec.builder import QueryBuilder from sqlspec.core import Statement, StatementFilter from sqlspec.typing import StatementParameters From 82227bd38f93f79bffe2b18ec93436eb438e4131 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 15:25:51 +0000 Subject: [PATCH 06/11] fix(mssql): fix event index regex (#816), encapsulate _bulk_copy, and move multi-row insert to dispatch_execute_many --- sqlspec/adapters/mssql_python/__init__.py | 7 +- .../adapters/mssql_python/data_dictionary.py | 52 +------- sqlspec/adapters/mssql_python/driver.py | 45 ++++--- sqlspec/adapters/mssql_python/events/store.py | 2 +- .../adapters/mssql_python/type_converter.py | 12 +- sqlspec/adapters/pymssql/_typing.py | 3 - sqlspec/adapters/pymssql/core.py | 45 +++++-- sqlspec/adapters/pymssql/data_dictionary.py | 52 +------- sqlspec/adapters/pymssql/driver.py | 122 ++++++------------ sqlspec/adapters/pymssql/events/store.py | 2 +- .../dialects/mssql/__init__.py | 4 + .../data_dictionary/dialects/mssql/config.py | 57 +++++++- .../adapters/test_mssql_python/test_arrow.py | 29 ++++- .../test_bulk_copy_result.py | 2 +- .../test_mssql_python/test_events_store.py | 5 + .../unit/adapters/test_pymssql/test_config.py | 4 +- tests/unit/adapters/test_pymssql/test_core.py | 43 +++--- .../test_pymssql/test_data_dictionary.py | 12 ++ .../unit/adapters/test_pymssql/test_driver.py | 86 ++++-------- .../adapters/test_pymssql/test_extensions.py | 4 + 20 files changed, 272 insertions(+), 316 deletions(-) diff --git a/sqlspec/adapters/mssql_python/__init__.py b/sqlspec/adapters/mssql_python/__init__.py index 4964a7b8a..a098ed261 100644 --- a/sqlspec/adapters/mssql_python/__init__.py +++ b/sqlspec/adapters/mssql_python/__init__.py @@ -9,17 +9,12 @@ ) from sqlspec.adapters.mssql_python.core import default_statement_config, driver_profile from sqlspec.adapters.mssql_python.data_dictionary import MssqlPythonSyncDataDictionary, MssqlVersionInfo -from sqlspec.adapters.mssql_python.driver import ( - MssqlPythonBulkCopyResult, - MssqlPythonDriver, - MssqlPythonExceptionHandler, -) +from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver, MssqlPythonExceptionHandler from sqlspec.adapters.mssql_python.migrations import MssqlPythonSyncMigrationTracker from sqlspec.adapters.mssql_python.pool import MssqlPythonConnectionPool from sqlspec.adapters.mssql_python.type_converter import MssqlPythonTypeConverter, mssql_type_to_arrow __all__ = ( - "MssqlPythonBulkCopyResult", "MssqlPythonConfig", "MssqlPythonConnection", "MssqlPythonConnectionParams", diff --git a/sqlspec/adapters/mssql_python/data_dictionary.py b/sqlspec/adapters/mssql_python/data_dictionary.py index cb4545719..f4a2bd50e 100644 --- a/sqlspec/adapters/mssql_python/data_dictionary.py +++ b/sqlspec/adapters/mssql_python/data_dictionary.py @@ -1,6 +1,6 @@ """mssql-python data dictionary.""" -from typing import TYPE_CHECKING, Any, ClassVar, Final, cast +from typing import TYPE_CHECKING, Any, ClassVar, cast from mypy_extensions import mypyc_attr @@ -14,23 +14,21 @@ SystemMetadataRequest, SystemMetadataResult, TableMetadata, - VersionInfo, ensure_system_metadata_request, get_data_dictionary_loader, get_dialect_config, system_metadata_gated_result, ) from sqlspec.data_dictionary.dialects.mssql import ( + MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, build_mssql_table_ddl_result, extract_mssql_version_value, get_mssql_data_dictionary_options, - is_mssql_azure_sql, list_mssql_available_features, merge_mssql_table_lists, - mssql_supports_native_json, parse_mssql_engine_edition, parse_mssql_version_components, resolve_mssql_feature_flag, @@ -50,50 +48,6 @@ logger = get_logger("sqlspec.adapters.mssql_python.data_dictionary") -MSSQL_VECTOR_MIN_MAJOR: Final[int] = 17 - - -class MssqlVersionInfo(VersionInfo): - """MSSQL database version info with build, revision, and Azure SQL detection.""" - - def __init__( - self, - major: int, - minor: int = 0, - build: int = 0, - revision: int = 0, - edition: str | None = None, - engine_edition: int | None = None, - ) -> None: - super().__init__(major, minor, 0) - self.build = build - self.revision = revision - self.edition = edition - self.engine_edition = engine_edition - self.is_azure_sql = is_mssql_azure_sql(engine_edition) - - def supports_native_json(self) -> bool: - """Return whether this server supports the native JSON type.""" - return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) - - def supports_vector(self) -> bool: - """Return whether this server supports native VECTOR data types and functions.""" - return self.is_azure_sql or self.major >= MSSQL_VECTOR_MIN_MAJOR - - @property - def version_tuple(self) -> "tuple[int, int, int]": - """Get version tuple using the MSSQL build number as the third component.""" - return (self.major, self.minor, self.build) - - def __str__(self) -> str: - """String representation of version info.""" - version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" - if self.edition: - version_str += f" ({self.edition})" - if self.is_azure_sql: - version_str += " [Azure]" - return version_str - class _MssqlDataDictionaryMixin: """Shared helpers for MSSQL data dictionaries.""" @@ -196,8 +150,6 @@ def get_version(self, driver: "MssqlPythonDriver") -> MssqlVersionInfo | None: def get_feature_flag(self, driver: "MssqlPythonDriver", feature: str) -> bool: """Check whether SQL Server supports a feature.""" version_info = self.get_version(driver) - if feature == "supports_vector": - return bool(version_info and version_info.supports_vector()) return resolve_mssql_feature_flag( feature, major=version_info.major if version_info is not None else 0, diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index 35f929d1b..f12dba242 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -49,13 +49,7 @@ from sqlspec.core import Statement, StatementFilter from sqlspec.typing import ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters -__all__ = ( - "MssqlPythonBulkCopyResult", - "MssqlPythonCursor", - "MssqlPythonDriver", - "MssqlPythonExceptionHandler", - "MssqlPythonSessionContext", -) +__all__ = ("MssqlPythonCursor", "MssqlPythonDriver", "MssqlPythonExceptionHandler", "MssqlPythonSessionContext") logger = get_logger("sqlspec.adapters.mssql_python") _COLUMN_CACHE_MAX_SIZE = 256 @@ -371,7 +365,7 @@ def select_to_arrow( prepared_statement, table, return_format=return_format, batch_size=batch_size, arrow_schema=arrow_schema ) - def bulk_copy( + def _bulk_copy( self, target_table: str, rows: Iterable[tuple[Any, ...]], @@ -479,24 +473,43 @@ def load_from_arrow( telemetry_payload = cast("StorageTelemetry", {"destination": table, "format": "arrow", "extra": {}}) else: arrow_table = self._coerce_arrow_table(source) - cols = column_mappings or list(arrow_table.column_names) + cols = column_mappings if column_mappings is not None else list(arrow_table.column_names) if arrow_table.num_rows: exc_handler = self.handle_database_exceptions() + use_fallback = False with exc_handler, self.with_cursor(self.connection) as cursor: - raw_result = cursor.bulkcopy_arrow( + if hasattr(cursor, "bulkcopy_arrow"): + raw_result = cursor.bulkcopy_arrow( + table, + arrow_table, + batch_size=batch_size, + timeout=timeout, + table_lock=table_lock, + check_constraints=check_constraints, + fire_triggers=fire_triggers, + keep_identity=keep_identity, + keep_nulls=keep_nulls, + use_internal_transaction=use_internal_transaction, + column_mappings=cols, + ) + else: + use_fallback = True + self._check_pending_exception(exc_handler) + if use_fallback: + _, records = self._arrow_table_to_rows(arrow_table) + raw_result = self._bulk_copy( table, - arrow_table, + records, batch_size=batch_size, timeout=timeout, - table_lock=table_lock, - check_constraints=check_constraints, - fire_triggers=fire_triggers, + column_mappings=cols, keep_identity=keep_identity, + check_constraints=check_constraints, + table_lock=table_lock, keep_nulls=keep_nulls, + fire_triggers=fire_triggers, use_internal_transaction=use_internal_transaction, - column_mappings=cols, ) - self._check_pending_exception(exc_handler) telemetry_payload = self._ingest_telemetry(arrow_table) extra = telemetry_payload.setdefault("extra", {}) diff --git a/sqlspec/adapters/mssql_python/events/store.py b/sqlspec/adapters/mssql_python/events/store.py index 09bee9b9f..854a9f22f 100644 --- a/sqlspec/adapters/mssql_python/events/store.py +++ b/sqlspec/adapters/mssql_python/events/store.py @@ -38,7 +38,7 @@ def _wrap_create_statement(self, statement: str, object_type: str) -> str: table_name = match.group(1) return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NULL BEGIN {statement}; END" if object_type == "index": - match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+(\S+)", statement, re.IGNORECASE) + match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+([^\s(]+)", statement, re.IGNORECASE) if match: index_name = match.group(1).strip("[]") table_name = match.group(2) diff --git a/sqlspec/adapters/mssql_python/type_converter.py b/sqlspec/adapters/mssql_python/type_converter.py index a1d482413..cf6ced2fe 100644 --- a/sqlspec/adapters/mssql_python/type_converter.py +++ b/sqlspec/adapters/mssql_python/type_converter.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Any, Final, cast from uuid import UUID -from sqlspec.utils.module_loader import ensure_pyarrow +from sqlspec.utils.module_loader import ensure_pyarrow, import_optional from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: @@ -79,9 +79,8 @@ def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: i normalized_type = sql_type.lower().split("(", 1)[0].strip() if normalized_type == "vector": ensure_pyarrow() - import pyarrow as pa - - return cast("pa.DataType", pa.list_(pa.float32())) + pyarrow_mod = cast("Any", import_optional("pyarrow")) + return cast("pa.DataType", pyarrow_mod.list_(pyarrow_mod.float32())) if normalized_type in {"decimal", "numeric"} and precision is not None and scale is not None: return _arrow_type("decimal128", (precision, scale)) spec = _MSSQL_ARROW_TYPE_SPECS.get(normalized_type) @@ -93,6 +92,5 @@ def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: i def _arrow_type(name: str, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None) -> "pa.DataType": ensure_pyarrow() - import pyarrow as pa - - return cast("pa.DataType", getattr(pa, name)(*args, **(kwargs or {}))) + pyarrow_mod = cast("Any", import_optional("pyarrow")) + return cast("pa.DataType", getattr(pyarrow_mod, name)(*args, **(kwargs or {}))) diff --git a/sqlspec/adapters/pymssql/_typing.py b/sqlspec/adapters/pymssql/_typing.py index de7afc100..9c6cea9cd 100644 --- a/sqlspec/adapters/pymssql/_typing.py +++ b/sqlspec/adapters/pymssql/_typing.py @@ -89,9 +89,6 @@ def __enter__(self) -> "PymssqlDriver": def __exit__( self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" ) -> "bool | None": - if exc_type is not None and self._driver is not None: - with contextlib.suppress(Exception): - self._driver.rollback() if self._connection is not None: self._release_connection(self._connection, exc_type=exc_type, exc_val=exc_val, exc_tb=exc_tb) self._connection = None diff --git a/sqlspec/adapters/pymssql/core.py b/sqlspec/adapters/pymssql/core.py index 77f5eeed4..aae8a71b2 100644 --- a/sqlspec/adapters/pymssql/core.py +++ b/sqlspec/adapters/pymssql/core.py @@ -5,6 +5,8 @@ from logging import Logger from typing import Any, Final, Literal +from sqlglot import exp + from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile from sqlspec.exceptions import ( CheckViolationError, @@ -38,6 +40,7 @@ "driver_profile", "extract_error_number", "format_identifier", + "is_plain_values_insert", "normalize_execute_many_parameters", "normalize_execute_parameters", "quote_tsql_identifier", @@ -46,7 +49,7 @@ "resolve_rowcount", ) -_ERROR_NUMBER_PATTERN: Final[re.Pattern[str]] = re.compile(r"\(([-]?\d+)(?:,|\))") +_ERROR_NUMBER_PATTERN: Final[re.Pattern[str]] = re.compile(r"(?:\(([-]?\d+)(?:,|\))|\bMsg\s+([-]?\d+)\b)") _MSSQL_CONSTRAINT_547: Final[int] = 547 _COLUMN_CACHE_MAX_SIZE: Final[int] = 256 _ERROR_CODE_MAPPING: Final[dict[int, tuple[type[SQLSpecError], str]]] = { @@ -88,21 +91,43 @@ def build_insert_statement(table: str, columns: list[str]) -> str: return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES ({placeholders})" -def build_multi_row_insert(table: str, columns: list[str], num_rows: int) -> str: +def build_multi_row_insert(table: str, columns: Sequence[str], num_rows: int, *, num_columns: int | None = None) -> str: """Build a multi-row VALUES (...), (...) batch INSERT statement. Args: table: Target table name. - columns: Column names to insert. + columns: Column names to insert (empty when inserting into all table columns). num_rows: Number of row tuples in the VALUES clause (up to 1,000). + num_columns: Explicit column count when ``columns`` is empty. Returns: Parameterized T-SQL INSERT statement. """ - column_clause = ", ".join(quote_tsql_identifier(column) for column in columns) - single_row = f"({', '.join('%s' for _ in columns)})" + col_count = len(columns) if columns else (num_columns or 0) + single_row = f"({', '.join('%s' for _ in range(col_count))})" values_clause = ", ".join(single_row for _ in range(num_rows)) - return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES {values_clause}" + if columns: + column_clause = ", ".join(quote_tsql_identifier(column) for column in columns) + return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES {values_clause}" + return f"INSERT INTO {format_identifier(table)} VALUES {values_clause}" + + +def is_plain_values_insert(expression: Any, expected_columns: int) -> bool: + """Return whether a parsed INSERT expression is a single-row plain VALUES insert without OUTPUT/RETURNING.""" + if not isinstance(expression, exp.Insert): + return False + if expression.args.get("output") or expression.args.get("returning"): + return False + values = expression.expression + if not isinstance(values, exp.Values): + return False + rows = values.expressions + if len(rows) != 1: + return False + row = rows[0] + if not isinstance(row, exp.Tuple): + return False + return len(row.expressions) == expected_columns def normalize_execute_parameters(parameters: Any) -> Any: @@ -292,16 +317,18 @@ def extract_error_number(exc: BaseException | None) -> int | None: return None for attr in ("number", "error_code", "errno"): val = getattr(exc, attr, None) - if isinstance(val, int) and val != 0: + if isinstance(val, int) and not isinstance(val, bool) and val != 0: return val if hasattr(exc, "args") and exc.args: first = exc.args[0] - if isinstance(first, int): + if isinstance(first, int) and not isinstance(first, bool): return first matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) if matches: + last_match = matches[-1] + raw_num = last_match[0] or last_match[1] if isinstance(last_match, tuple) else last_match try: - return int(matches[-1]) + return int(raw_num) except ValueError: pass return None diff --git a/sqlspec/adapters/pymssql/data_dictionary.py b/sqlspec/adapters/pymssql/data_dictionary.py index 9d5a5de9e..b257e4f60 100644 --- a/sqlspec/adapters/pymssql/data_dictionary.py +++ b/sqlspec/adapters/pymssql/data_dictionary.py @@ -1,6 +1,6 @@ """pymssql data dictionary.""" -from typing import TYPE_CHECKING, Any, ClassVar, Final, cast +from typing import TYPE_CHECKING, Any, ClassVar, cast from mypy_extensions import mypyc_attr @@ -14,23 +14,21 @@ SystemMetadataRequest, SystemMetadataResult, TableMetadata, - VersionInfo, ensure_system_metadata_request, get_data_dictionary_loader, get_dialect_config, system_metadata_gated_result, ) from sqlspec.data_dictionary.dialects.mssql import ( + MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, build_mssql_table_ddl_result, extract_mssql_version_value, get_mssql_data_dictionary_options, - is_mssql_azure_sql, list_mssql_available_features, merge_mssql_table_lists, - mssql_supports_native_json, parse_mssql_engine_edition, parse_mssql_version_components, resolve_mssql_feature_flag, @@ -50,50 +48,6 @@ logger = get_logger("sqlspec.adapters.pymssql.data_dictionary") -MSSQL_VECTOR_MIN_MAJOR: Final[int] = 17 - - -class MssqlVersionInfo(VersionInfo): - """MSSQL database version info with build, revision, and Azure SQL detection.""" - - def __init__( - self, - major: int, - minor: int = 0, - build: int = 0, - revision: int = 0, - edition: str | None = None, - engine_edition: int | None = None, - ) -> None: - super().__init__(major, minor, 0) - self.build = build - self.revision = revision - self.edition = edition - self.engine_edition = engine_edition - self.is_azure_sql = is_mssql_azure_sql(engine_edition) - - def supports_native_json(self) -> bool: - """Return whether this server supports the native JSON type.""" - return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) - - def supports_vector(self) -> bool: - """Return whether this server supports native VECTOR data types and functions.""" - return self.is_azure_sql or self.major >= MSSQL_VECTOR_MIN_MAJOR - - @property - def version_tuple(self) -> tuple[int, int, int]: - """Get version tuple using the MSSQL build number as the third component.""" - return (self.major, self.minor, self.build) - - def __str__(self) -> str: - """String representation of version info.""" - version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" - if self.edition: - version_str += f" ({self.edition})" - if self.is_azure_sql: - version_str += " [Azure]" - return version_str - class _MssqlDataDictionaryMixin: """Shared helpers for MSSQL data dictionaries.""" @@ -196,8 +150,6 @@ def get_version(self, driver: "PymssqlDriver") -> MssqlVersionInfo | None: def get_feature_flag(self, driver: "PymssqlDriver", feature: str) -> bool: """Check whether SQL Server supports a feature.""" version_info = self.get_version(driver) - if feature == "supports_vector": - return bool(version_info and version_info.supports_vector()) return resolve_mssql_feature_flag( feature, major=version_info.major if version_info is not None else 0, diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 22658251b..00c57827b 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -21,6 +21,7 @@ default_statement_config, driver_profile, format_identifier, + is_plain_values_insert, normalize_execute_many_parameters, normalize_execute_parameters, quote_tsql_identifier, @@ -30,7 +31,6 @@ ) from sqlspec.adapters.pymssql.data_dictionary import PymssqlSyncDataDictionary from sqlspec.core import SQL, ArrowResult, StatementConfig, get_cache_config, register_driver_profile -from sqlspec.core.result import DMLResult, SQLResult from sqlspec.driver import ( BaseSyncExceptionHandler, ExecutionResult, @@ -46,10 +46,6 @@ if TYPE_CHECKING: from pymssql._pymssql import QueryParams - from sqlspec.builder import QueryBuilder - from sqlspec.core import Statement, StatementFilter - from sqlspec.typing import StatementParameters - __all__ = ("PymssqlCursor", "PymssqlDriver", "PymssqlExceptionHandler", "PymssqlSessionContext") logger = get_logger("sqlspec.adapters.pymssql") @@ -144,15 +140,6 @@ def __init__( statement_config = default_statement_config.replace( enable_caching=get_cache_config().compiled_cache_enabled ) - if driver_features is None or "storage_capabilities" not in driver_features: - driver_features = dict(driver_features) if driver_features else {} - driver_features["storage_capabilities"] = { - "arrow_export_enabled": False, - "arrow_import_enabled": True, - "parquet_export_enabled": False, - "parquet_import_enabled": False, - "partition_strategies": [], - } super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features) self._data_dictionary: PymssqlSyncDataDictionary | None = None @@ -181,7 +168,16 @@ def dispatch_execute(self, cursor: PymssqlRawCursor, statement: SQL) -> Executio return self.create_execution_result(cursor, rowcount_override=resolve_rowcount(cursor)) def dispatch_execute_many(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: - sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) + cached_statement, prepared_parameters = self._compiled_statement(statement, self.statement_config) + sql = cached_statement.compiled_sql + parsed_expression = cached_statement.expression + if parsed_expression is None and statement.raw_sql.lstrip().upper().startswith("INSERT"): + with contextlib.suppress(Exception): + parsed_expression = sqlglot.parse_one(statement.raw_sql, read="tsql") + if isinstance(parsed_expression, exp.Insert): + bulk_result = self._execute_bulk_insert_many(cursor, parsed_expression, prepared_parameters) + if bulk_result is not None: + return bulk_result prepared_parameters = normalize_execute_many_parameters(prepared_parameters) parameter_count = len(prepared_parameters) if isinstance(prepared_parameters, Sized) else None @@ -312,70 +308,51 @@ def data_dictionary(self) -> PymssqlSyncDataDictionary: self._data_dictionary = PymssqlSyncDataDictionary() return self._data_dictionary - def execute_many( - self, - statement: "SQL | Statement | QueryBuilder", - /, - parameters: "Sequence[StatementParameters]", - *filters: "StatementParameters | StatementFilter", - statement_config: StatementConfig | None = None, - **kwargs: Any, - ) -> SQLResult: - """Execute a statement across parameter sets with multi-row batching.""" - config = statement_config or self.statement_config - if isinstance(statement, str) and not filters and not kwargs and config is self.statement_config: - prepared_statement = SQL( - statement, - tuple(parameters) if isinstance(parameters, list) else parameters, - statement_config=config, - is_many=True, - ) - cached_statement, prepared_parameters = self._compiled_statement(prepared_statement, config) - parsed_expression = cached_statement.expression - if parsed_expression is None and statement.lstrip().upper().startswith("INSERT"): - with contextlib.suppress(Exception): - parsed_expression = sqlglot.parse_one(statement, read="tsql") - if isinstance(parsed_expression, exp.Insert) and not parsed_expression.args.get("returning"): - bulk_result = self._execute_bulk_insert_many(parsed_expression, prepared_parameters) - if bulk_result is not None: - return bulk_result - return super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - - def _execute_bulk_insert_many(self, expression: exp.Insert, prepared_parameters: Any) -> DMLResult | None: + def _execute_bulk_insert_many( + self, cursor: PymssqlRawCursor, expression: exp.Insert, prepared_parameters: Any + ) -> ExecutionResult | None: """Execute a batch INSERT via multi-row VALUES chunking up to 1,000 rows.""" if not isinstance(prepared_parameters, (list, tuple)) or not prepared_parameters: return None - if not isinstance(expression.this, exp.Schema): - return None - if not _is_plain_values_insert(expression, len(expression.this.expressions)): + first_row = prepared_parameters[0] + if not isinstance(first_row, (list, tuple)) or not first_row: return None - if not isinstance(prepared_parameters[0], (list, tuple)): + + target = expression.this + if isinstance(target, exp.Schema): + table_expr = target.this + column_names = [column.name for column in target.expressions] + elif isinstance(target, exp.Table): + table_expr = target + column_names = [] + else: return None - table_expr = expression.this.this if not isinstance(table_expr, exp.Table) or table_expr.alias: return None - column_names = [column.name for column in expression.this.expressions] + expected_columns = len(column_names) if column_names else len(first_row) + if expected_columns <= 0 or not is_plain_values_insert(expression, expected_columns): + return None + target_table = table_expr.sql(dialect="tsql") rows = prepared_parameters total_affected = 0 - chunk_size = 1000 + chunk_size = max(1, min(1000, 2000 // expected_columns)) - handler = self.handle_database_exceptions() - with handler, self.with_cursor(self.connection) as cursor: - for i in range(0, len(rows), chunk_size): - chunk = rows[i : i + chunk_size] - chunk_sql = build_multi_row_insert(target_table, column_names, len(chunk)) - flat_params: list[Any] = [] - for row in chunk: - flat_params.extend(row) - cursor.execute(chunk_sql, tuple(flat_params)) - total_affected += len(chunk) - self._check_pending_exception(handler) - return DMLResult("INSERT", total_affected) + for i in range(0, len(rows), chunk_size): + chunk = rows[i : i + chunk_size] + chunk_sql = build_multi_row_insert(target_table, column_names, len(chunk), num_columns=expected_columns) + flat_params: list[Any] = [] + for row in chunk: + flat_params.extend(row) + cursor.execute(chunk_sql, tuple(flat_params)) + rowcount = resolve_rowcount(cursor) + total_affected += rowcount if rowcount > 0 else len(chunk) - def bulk_copy( + return self.create_execution_result(cursor, rowcount_override=total_affected, is_many_result=True) + + def _bulk_copy( self, table_name: str, rows: Sequence[Sequence[Any]] | Iterable[Sequence[Any]], @@ -453,7 +430,7 @@ def load_from_arrow( for batch in arrow_table.to_batches(): pydict = batch.to_pydict() rows = list(zip(*pydict.values(), strict=False)) - self.bulk_copy( + self._bulk_copy( table, rows, column_ids=column_ids, @@ -489,19 +466,6 @@ def _connection_in_transaction(self) -> bool: return self._transaction_active -def _is_plain_values_insert(expression: exp.Insert, expected_columns: int) -> bool: - values = expression.expression - if not isinstance(values, exp.Values): - return False - rows = values.expressions - if len(rows) != 1: - return False - row = rows[0] - if not isinstance(row, exp.Tuple): - return False - return len(row.expressions) == expected_columns - - def _alter_default_schema_sql(user_name: str, schema: str) -> str: return f"ALTER USER {quote_tsql_identifier(user_name)} WITH DEFAULT_SCHEMA = {quote_tsql_identifier(schema)};" diff --git a/sqlspec/adapters/pymssql/events/store.py b/sqlspec/adapters/pymssql/events/store.py index 276af2275..9a40f6b09 100644 --- a/sqlspec/adapters/pymssql/events/store.py +++ b/sqlspec/adapters/pymssql/events/store.py @@ -39,7 +39,7 @@ def _wrap_create_statement(self, statement: str, object_type: str) -> str: table_name = match.group(1) return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NULL BEGIN {statement}; END" if object_type == "index": - match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+(\S+)", statement, re.IGNORECASE) + match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+([^\s(]+)", statement, re.IGNORECASE) if match: index_name = match.group(1).strip("[]") table_name = match.group(2) diff --git a/sqlspec/data_dictionary/dialects/mssql/__init__.py b/sqlspec/data_dictionary/dialects/mssql/__init__.py index 43b098e76..85f31759a 100644 --- a/sqlspec/data_dictionary/dialects/mssql/__init__.py +++ b/sqlspec/data_dictionary/dialects/mssql/__init__.py @@ -4,6 +4,7 @@ MSSQL_CONFIG, MSSQL_PRODUCT_VERSION_PATTERN, MSSQL_VERSION_PATTERN, + MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, @@ -17,6 +18,7 @@ mssql_supports_json_functions, mssql_supports_native_json, mssql_supports_string_agg, + mssql_supports_vector, mssql_system_metadata_denied, parse_mssql_engine_edition, parse_mssql_version_components, @@ -28,6 +30,7 @@ "MSSQL_CONFIG", "MSSQL_PRODUCT_VERSION_PATTERN", "MSSQL_VERSION_PATTERN", + "MssqlVersionInfo", "build_mssql_metadata_capability_profile", "build_mssql_system_metadata_capability", "build_mssql_system_metadata_result", @@ -41,6 +44,7 @@ "mssql_supports_json_functions", "mssql_supports_native_json", "mssql_supports_string_agg", + "mssql_supports_vector", "mssql_system_metadata_denied", "parse_mssql_engine_edition", "parse_mssql_version_components", diff --git a/sqlspec/data_dictionary/dialects/mssql/config.py b/sqlspec/data_dictionary/dialects/mssql/config.py index 622e3dea7..3b8ae1834 100644 --- a/sqlspec/data_dictionary/dialects/mssql/config.py +++ b/sqlspec/data_dictionary/dialects/mssql/config.py @@ -18,14 +18,16 @@ SystemMetadataRedactionPolicy, SystemMetadataRequest, SystemMetadataResult, + VersionInfo, register_dialect, system_metadata_gated_result, ) if TYPE_CHECKING: - from sqlspec.data_dictionary import TableMetadata, VersionInfo + from sqlspec.data_dictionary import TableMetadata __all__ = ( + "MssqlVersionInfo", "build_mssql_metadata_capability_profile", "build_mssql_system_metadata_capability", "build_mssql_system_metadata_result", @@ -39,6 +41,7 @@ "mssql_supports_json_functions", "mssql_supports_native_json", "mssql_supports_string_agg", + "mssql_supports_vector", "mssql_system_metadata_denied", "parse_mssql_engine_edition", "parse_mssql_version_components", @@ -53,6 +56,7 @@ MSSQL_MIN_STRING_AGG_VERSION: Final[int] = 14 MSSQL_MIN_GREATEST_LEAST_VERSION: Final[int] = 16 MSSQL_MIN_NATIVE_JSON_VERSION: Final[int] = 17 +MSSQL_VECTOR_MIN_MAJOR: Final[int] = 17 MSSQL_ENGINE_EDITION_AZURE_SET: Final[frozenset[int]] = frozenset({5, 8, 11}) MSSQL_DYNAMIC_FEATURES: Final[tuple[str, ...]] = ( @@ -61,6 +65,7 @@ "supports_string_agg", "supports_greatest_least", "supports_native_json", + "supports_vector", ) MSSQL_REPLACEMENT_DOMAINS: Final[tuple[str, ...]] = ( @@ -133,6 +138,7 @@ "text": "NVARCHAR(MAX)", "json": "NVARCHAR(MAX)", "jsonb": "NVARCHAR(MAX)", + "vector": "VARBINARY(MAX)", "timestamp": "DATETIME2(6)", "timestamptz": "DATETIMEOFFSET(6)", "bytea": "VARBINARY(MAX)", @@ -156,6 +162,48 @@ register_dialect(MSSQL_CONFIG) +class MssqlVersionInfo(VersionInfo): + """MSSQL database version info with build, revision, and Azure SQL detection.""" + + def __init__( + self, + major: int, + minor: int = 0, + build: int = 0, + revision: int = 0, + edition: str | None = None, + engine_edition: int | None = None, + ) -> None: + super().__init__(major, minor, 0) + self.build = build + self.revision = revision + self.edition = edition + self.engine_edition = engine_edition + self.is_azure_sql = is_mssql_azure_sql(engine_edition) + + def supports_native_json(self) -> bool: + """Return whether this server supports the native JSON type.""" + return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) + + def supports_vector(self) -> bool: + """Return whether this server supports native VECTOR data types and functions.""" + return mssql_supports_vector(self.major, is_azure_sql=self.is_azure_sql) + + @property + def version_tuple(self) -> tuple[int, int, int]: + """Get version tuple using the MSSQL build number as the third component.""" + return (self.major, self.minor, self.build) + + def __str__(self) -> str: + """String representation of version info.""" + version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" + if self.edition: + version_str += f" ({self.edition})" + if self.is_azure_sql: + version_str += " [Azure]" + return version_str + + def extract_mssql_version_value(row: object) -> "str | None": """Extract a SQL Server version string from a row-like object.""" if isinstance(row, dict): @@ -222,6 +270,11 @@ def mssql_supports_native_json(major: int, is_azure_sql: bool = False) -> bool: return is_azure_sql or major >= MSSQL_MIN_NATIVE_JSON_VERSION +def mssql_supports_vector(major: int, is_azure_sql: bool = False) -> bool: + """Return whether the SQL Server version supports native VECTOR data types and functions.""" + return is_azure_sql or major >= MSSQL_VECTOR_MIN_MAJOR + + def resolve_mssql_feature_flag( feature: str, *, @@ -243,6 +296,8 @@ def resolve_mssql_feature_flag( return mssql_supports_greatest_least(major) if feature == "supports_native_json": return mssql_supports_native_json(major, is_azure_sql=is_azure_sql) + if feature == "supports_vector": + return mssql_supports_vector(major, is_azure_sql=is_azure_sql) dialect_config = config or MSSQL_CONFIG flag = dialect_config.get_feature_flag(feature) diff --git a/tests/unit/adapters/test_mssql_python/test_arrow.py b/tests/unit/adapters/test_mssql_python/test_arrow.py index 7a0f24a26..8ab601875 100644 --- a/tests/unit/adapters/test_mssql_python/test_arrow.py +++ b/tests/unit/adapters/test_mssql_python/test_arrow.py @@ -3,6 +3,7 @@ from collections.abc import Iterable from typing import TYPE_CHECKING, cast +import pyarrow as pa import pytest from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver @@ -36,13 +37,9 @@ def fetchmany(self, size: int) -> list[tuple[int, str]]: return chunk def arrow(self, batch_size: int = 8192) -> object: - import pyarrow as pa - return pa.table({"x": [1, 2, 3]}) def arrow_reader(self, batch_size: int = 8192) -> object: - import pyarrow as pa - table = pa.table({"x": [1, 2, 3]}) return pa.RecordBatchReader.from_batches(table.schema, table.to_batches(max_chunksize=batch_size)) @@ -158,7 +155,7 @@ def test_bulk_copy_forwards_options_to_cursor_bulkcopy() -> None: connection = ArrowConnection() driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) - result = driver.bulk_copy( + result = driver._bulk_copy( "dbo.target", [(1, "a"), (2, "b")], batch_size=1000, timeout=30, table_lock=True, keep_nulls=True ) @@ -178,7 +175,7 @@ def test_bulk_copy_defaults_match_mssql_python_runtime() -> None: connection = ArrowConnection() driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) - result = driver.bulk_copy("dbo.target", [(1,)]) + result = driver._bulk_copy("dbo.target", [(1,)]) _, _, options = connection.cursor_obj.bulkcopy_calls[0] assert result["rows_copied"] == 1 @@ -193,6 +190,24 @@ def test_bulk_copy_raises_mapped_driver_exception() -> None: driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) with pytest.raises(UniqueViolationError): - driver.bulk_copy("dbo.target", [(1,)]) + driver._bulk_copy("dbo.target", [(1,)]) assert connection.cursor_obj.closed is True + + +def test_load_from_arrow_falls_back_to_bulk_copy_when_bulkcopy_arrow_absent() -> None: + """load_from_arrow should fall back to _bulk_copy when cursor lacks bulkcopy_arrow.""" + connection = ArrowConnection() + driver = MssqlPythonDriver( + cast("MssqlPythonConnection", connection), + driver_features={"storage_capabilities": {"arrow_import_enabled": True}}, + ) + table = pa.table({"id": [1, 2], "name": ["Ada", "Grace"]}) + + job = driver.load_from_arrow("dbo.target", table, column_mappings=[]) + + assert job.telemetry["rows_processed"] == 2 + target_table, rows, options = connection.cursor_obj.bulkcopy_calls[0] + assert target_table == "dbo.target" + assert rows == [(1, "Ada"), (2, "Grace")] + assert options["column_mappings"] == [] diff --git a/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py b/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py index 462f74757..511cef051 100644 --- a/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py +++ b/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py @@ -36,7 +36,7 @@ def test_bulk_copy_defaults_match_upstream(driver_cls: Any) -> None: from mssql_python.cursor import Cursor upstream = inspect.signature(Cursor.bulkcopy).parameters - wrapper = inspect.signature(driver_cls.bulk_copy).parameters + wrapper = inspect.signature(driver_cls._bulk_copy).parameters for name in ( "batch_size", "timeout", diff --git a/tests/unit/adapters/test_mssql_python/test_events_store.py b/tests/unit/adapters/test_mssql_python/test_events_store.py index bfa11578b..ffca80002 100644 --- a/tests/unit/adapters/test_mssql_python/test_events_store.py +++ b/tests/unit/adapters/test_mssql_python/test_events_store.py @@ -22,6 +22,11 @@ def test_event_queue_store_uses_tsql_column_types_and_idempotency() -> None: assert "payload_json NVARCHAR(MAX) NOT NULL" in ddl assert "available_at DATETIME2(6) NOT NULL DEFAULT SYSUTCDATETIME()" in ddl assert "IF NOT EXISTS (SELECT 1 FROM sys.indexes" in ddl + assert "OBJECT_ID(N'[dbo].[sqlspec_event_queue]')" in ddl + wrapped_no_space = store._wrap_create_statement( + "CREATE INDEX idx_events_channel_status ON app_events(channel, status, available_at)", "index" + ) + assert "OBJECT_ID(N'[dbo].[app_events]')" in wrapped_no_space def test_event_queue_store_drop_uses_object_id_guard() -> None: diff --git a/tests/unit/adapters/test_pymssql/test_config.py b/tests/unit/adapters/test_pymssql/test_config.py index cffd9fb35..53d5ee1c1 100644 --- a/tests/unit/adapters/test_pymssql/test_config.py +++ b/tests/unit/adapters/test_pymssql/test_config.py @@ -4,7 +4,9 @@ import pytest +from sqlspec.adapters.pymssql import PymssqlConnection as PublicPymssqlConnection from sqlspec.adapters.pymssql import build_connection_config +from sqlspec.adapters.pymssql._typing import PymssqlConnection, PymssqlRawCursor, pymssql_module from sqlspec.adapters.pymssql.config import PymssqlConfig, PymssqlConnectionParams from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.adapters.pymssql.pool import PymssqlConnectionPool @@ -121,8 +123,6 @@ def test_signature_namespace_exposes_public_adapter_types() -> None: def test_pymssql_runtime_aliases_resolve_to_installed_classes() -> None: """pymssql public runtime aliases should expose installed pymssql classes.""" pymssql = pytest.importorskip("pymssql") - from sqlspec.adapters.pymssql import PymssqlConnection as PublicPymssqlConnection - from sqlspec.adapters.pymssql._typing import PymssqlConnection, PymssqlRawCursor, pymssql_module namespace = PymssqlConfig().get_signature_namespace() diff --git a/tests/unit/adapters/test_pymssql/test_core.py b/tests/unit/adapters/test_pymssql/test_core.py index 67d24143d..d7aaedd94 100644 --- a/tests/unit/adapters/test_pymssql/test_core.py +++ b/tests/unit/adapters/test_pymssql/test_core.py @@ -4,6 +4,19 @@ import pytest +from sqlspec.adapters.pymssql.core import ( + build_insert_statement, + build_multi_row_insert, + collect_rows, + create_mapped_exception, + default_statement_config, + driver_profile, + extract_error_number, + format_identifier, + normalize_execute_many_parameters, + normalize_execute_parameters, + quote_tsql_identifier, +) from sqlspec.core import SQL, ParameterStyle from sqlspec.exceptions import ( CheckViolationError, @@ -16,8 +29,6 @@ def test_profile_uses_tsql_and_pyformat_execution() -> None: """The pymssql profile should compile T-SQL to pyformat placeholders.""" - from sqlspec.adapters.pymssql.core import default_statement_config, driver_profile - parameter_config = default_statement_config.parameter_config assert default_statement_config.dialect == "tsql" @@ -35,8 +46,6 @@ def test_profile_uses_tsql_and_pyformat_execution() -> None: def test_statement_config_compiles_qmark_input_to_percent_s() -> None: """Qmark input should execute as positional pyformat for pymssql.""" - from sqlspec.adapters.pymssql.core import default_statement_config - statement = SQL("SELECT * FROM dbo.users WHERE id = ?", 3, statement_config=default_statement_config) compiled_sql, parameters = statement.compile() @@ -47,8 +56,6 @@ def test_statement_config_compiles_qmark_input_to_percent_s() -> None: def test_statement_config_compiles_named_pyformat_input_to_positional() -> None: """Named pyformat input should compile to pymssql's supported positional style.""" - from sqlspec.adapters.pymssql.core import default_statement_config - statement = SQL( "SELECT * FROM dbo.users WHERE id = %(user_id)s", {"user_id": 3}, statement_config=default_statement_config ) @@ -61,8 +68,6 @@ def test_statement_config_compiles_named_pyformat_input_to_positional() -> None: def test_format_identifier_and_insert_statement_use_tsql_identifiers() -> None: """Generated DML helpers should quote T-SQL identifiers and use %s placeholders.""" - from sqlspec.adapters.pymssql.core import build_insert_statement, format_identifier - assert format_identifier("dbo.users") == "[dbo].[users]" assert format_identifier("[sales].[order]]items]") == "[sales].[order]]items]" assert build_insert_statement("dbo.users", ["id", "display_name"]) == ( @@ -79,8 +84,6 @@ def test_format_identifier_and_insert_statement_use_tsql_identifiers() -> None: ) def test_create_mapped_exception_maps_tsql_error_numbers(message: str, expected_type: type[Exception]) -> None: """SQL Server error numbers should map to SQLSpec exceptions.""" - from sqlspec.adapters.pymssql.core import create_mapped_exception - exc = create_mapped_exception(Exception(message)) assert isinstance(exc, expected_type) @@ -116,8 +119,6 @@ def test_create_mapped_exception_disambiguates_547_check_vs_foreign_key( message: str, expected_type: type[Exception], expected_detail: str ) -> None: """SQL Server 547 distinguishes CHECK from foreign-key constraint violations.""" - from sqlspec.adapters.pymssql.core import create_mapped_exception - mapped = create_mapped_exception(Exception(message)) assert isinstance(mapped, expected_type) @@ -137,15 +138,11 @@ def test_create_mapped_exception_classifies_native_constraint_shapes( error: Exception, expected_type: type[Exception] ) -> None: """Native pymssql argument and message shapes map to specific constraint exceptions.""" - from sqlspec.adapters.pymssql.core import create_mapped_exception - assert isinstance(create_mapped_exception(error), expected_type) def test_normalize_execute_many_parameters_passes_through() -> None: """normalize_execute_many_parameters returns the batch payload unchanged.""" - from sqlspec.adapters.pymssql.core import normalize_execute_many_parameters - assert normalize_execute_many_parameters([]) == [] rows: list[tuple[Any, ...]] = [(1,), (2,)] @@ -154,8 +151,6 @@ def test_normalize_execute_many_parameters_passes_through() -> None: def test_quote_tsql_identifier() -> None: """quote_tsql_identifier wraps identifiers in brackets and escapes closing brackets.""" - from sqlspec.adapters.pymssql.core import quote_tsql_identifier - assert quote_tsql_identifier("users") == "[users]" assert quote_tsql_identifier("[users]") == "[users]" assert quote_tsql_identifier("dbo.users") == "[dbo.users]" @@ -164,12 +159,16 @@ def test_quote_tsql_identifier() -> None: def test_extract_error_number() -> None: """extract_error_number detects error number from attribute, tuple, or regex.""" - from sqlspec.adapters.pymssql.core import extract_error_number class AttributeException(Exception): number = 2627 + class BoolAttributeException(Exception): + number = True + assert extract_error_number(AttributeException("duplicate key")) == 2627 + assert extract_error_number(BoolAttributeException("Msg 2627, Level 14")) == 2627 + assert extract_error_number(Exception(True, "Msg 1205, Level 13")) == 1205 assert extract_error_number(Exception(1205, "Deadlock found")) == 1205 assert extract_error_number(Exception("Violation of UNIQUE KEY constraint (2627)")) == 2627 assert extract_error_number(Exception("Plain error")) is None @@ -177,16 +176,12 @@ class AttributeException(Exception): def test_build_multi_row_insert() -> None: """build_multi_row_insert generates a multi-row VALUES INSERT statement.""" - from sqlspec.adapters.pymssql.core import build_multi_row_insert - sql = build_multi_row_insert("dbo.users", ["id", "name"], 3) assert sql == "INSERT INTO [dbo].[users] ([id], [name]) VALUES (%s, %s), (%s, %s), (%s, %s)" def test_collect_rows_preserves_list_identity() -> None: """collect_rows avoids copying when the input rows are already a list.""" - from sqlspec.adapters.pymssql.core import collect_rows - input_rows = [(1, "Alice"), (2, "Bob")] description = [("id",), ("name",)] rows, column_names, row_format = collect_rows(input_rows, description) @@ -198,7 +193,5 @@ def test_collect_rows_preserves_list_identity() -> None: def test_normalize_execute_parameters_preserves_tuples() -> None: """normalize_execute_parameters passes tuples through directly.""" - from sqlspec.adapters.pymssql.core import normalize_execute_parameters - params = (1, "Alice") assert normalize_execute_parameters(params) is params diff --git a/tests/unit/adapters/test_pymssql/test_data_dictionary.py b/tests/unit/adapters/test_pymssql/test_data_dictionary.py index 2102a4d5b..8b9e08956 100644 --- a/tests/unit/adapters/test_pymssql/test_data_dictionary.py +++ b/tests/unit/adapters/test_pymssql/test_data_dictionary.py @@ -2,7 +2,9 @@ from typing import Any, cast +from sqlspec.adapters.mssql_python.data_dictionary import MssqlVersionInfo as MssqlPythonVersionInfo from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo, PymssqlSyncDataDictionary +from sqlspec.data_dictionary.dialects.mssql import MssqlVersionInfo as DialectMssqlVersionInfo class FakeSyncDriver: @@ -144,13 +146,23 @@ def test_mssql_version_info_supports_vector() -> None: def test_data_dictionary_vector_feature_flag_and_optimal_type() -> None: """Sync data dictionary resolves supports_vector and optimal type for vector.""" + assert MssqlVersionInfo is MssqlPythonVersionInfo + assert MssqlVersionInfo is DialectMssqlVersionInfo class VectorDriver: def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: return {"product_version": "17.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} + class NonVectorDriver: + def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: + return {"product_version": "16.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} + data_dictionary = PymssqlSyncDataDictionary() driver = VectorDriver() + old_driver = NonVectorDriver() + assert "supports_vector" in data_dictionary.list_available_features() assert data_dictionary.get_feature_flag(cast(Any, driver), "supports_vector") is True assert data_dictionary.get_optimal_type(cast(Any, driver), "vector") == "VECTOR" + assert data_dictionary.get_feature_flag(cast(Any, old_driver), "supports_vector") is False + assert data_dictionary.get_optimal_type(cast(Any, old_driver), "vector") == "VARBINARY(MAX)" diff --git a/tests/unit/adapters/test_pymssql/test_driver.py b/tests/unit/adapters/test_pymssql/test_driver.py index ece0e7eac..4cf1f848d 100644 --- a/tests/unit/adapters/test_pymssql/test_driver.py +++ b/tests/unit/adapters/test_pymssql/test_driver.py @@ -2,11 +2,14 @@ from typing import Any, cast +import pyarrow as pa import pytest from pymssql import IntegrityError as PymssqlIntegrityError from sqlspec import StatementStack from sqlspec.adapters.pymssql._typing import PymssqlConnection, PymssqlRawCursor +from sqlspec.adapters.pymssql.core import default_statement_config +from sqlspec.adapters.pymssql.driver import PymssqlDriver, PymssqlExceptionHandler from sqlspec.core import SQL from sqlspec.exceptions import SQLSpecError, StackExecutionError, TransactionError, UniqueViolationError from tests.unit.adapters.test_pymssql._fakes import FakeConnection, FakeCursor @@ -29,8 +32,6 @@ def test_execute_maps_pymssql_row_formats( rows: list[tuple[int, str] | dict[str, int | str]], expected: list[dict[str, int | str]] ) -> None: - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor(rows=rows, description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor))) @@ -44,8 +45,6 @@ def test_execute_maps_pymssql_row_formats( @pytest.mark.parametrize("bad_name", UNSAFE_SAVEPOINT_NAMES) def test_pymssql_savepoint_overrides_reject_unsafe_names(bad_name: str) -> None: """The T-SQL savepoint overrides must reject unsafe identifiers before interpolation.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection())) with pytest.raises(TransactionError): @@ -58,8 +57,6 @@ def test_pymssql_savepoint_overrides_reject_unsafe_names(bad_name: str) -> None: def test_pymssql_savepoint_overrides_accept_valid_name() -> None: """A safe savepoint name should pass validation and reach the underlying execute path.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor() connection = FakeConnection(cursor) driver = PymssqlDriver(cast("PymssqlConnection", connection)) @@ -74,9 +71,6 @@ def test_pymssql_savepoint_overrides_accept_valid_name() -> None: def test_dispatch_execute_select_compiles_to_pyformat_and_collects_rows() -> None: """SELECT dispatch should execute pyformat SQL and return fetched rows.""" - from sqlspec.adapters.pymssql.core import default_statement_config - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor(rows=[(1, "Ada")], description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor)), statement_config=default_statement_config) statement = SQL("SELECT id, name FROM dbo.users WHERE id = ?", 1, statement_config=default_statement_config) @@ -90,19 +84,19 @@ def test_dispatch_execute_select_compiles_to_pyformat_and_collects_rows() -> Non def test_dispatch_execute_many_uses_executemany_and_rowcount() -> None: - """execute_many dispatch should forward batch parameters to pymssql.""" - from sqlspec.adapters.pymssql.core import default_statement_config - from sqlspec.adapters.pymssql.driver import PymssqlDriver - + """execute_many dispatch should forward non-plain-INSERT batch parameters to pymssql executemany.""" cursor = FakeCursor(rowcount=2) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor)), statement_config=default_statement_config) statement = SQL( - "INSERT INTO dbo.users (id) VALUES (?)", [(1,), (2,)], statement_config=default_statement_config, is_many=True + "UPDATE dbo.users SET name = ? WHERE id = ?", + [("Ada", 1), ("Grace", 2)], + statement_config=default_statement_config, + is_many=True, ) result = driver.dispatch_execute_many(cast("PymssqlRawCursor", cursor), statement) - assert cursor.many_calls == [("INSERT INTO dbo.users (id) VALUES (%s)", [(1,), (2,)])] + assert cursor.many_calls == [("UPDATE dbo.users SET name = %s WHERE id = %s", [("Ada", 1), ("Grace", 2)])] assert result.rowcount_override == 2 assert result.is_many_result is True @@ -110,8 +104,6 @@ def test_dispatch_execute_many_uses_executemany_and_rowcount() -> None: @pytest.mark.parametrize("finish", ["commit", "rollback"]) def test_autocommit_transaction_is_ended_with_tsql(finish: str) -> None: """pymssql ignores commit() and rollback() under autocommit, so the driver ends its own transaction.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor() connection = FakeConnection(cursor) driver = PymssqlDriver(cast("PymssqlConnection", connection)) @@ -127,8 +119,6 @@ def test_autocommit_transaction_is_ended_with_tsql(finish: str) -> None: @pytest.mark.parametrize("finish", ["commit", "rollback"]) def test_non_autocommit_transaction_uses_connection_boundaries(finish: str) -> None: """Without autocommit, pymssql's connection commit() and rollback() end the open transaction.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor() connection = FakeConnection(cursor) connection.autocommit(False) @@ -143,8 +133,6 @@ def test_non_autocommit_transaction_uses_connection_boundaries(finish: str) -> N def test_begin_reuses_the_open_transaction_without_autocommit() -> None: """A connection with autocommit disabled already holds a transaction, so begin issues no SQL.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor() connection = FakeConnection(cursor) connection.autocommit(False) @@ -160,8 +148,6 @@ def test_begin_reuses_the_open_transaction_without_autocommit() -> None: def test_exception_handler_maps_pymssql_errors() -> None: """pymssql exception handlers should surface mapped SQLSpec exceptions.""" - from sqlspec.adapters.pymssql.driver import PymssqlExceptionHandler - handler = PymssqlExceptionHandler() handled = handler._handle_exception( @@ -174,7 +160,6 @@ def test_exception_handler_maps_pymssql_errors() -> None: def test_commit_wraps_driver_errors() -> None: """Commit failures should be wrapped in SQLSpecError.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver class FailingConnection(FakeConnection): def commit(self) -> None: @@ -188,8 +173,6 @@ def commit(self) -> None: def test_collect_rows_returns_column_names() -> None: """The direct row collection hook should match SyncDriverAdapterBase expectations.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor(description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor))) @@ -202,8 +185,6 @@ def test_collect_rows_returns_column_names() -> None: def test_select_stream_uses_fetchmany_chunks() -> None: """The pymssql driver should stream rows with cursor.fetchmany().""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor(rows=[(1, "Ada"), (2, "Grace"), (3, "Linus")], description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor))) @@ -218,8 +199,6 @@ def test_select_stream_uses_fetchmany_chunks() -> None: @pytest.mark.parametrize("finish", ["commit", "rollback"]) def test_connection_in_transaction_tracks_successful_boundaries(finish: str) -> None: - from sqlspec.adapters.pymssql.driver import PymssqlDriver - driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection())) assert driver._connection_in_transaction() is False driver.begin() @@ -230,8 +209,6 @@ def test_connection_in_transaction_tracks_successful_boundaries(finish: str) -> @pytest.mark.parametrize("operation", ["begin", "commit", "rollback"]) def test_failed_transaction_boundary_preserves_state(operation: str, monkeypatch: pytest.MonkeyPatch) -> None: - from sqlspec.adapters.pymssql.driver import PymssqlDriver - connection = FakeConnection() driver = PymssqlDriver(cast("PymssqlConnection", connection)) if operation != "begin": @@ -250,8 +227,6 @@ def fail(*_args: object) -> None: @pytest.mark.parametrize("fails", [False, True]) def test_execute_stack_preserves_caller_transaction(fails: bool, monkeypatch: pytest.MonkeyPatch) -> None: - from sqlspec.adapters.pymssql.driver import PymssqlDriver - connection = FakeConnection(FakeCursor(rowcount=1)) driver = PymssqlDriver(cast("PymssqlConnection", connection)) driver.begin() @@ -281,13 +256,11 @@ def fail(sql: str, parameters: object = None) -> None: def test_driver_bulk_copy_forwards_options() -> None: - """bulk_copy forwards batch options to underlying connection.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - + """_bulk_copy forwards batch options to underlying connection.""" connection = FakeConnection() driver = PymssqlDriver(cast("PymssqlConnection", connection)) - result = driver.bulk_copy( + result = driver._bulk_copy( "dbo.users", [(1, "Ada"), (2, "Grace")], column_ids=[1, 2], @@ -310,13 +283,11 @@ def test_driver_bulk_copy_forwards_options() -> None: def test_load_from_arrow_bulk_copies_batches() -> None: - """load_from_arrow processes Arrow table in batches via bulk_copy.""" - import pyarrow as pa - - from sqlspec.adapters.pymssql.driver import PymssqlDriver - + """load_from_arrow processes Arrow table in batches via _bulk_copy.""" connection = FakeConnection() - driver = PymssqlDriver(cast("PymssqlConnection", connection)) + driver = PymssqlDriver( + cast("PymssqlConnection", connection), driver_features={"storage_capabilities": {"arrow_import_enabled": True}} + ) table = pa.table({"id": [1, 2], "name": ["Ada", "Grace"]}) job = driver.load_from_arrow("dbo.users", table, batch_size=500) @@ -327,13 +298,11 @@ def test_load_from_arrow_bulk_copies_batches() -> None: def test_load_from_arrow_overwrite_truncates_first() -> None: """load_from_arrow with overwrite=True executes TRUNCATE TABLE.""" - import pyarrow as pa - - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor() connection = FakeConnection(cursor) - driver = PymssqlDriver(cast("PymssqlConnection", connection)) + driver = PymssqlDriver( + cast("PymssqlConnection", connection), driver_features={"storage_capabilities": {"arrow_import_enabled": True}} + ) table = pa.table({"id": [1], "name": ["Ada"]}) driver.load_from_arrow("dbo.users", table, overwrite=True) @@ -344,16 +313,15 @@ def test_load_from_arrow_overwrite_truncates_first() -> None: def test_load_from_arrow_overwrite_falls_back_on_fk_error() -> None: """load_from_arrow falls back to DELETE FROM when error 4712 is encountered.""" - import pyarrow as pa - - from sqlspec.adapters.pymssql.driver import PymssqlDriver class FkError(Exception): number = 4712 cursor = FakeCursor() connection = FakeConnection(cursor) - driver = PymssqlDriver(cast("PymssqlConnection", connection)) + driver = PymssqlDriver( + cast("PymssqlConnection", connection), driver_features={"storage_capabilities": {"arrow_import_enabled": True}} + ) def execute_with_fk(sql: str, *args: Any) -> None: cursor.calls.append((sql, args)) @@ -371,9 +339,7 @@ def execute_with_fk(sql: str, *args: Any) -> None: def test_execute_many_plain_values_chunks_into_multi_row_insert() -> None: - """execute_many with plain VALUES uses multi-row INSERT.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - + """execute_many with plain VALUES uses multi-row INSERT for both str and SQL objects.""" cursor = FakeCursor(rowcount=3) connection = FakeConnection(cursor) driver = PymssqlDriver(cast("PymssqlConnection", connection)) @@ -386,11 +352,15 @@ def test_execute_many_plain_values_chunks_into_multi_row_insert() -> None: assert len(executed_sqls) == 1 assert "VALUES (%s, %s), (%s, %s), (%s, %s)" in executed_sqls[0] + cursor.calls.clear() + sql_obj_result = driver.execute_many(SQL("INSERT INTO dbo.users VALUES (?, ?)"), params) + assert sql_obj_result.rows_affected == 3 + assert len(cursor.calls) == 1 + assert "INSERT INTO [dbo].[users] VALUES (%s, %s), (%s, %s), (%s, %s)" in cursor.calls[0][0] + def test_execute_many_non_plain_values_uses_standard_executemany() -> None: """execute_many with non-plain SQL uses cursor.executemany.""" - from sqlspec.adapters.pymssql.driver import PymssqlDriver - cursor = FakeCursor(rowcount=2) connection = FakeConnection(cursor) driver = PymssqlDriver(cast("PymssqlConnection", connection)) diff --git a/tests/unit/adapters/test_pymssql/test_extensions.py b/tests/unit/adapters/test_pymssql/test_extensions.py index f70b0c2a5..36837fa72 100644 --- a/tests/unit/adapters/test_pymssql/test_extensions.py +++ b/tests/unit/adapters/test_pymssql/test_extensions.py @@ -17,6 +17,10 @@ def test_event_store_uses_tsql_column_types_and_idempotent_wrappers() -> None: assert store._timestamp_default() == "SYSUTCDATETIME()" assert "OBJECT_ID" in store._wrap_create_statement("CREATE TABLE event_queue (id INT)", "table") assert "sys.indexes" in store._wrap_create_statement("CREATE INDEX idx_events ON event_queue (channel)", "index") + wrapped_no_space = store._wrap_create_statement( + "CREATE INDEX idx_events_channel_status ON app_events(channel, status, available_at)", "index" + ) + assert "OBJECT_ID(N'[dbo].[app_events]')" in wrapped_no_space def test_litestar_store_ddl_is_tsql_idempotent() -> None: From cdc4bbe1a00c99f8ee0a7f618f06cfbcbc4618f1 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 20:53:12 +0000 Subject: [PATCH 07/11] fix(mssql): preserve native execution and narrow adapter cleanup --- sqlspec/adapters/mssql_python/__init__.py | 7 +- sqlspec/adapters/mssql_python/_typing.py | 19 +- sqlspec/adapters/mssql_python/adk/store.py | 152 +++++------ sqlspec/adapters/mssql_python/config.py | 96 +++---- sqlspec/adapters/mssql_python/core.py | 48 ++-- .../adapters/mssql_python/data_dictionary.py | 46 +++- sqlspec/adapters/mssql_python/driver.py | 236 ++++++------------ .../adapters/mssql_python/litestar/store.py | 24 +- sqlspec/adapters/mssql_python/pool.py | 34 ++- .../adapters/mssql_python/type_converter.py | 21 +- sqlspec/adapters/pymssql/_typing.py | 22 +- sqlspec/adapters/pymssql/adk/store.py | 150 +++++------ sqlspec/adapters/pymssql/config.py | 70 +++--- sqlspec/adapters/pymssql/core.py | 107 +++----- sqlspec/adapters/pymssql/data_dictionary.py | 46 +++- sqlspec/adapters/pymssql/driver.py | 221 +++------------- sqlspec/adapters/pymssql/events/store.py | 7 +- sqlspec/adapters/pymssql/litestar/store.py | 24 +- sqlspec/adapters/pymssql/pool.py | 4 +- .../dialects/mssql/__init__.py | 4 - .../data_dictionary/dialects/mssql/config.py | 57 +---- .../adapters/test_mssql_python/test_arrow.py | 29 +-- .../test_bulk_copy_result.py | 2 +- .../adapters/test_mssql_python/test_config.py | 45 ++-- .../adapters/test_mssql_python/test_core.py | 8 +- .../test_mssql_python/test_data_dictionary.py | 25 -- .../test_mssql_python/test_load_from_arrow.py | 45 +--- .../test_mssql_python/test_type_converter.py | 6 - tests/unit/adapters/test_pymssql/_fakes.py | 21 -- .../unit/adapters/test_pymssql/test_config.py | 8 +- tests/unit/adapters/test_pymssql/test_core.py | 16 -- .../test_pymssql/test_data_dictionary.py | 37 --- .../unit/adapters/test_pymssql/test_driver.py | 162 +++--------- 33 files changed, 643 insertions(+), 1156 deletions(-) diff --git a/sqlspec/adapters/mssql_python/__init__.py b/sqlspec/adapters/mssql_python/__init__.py index a098ed261..4964a7b8a 100644 --- a/sqlspec/adapters/mssql_python/__init__.py +++ b/sqlspec/adapters/mssql_python/__init__.py @@ -9,12 +9,17 @@ ) from sqlspec.adapters.mssql_python.core import default_statement_config, driver_profile from sqlspec.adapters.mssql_python.data_dictionary import MssqlPythonSyncDataDictionary, MssqlVersionInfo -from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver, MssqlPythonExceptionHandler +from sqlspec.adapters.mssql_python.driver import ( + MssqlPythonBulkCopyResult, + MssqlPythonDriver, + MssqlPythonExceptionHandler, +) from sqlspec.adapters.mssql_python.migrations import MssqlPythonSyncMigrationTracker from sqlspec.adapters.mssql_python.pool import MssqlPythonConnectionPool from sqlspec.adapters.mssql_python.type_converter import MssqlPythonTypeConverter, mssql_type_to_arrow __all__ = ( + "MssqlPythonBulkCopyResult", "MssqlPythonConfig", "MssqlPythonConnection", "MssqlPythonConnectionParams", diff --git a/sqlspec/adapters/mssql_python/_typing.py b/sqlspec/adapters/mssql_python/_typing.py index bf9b1dfd5..27f48317c 100644 --- a/sqlspec/adapters/mssql_python/_typing.py +++ b/sqlspec/adapters/mssql_python/_typing.py @@ -3,27 +3,36 @@ import contextlib from typing import TYPE_CHECKING, Any -import mssql_python as mssql_python_module +import mssql_python as _mssql_python # pyright: ignore[reportMissingImports] from mssql_python import Error as MssqlPythonError -from mssql_python.connection import Connection as MssqlPythonConnection -from mssql_python.connection import TokenProvider -from mssql_python.cursor import Cursor as MssqlPythonRawCursor +from mssql_python.connection import Connection, TokenProvider # pyright: ignore +from mssql_python.cursor import Cursor # pyright: ignore + +MSSQL_PYTHON_MODULE: Any = _mssql_python if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType + from typing import TypeAlias from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver from sqlspec.core import StatementConfig + MssqlPythonConnection: TypeAlias = Connection + MssqlPythonRawCursor: TypeAlias = Cursor + +if not TYPE_CHECKING: + MssqlPythonConnection = Connection + MssqlPythonRawCursor = Cursor + __all__ = ( + "MSSQL_PYTHON_MODULE", "MssqlPythonConnection", "MssqlPythonCursor", "MssqlPythonError", "MssqlPythonRawCursor", "MssqlPythonSessionContext", "TokenProvider", - "mssql_python_module", ) diff --git a/sqlspec/adapters/mssql_python/adk/store.py b/sqlspec/adapters/mssql_python/adk/store.py index 629611567..1626a210a 100644 --- a/sqlspec/adapters/mssql_python/adk/store.py +++ b/sqlspec/adapters/mssql_python/adk/store.py @@ -1,26 +1,25 @@ """mssql-python ADK stores for Google Agent Development Kit session storage.""" -from collections.abc import Sequence -from datetime import datetime, timedelta -from typing import Any, ClassVar, Final, Literal, cast +from datetime import datetime +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast from typing_extensions import NotRequired from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor, MssqlPythonError -from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.adapters.mssql_python.core import extract_error_number from sqlspec.config import ADKConfig -from sqlspec.extensions.adk import ( - BaseSyncADKMemoryStore, - BaseSyncADKStore, - SessionOrderBy, - StoredEvent, - StoredMemory, - StoredSession, - normalize_session_list_options, -) +from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options +from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore from sqlspec.utils.serializers import from_json, to_json +if TYPE_CHECKING: + from collections.abc import Sequence + from datetime import timedelta + + from sqlspec.adapters.mssql_python.config import MssqlPythonConfig + from sqlspec.extensions.adk import SessionOrderBy + from sqlspec.extensions.adk.memory._types import StoredMemory + __all__ = ("MssqlPythonADKConfig", "MssqlPythonADKMemoryStore", "MssqlPythonADKStore") MSSQL_TABLE_NOT_FOUND_ERROR: Final[int] = 208 @@ -44,7 +43,7 @@ class MssqlPythonADKStore(BaseSyncADKStore["MssqlPythonConfig"]): connector_name: ClassVar[str] = "mssql_python" __slots__ = ("_json_column_type",) - def __init__(self, config: MssqlPythonConfig) -> None: + def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) adk_config = _adk_config(config) native_json = adk_config.get("native_json") @@ -71,7 +70,7 @@ def create_tables(self) -> None: driver.commit() def create_session( - self, session_id: str, app_name: str, user_id: str, state: dict[str, Any], owner_id: Any | None = None + self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None ) -> StoredSession: """Create a new ADK session.""" owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" @@ -95,8 +94,8 @@ def create_session( return _session_record_from_row(row) def get_session( - self, app_name: str, user_id: str, session_id: str, *, renew_for: int | timedelta | None = None - ) -> StoredSession | None: + self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None + ) -> "StoredSession | None": """Return a scoped session or ``None`` if absent.""" try: if renew_for is not None and self._calculate_expires_at(renew_for) is not None: @@ -123,7 +122,7 @@ def get_session( raise return _session_record_from_row(row) if row is not None else None - def update_session_state(self, app_name: str, user_id: str, session_id: str, state: dict[str, Any]) -> None: + def update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None: """Replace a session's durable state.""" self._execute( f""" @@ -138,13 +137,13 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta def list_sessions( self, app_name: str, - user_id: str | None = None, + user_id: "str | None" = None, *, - order_by: SessionOrderBy = "update_time", + order_by: "SessionOrderBy" = "update_time", descending: bool = True, - limit: int | None = None, - offset: int | None = None, - ) -> list[StoredSession]: + limit: "int | None" = None, + offset: "int | None" = None, + ) -> "list[StoredSession]": """List ADK sessions for an application, optionally scoped to a user.""" column, direction, page_limit, page_offset = normalize_session_list_options(order_by, descending, limit, offset) if page_limit == 0: @@ -179,10 +178,10 @@ def append_event_and_update_state( app_name: str, user_id: str, session_id: str, - state: dict[str, Any], + state: "dict[str, Any]", *, - app_state: dict[str, Any] | None = None, - user_state: dict[str, Any] | None = None, + app_state: "dict[str, Any] | None" = None, + user_state: "dict[str, Any] | None" = None, ) -> StoredSession: """Atomically append an event and update durable session/scoped state.""" update_sql = f""" @@ -213,9 +212,9 @@ def get_events( app_name: str, user_id: str, session_id: str, - after_timestamp: datetime | None = None, - limit: int | None = None, - ) -> list[StoredEvent]: + after_timestamp: "datetime | None" = None, + limit: "int | None" = None, + ) -> "list[StoredEvent]": """Return events for a scoped session ordered by event timestamp.""" if limit == 0: return [] @@ -228,7 +227,7 @@ def get_events( raise return [_event_record_from_row(row) for row in rows] - def delete_expired_events(self, before: datetime, app_name: str | None = None) -> int: + def delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: """Delete events older than ``before``.""" sql = f"DELETE FROM {_table_ref(self._events_table)} WHERE timestamp < ?" params: list[Any] = [before] @@ -242,7 +241,7 @@ def delete_expired_events(self, before: datetime, app_name: str | None = None) - return 0 raise - def delete_idle_sessions(self, updated_before: datetime, app_name: str | None = None) -> int: + def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: """Delete sessions whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._session_table)} WHERE update_time < ?" params: list[Any] = [updated_before] @@ -256,7 +255,7 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: str | None = return 0 raise - def delete_idle_user_states(self, updated_before: datetime, app_name: str | None = None) -> int: + def delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: """Delete user state rows whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._user_state_table)} WHERE update_time < ?" params: list[Any] = [updated_before] @@ -270,7 +269,7 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: str | None return 0 raise - def get_app_state(self, app_name: str) -> dict[str, Any] | None: + def get_app_state(self, app_name: str) -> "dict[str, Any] | None": """Return app-scoped state.""" try: row = self._execute_fetchone( @@ -282,7 +281,7 @@ def get_app_state(self, app_name: str) -> dict[str, Any] | None: raise return _json_dict(row[0]) if row is not None else None - def get_user_state(self, app_name: str, user_id: str) -> dict[str, Any] | None: + def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None": """Return user-scoped state.""" try: row = self._execute_fetchone( @@ -299,15 +298,15 @@ def get_user_state(self, app_name: str, user_id: str) -> dict[str, Any] | None: raise return _json_dict(row[0]) if row is not None else None - def upsert_app_state(self, app_name: str, state: dict[str, Any]) -> None: + def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """Insert or replace app-scoped state.""" self._execute(self._upsert_app_state_sql(), (app_name, to_json(state)), commit=True) - def upsert_user_state(self, app_name: str, user_id: str, state: dict[str, Any]) -> None: + def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: """Insert or replace user-scoped state.""" self._execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(state)), commit=True) - def get_metadata(self, key: str) -> str | None: + def get_metadata(self, key: str) -> "str | None": """Return an ADK metadata value.""" try: row = self._execute_fetchone( @@ -323,7 +322,7 @@ def set_metadata(self, key: str, value: str) -> None: """Set an ADK metadata value.""" self._execute(_upsert_metadata_sql(self._metadata_table), (key, value), commit=True) - def _index_specs(self) -> list[tuple[str, str, str]]: + def _index_specs(self) -> "list[tuple[str, str, str]]": """Return ``(index_name, table, columns)`` specs for session and event indexes.""" return [*_sessions_index_specs(self._session_table), *_events_index_specs(self._events_table)] @@ -356,7 +355,7 @@ def _drop_user_states_table_sql(self) -> str: def _drop_metadata_table_sql(self) -> str: return f"DROP TABLE IF EXISTS {_table_ref(self._metadata_table)}" - def _drop_tables_sql(self) -> list[str]: + def _drop_tables_sql(self) -> "list[str]": return [ self._drop_metadata_table_sql(), self._drop_user_states_table_sql(), @@ -376,15 +375,15 @@ def _events_query( app_name: str, user_id: str, session_id: str, - after_timestamp: datetime | None = None, - limit: int | None = None, - ) -> tuple[str, tuple[Any, ...]]: + after_timestamp: "datetime | None" = None, + limit: "int | None" = None, + ) -> "tuple[str, tuple[Any, ...]]": return _events_query(self._events_table, app_name, user_id, session_id, after_timestamp, limit) def _json_column_type_sync(self) -> str: return self._json_column_type - def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> Any | None: + def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None": with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: cursor.execute(sql, params) row = cursor.fetchone() @@ -392,12 +391,12 @@ def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: b conn.commit() return row - def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: cursor.execute(sql, params) return list(cursor.fetchall()) - def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: cursor.execute(sql, params) rowcount = _cursor_rowcount(cursor) @@ -411,7 +410,7 @@ class MssqlPythonADKMemoryStore(BaseSyncADKMemoryStore["MssqlPythonConfig"]): __slots__ = () - def __init__(self, config: MssqlPythonConfig) -> None: + def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) def create_tables(self) -> None: @@ -432,7 +431,7 @@ def create_tables(self) -> None: driver.execute(_create_index_sql(index_table, index_name, columns)) driver.commit() - def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | None = None) -> int: + def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: """Bulk insert memory entries with event-id deduplication.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -442,6 +441,7 @@ def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" owner_value = ", ?" if self._owner_id_column_name else "" + # Keep the key-range lock and insertion in one statement, including autocommit. sql = f""" INSERT INTO {_table_ref(self._memory_table)} ( id, session_id, app_name, user_id, scope, event_id, author, timestamp, @@ -481,10 +481,10 @@ def search_entries( query: str, app_name: str, user_id: str, - limit: int | None = None, + limit: "int | None" = None, scope_filter: Literal["all", "user", "app"] = "all", - embedding: Sequence[float] | None = None, - ) -> list[StoredMemory]: + embedding: "Sequence[float] | None" = None, + ) -> "list[StoredMemory]": """Search memory entries by text query.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -508,7 +508,7 @@ def delete_entries_by_session(self, session_id: str) -> int: f"DELETE FROM {_table_ref(self._memory_table)} WHERE session_id = ?", (session_id,), commit=True ) - def delete_entries_older_than(self, days: int, app_name: str | None = None, scope: str | None = None) -> int: + def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: """Delete memory entries older than the retention window.""" clauses = ["inserted_at < DATEADD(day, -?, SYSUTCDATETIME())"] params: list[Any] = [days] @@ -549,7 +549,7 @@ def _memory_table_ddl(self) -> str: END; """ - def _memory_index_specs(self) -> list[tuple[str, str, str]]: + def _memory_index_specs(self) -> "list[tuple[str, str, str]]": """Return ``(index_name, table, columns)`` specs for memory-table indexes.""" return [ ( @@ -562,15 +562,15 @@ def _memory_index_specs(self) -> list[tuple[str, str, str]]: (f"idx_{self._memory_table}_timestamp", self._memory_table, "timestamp DESC"), ] - def _drop_memory_table_sql(self) -> list[str]: + def _drop_memory_table_sql(self) -> "list[str]": return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"] - def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: cursor.execute(sql, params) return list(cursor.fetchall()) - def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: cursor.execute(sql, params) rowcount = _cursor_rowcount(cursor) @@ -579,11 +579,6 @@ def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = Fal return rowcount -def _cursor_rowcount(cursor: Any) -> int: - rowcount = getattr(cursor, "rowcount", 0) - return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 - - def _adk_config(config: Any) -> MssqlPythonADKConfig: extension_config = getattr(config, "extension_config", {}) if not isinstance(extension_config, dict): @@ -594,7 +589,7 @@ def _adk_config(config: Any) -> MssqlPythonADKConfig: return cast("MssqlPythonADKConfig", adk_config) -def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: str | None) -> str: +def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: "str | None") -> str: owner_line = f",\n {owner_id_column_ddl}" if owner_id_column_ddl else "" return f""" IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo')) @@ -614,7 +609,7 @@ def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: """ -def _sessions_index_specs(table: str) -> list[tuple[str, str, str]]: +def _sessions_index_specs(table: str) -> "list[tuple[str, str, str]]": return [ (f"idx_{table}_app_user", table, "app_name, user_id"), (f"idx_{table}_update_time", table, "update_time DESC"), @@ -643,7 +638,7 @@ def _events_table_ddl(table: str, session_table: str, json_column_type: str) -> """ -def _events_index_specs(table: str) -> list[tuple[str, str, str]]: +def _events_index_specs(table: str) -> "list[tuple[str, str, str]]": return [ (f"idx_{table}_scope", table, "app_name, user_id, session_id, timestamp ASC"), (f"idx_{table}_session", table, "session_id, timestamp ASC"), @@ -699,7 +694,7 @@ def _create_index_sql(table: str, index_name: str, columns: str) -> str: return f"CREATE INDEX {_quote_identifier(index_name)} ON {_table_ref(table)} ({columns})" -def _casefold_names(rows: list[Any], key: str) -> set[str]: +def _casefold_names(rows: "list[Any]", key: str) -> "set[str]": """Collapse data-dictionary rows into a case-folded, schema-stripped name set.""" return {str(row.get(key, "")).rsplit(".", 1)[-1].casefold() for row in rows} @@ -718,7 +713,7 @@ def _insert_event_sql(table: str) -> str: """ -def _upsert_state_sql(table: str, key_columns: tuple[str, ...], key_params: tuple[str, ...]) -> str: +def _upsert_state_sql(table: str, key_columns: "tuple[str, ...]", key_params: "tuple[str, ...]") -> str: source_columns = ", ".join( f"{param} AS {_quote_identifier(column)}" for column, param in zip(key_columns, key_params, strict=False) ) @@ -754,8 +749,8 @@ def _upsert_metadata_sql(table: str) -> str: def _events_query( - table: str, app_name: str, user_id: str, session_id: str, after_timestamp: datetime | None, limit: int | None -) -> tuple[str, tuple[Any, ...]]: + table: str, app_name: str, user_id: str, session_id: str, after_timestamp: "datetime | None", limit: "int | None" +) -> "tuple[str, tuple[Any, ...]]": top_clause = "TOP (?) " if limit is not None else "" params: list[Any] = [limit] if limit is not None else [] params.extend([app_name, user_id, session_id]) @@ -772,7 +767,7 @@ def _events_query( return sql, tuple(params) -def _event_insert_params(event_record: StoredEvent) -> tuple[Any, ...]: +def _event_insert_params(event_record: StoredEvent) -> "tuple[Any, ...]": return ( event_record["id"], event_record["app_name"], @@ -802,7 +797,7 @@ def _event_record_from_row(row: Any) -> StoredEvent: ) -def _memory_record_from_row(row: Any) -> StoredMemory: +def _memory_record_from_row(row: Any) -> "StoredMemory": return cast( "StoredMemory", { @@ -825,7 +820,7 @@ def _memory_record_from_row(row: Any) -> StoredMemory: def _build_mssql_scope_where( app_name: str, user_id: str, scope_filter: Literal["all", "user", "app"] -) -> tuple[str, tuple[Any, ...]]: +) -> "tuple[str, tuple[Any, ...]]": if scope_filter == "all": return "app_name = ? AND ((scope = 'user' AND user_id = ?) OR scope = 'app')", (app_name, user_id) if scope_filter == "user": @@ -833,7 +828,7 @@ def _build_mssql_scope_where( return "app_name = ? AND scope = 'app'", (app_name,) -def _json_dict(value: Any) -> dict[str, Any]: +def _json_dict(value: Any) -> "dict[str, Any]": if value is None: return {} if isinstance(value, dict): @@ -845,6 +840,11 @@ def _json_dict(value: Any) -> dict[str, Any]: return cast("dict[str, Any]", from_json(str(value))) +def _cursor_rowcount(cursor: Any) -> int: + rowcount = getattr(cursor, "rowcount", 0) + return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 + + def _is_mssql_table_missing(exc: BaseException) -> bool: text = str(exc).lower() return "invalid object name" in text or extract_error_number(exc) == MSSQL_TABLE_NOT_FOUND_ERROR @@ -872,8 +872,14 @@ def _raise_session_not_found(session_id: str) -> None: def _session_list_query( - session_table: str, app_name: str, user_id: str | None, column: str, direction: str, limit: int | None, offset: int -) -> tuple[str, tuple[Any, ...]]: + session_table: str, + app_name: str, + user_id: "str | None", + column: str, + direction: str, + limit: "int | None", + offset: int, +) -> "tuple[str, tuple[Any, ...]]": """Return the bounded session-list query and its bound values.""" params: list[Any] = [app_name] where_clause = "app_name = ?" diff --git a/sqlspec/adapters/mssql_python/config.py b/sqlspec/adapters/mssql_python/config.py index f40982ba5..0fda4dc80 100644 --- a/sqlspec/adapters/mssql_python/config.py +++ b/sqlspec/adapters/mssql_python/config.py @@ -1,8 +1,6 @@ """mssql-python database configuration.""" -from collections.abc import Callable -from types import TracebackType -from typing import Any, ClassVar, TypedDict, cast +from typing import TYPE_CHECKING, Any, ClassVar, TypedDict, cast from typing_extensions import NotRequired @@ -12,12 +10,18 @@ from sqlspec.adapters.mssql_python.migrations import MssqlPythonSyncMigrationTracker from sqlspec.adapters.mssql_python.pool import MssqlPythonConnectionPool from sqlspec.config import ExtensionConfigs, SyncDatabaseConfig -from sqlspec.core import StatementConfig, TypeCoercionCapabilities +from sqlspec.core import TypeCoercionCapabilities from sqlspec.driver import SyncPoolConnectionContext, SyncPoolSessionFactory -from sqlspec.observability import ObservabilityConfig from sqlspec.utils.config_tools import normalize_connection_config from sqlspec.utils.serializers import from_json, to_json +if TYPE_CHECKING: + from collections.abc import Callable + from types import TracebackType + + from sqlspec.core import StatementConfig + from sqlspec.observability import ObservabilityConfig + __all__ = ( "MssqlPythonConfig", "MssqlPythonConnectionParams", @@ -92,9 +96,9 @@ class MssqlPythonDriverFeatures(TypedDict): """mssql-python driver feature flags.""" use_pool: NotRequired[bool] - json_serializer: NotRequired[Callable[[Any], str]] - json_deserializer: NotRequired[Callable[[str], Any]] - on_connection_create: NotRequired[Callable[[MssqlPythonConnection], None]] + json_serializer: "NotRequired[Callable[[Any], str]]" + json_deserializer: "NotRequired[Callable[[str], Any]]" + on_connection_create: "NotRequired[Callable[[MssqlPythonConnection], None]]" enable_events: NotRequired[bool] @@ -107,15 +111,15 @@ def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) self._conn: MssqlPythonConnection | None = None - def __enter__(self) -> MssqlPythonConnection: + def __enter__(self) -> "MssqlPythonConnection": pool = self._config.provide_pool() conn = pool.acquire() self._conn = conn return cast("MssqlPythonConnection", conn) def __exit__( - self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None - ) -> bool | None: + self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" + ) -> "bool | None": if self._conn is not None: self._config.provide_pool().release(self._conn) self._conn = None @@ -129,13 +133,13 @@ def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) self._conn: MssqlPythonConnection | None = None - def acquire_connection(self) -> MssqlPythonConnection: + def acquire_connection(self) -> "MssqlPythonConnection": pool = self._config.provide_pool() conn = pool.acquire() self._conn = conn return cast("MssqlPythonConnection", conn) - def release_connection(self, _conn: MssqlPythonConnection, **kwargs: Any) -> None: + def release_connection(self, _conn: "MssqlPythonConnection", **kwargs: Any) -> None: if self._conn is None: return self._config.provide_pool().release(self._conn) @@ -147,38 +151,38 @@ class MssqlPythonConfig(SyncDatabaseConfig[MssqlPythonConnection, MssqlPythonCon __slots__ = ("_user_connection_hook",) - driver_type: ClassVar[type[MssqlPythonDriver]] = MssqlPythonDriver - connection_type: ClassVar[type[MssqlPythonConnection]] = MssqlPythonConnection - migration_tracker_type: ClassVar[type[MssqlPythonSyncMigrationTracker]] = MssqlPythonSyncMigrationTracker - supports_transactional_ddl: ClassVar[bool] = True - supports_migration_schemas: ClassVar[bool] = True - supports_native_arrow_export: ClassVar[bool] = True - supports_native_arrow_import: ClassVar[bool] = True - supports_arrow_streaming: ClassVar[bool] = True - supports_native_row_streaming: ClassVar[bool] = True - supports_native_parquet_export: ClassVar[bool] = False - supports_native_parquet_import: ClassVar[bool] = False - type_coercion_capabilities: ClassVar[TypeCoercionCapabilities] = TypeCoercionCapabilities( + driver_type: "ClassVar[type[MssqlPythonDriver]]" = MssqlPythonDriver + connection_type: "ClassVar[type[MssqlPythonConnection]]" = MssqlPythonConnection + migration_tracker_type: "ClassVar[type[MssqlPythonSyncMigrationTracker]]" = MssqlPythonSyncMigrationTracker + supports_transactional_ddl: "ClassVar[bool]" = True + supports_migration_schemas: "ClassVar[bool]" = True + supports_native_arrow_export: "ClassVar[bool]" = True + supports_native_arrow_import: "ClassVar[bool]" = True + supports_arrow_streaming: "ClassVar[bool]" = True + supports_native_row_streaming: "ClassVar[bool]" = True + supports_native_parquet_export: "ClassVar[bool]" = False + supports_native_parquet_import: "ClassVar[bool]" = False + type_coercion_capabilities: "ClassVar[TypeCoercionCapabilities]" = TypeCoercionCapabilities( datetime_binding="native", timestamp_precision="microsecond", json_columns_decoded=False, uuid_binding="native" ) - _connection_context_class: ClassVar[type[MssqlPythonConnectionContext]] = MssqlPythonConnectionContext - _session_factory_class: ClassVar[type[_MssqlPythonSyncSessionConnectionHandler]] = ( + _connection_context_class: "ClassVar[type[MssqlPythonConnectionContext]]" = MssqlPythonConnectionContext + _session_factory_class: "ClassVar[type[_MssqlPythonSyncSessionConnectionHandler]]" = ( _MssqlPythonSyncSessionConnectionHandler ) - _session_context_class: ClassVar[type[MssqlPythonSessionContext]] = MssqlPythonSessionContext + _session_context_class: "ClassVar[type[MssqlPythonSessionContext]]" = MssqlPythonSessionContext _default_statement_config = default_statement_config def __init__( self, *, - connection_config: MssqlPythonPoolParams | dict[str, Any] | None = None, - connection_instance: MssqlPythonConnectionPool | None = None, - migration_config: dict[str, Any] | None = None, - statement_config: StatementConfig | None = None, - driver_features: MssqlPythonDriverFeatures | dict[str, Any] | None = None, - bind_key: str | None = None, - extension_config: ExtensionConfigs | None = None, - observability_config: ObservabilityConfig | None = None, + connection_config: "MssqlPythonPoolParams | dict[str, Any] | None" = None, + connection_instance: "MssqlPythonConnectionPool | None" = None, + migration_config: "dict[str, Any] | None" = None, + statement_config: "StatementConfig | None" = None, + driver_features: "MssqlPythonDriverFeatures | dict[str, Any] | None" = None, + bind_key: "str | None" = None, + extension_config: "ExtensionConfigs | None" = None, + observability_config: "ObservabilityConfig | None" = None, **kwargs: Any, ) -> None: normalized, features_dict, user_connection_hook = _normalize_mssql_python_init( @@ -198,11 +202,11 @@ def __init__( **kwargs, ) - def create_connection(self) -> MssqlPythonConnection: + def create_connection(self) -> "MssqlPythonConnection": pool = self.provide_pool() return pool.acquire() - def get_signature_namespace(self) -> dict[str, Any]: + def get_signature_namespace(self) -> "dict[str, Any]": namespace = super().get_signature_namespace() namespace.update({ "MssqlPythonConfig": MssqlPythonConfig, @@ -217,7 +221,7 @@ def get_signature_namespace(self) -> dict[str, Any]: }) return namespace - def _create_pool(self) -> MssqlPythonConnectionPool: + def _create_pool(self) -> "MssqlPythonConnectionPool": return _create_mssql_python_pool(dict(self.connection_config), self.driver_features, self._user_connection_hook) def _close_pool(self) -> None: @@ -240,10 +244,10 @@ def _apply_json_serializer_override(statement_config: Any, features_dict: dict[s def _create_mssql_python_pool( - connection_config: dict[str, Any], - driver_features: dict[str, Any], - on_connection_create: Callable[[MssqlPythonConnection], None] | None = None, -) -> MssqlPythonConnectionPool: + connection_config: "dict[str, Any]", + driver_features: "dict[str, Any]", + on_connection_create: "Callable[[MssqlPythonConnection], None] | None" = None, +) -> "MssqlPythonConnectionPool": pool_size = int(connection_config.get("pool_size", 100)) pool_idle_timeout = int(connection_config.get("pool_idle_timeout", 600)) pool_enabled = bool(connection_config.get("pool_enabled", driver_features.get("use_pool", True))) @@ -259,9 +263,9 @@ def _create_mssql_python_pool( def _normalize_mssql_python_init( - connection_config: MssqlPythonPoolParams | dict[str, Any] | None, - driver_features: MssqlPythonDriverFeatures | dict[str, Any] | None, -) -> tuple[dict[str, Any], dict[str, Any], Callable[[MssqlPythonConnection], None] | None]: + connection_config: "MssqlPythonPoolParams | dict[str, Any] | None", + driver_features: "MssqlPythonDriverFeatures | dict[str, Any] | None", +) -> "tuple[dict[str, Any], dict[str, Any], Callable[[MssqlPythonConnection], None] | None]": normalized = normalize_connection_config(connection_config) _, features_dict = apply_driver_features(default_statement_config, driver_features) hook = cast("Callable[[MssqlPythonConnection], None] | None", features_dict.pop("on_connection_create", None)) diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index c89a16a8d..6aa702a78 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -1,12 +1,9 @@ """mssql-python adapter core helpers.""" import re -from collections.abc import Callable, Mapping, Sequence from importlib.metadata import PackageNotFoundError, version -from logging import Logger -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final -from sqlspec.core import StatementConfig from sqlspec.core.parameters import ParameterStyle from sqlspec.core.parameters._registry import build_statement_config_from_profile from sqlspec.core.parameters._types import DriverParameterProfile @@ -28,6 +25,12 @@ from sqlspec.utils.serializers import from_json, to_json from sqlspec.utils.type_converters import build_uuid_coercions +if TYPE_CHECKING: + from collections.abc import Callable, Mapping, Sequence + from logging import Logger + + from sqlspec.core import StatementConfig + __all__ = ( "MSSQL_PYTHON_VERSION", "apply_driver_features", @@ -138,14 +141,10 @@ def extract_error_number(exc: BaseException | None) -> int | None: if not matches: return None last_match = matches[-1] - raw_num = last_match[0] or last_match[1] if isinstance(last_match, tuple) else last_match - try: - return int(raw_num) - except ValueError: - return None + return int(last_match[0] or last_match[1]) -def create_mapped_exception(error: Exception, *, logger: Logger | None = None) -> SQLSpecError: +def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) -> SQLSpecError: """Map a mssql-python exception to SQLSpec's exception hierarchy.""" error_number = extract_error_number(error) if error_number == _MSSQL_CONSTRAINT_547: @@ -177,25 +176,22 @@ def create_mapped_exception(error: Exception, *, logger: Logger | None = None) - return SQLSpecError(f"SQL Server database error. Original error: {error}") -def materialize_tuple_rows(fetched: Sequence[Any] | None) -> list[tuple[Any, ...]]: - """Materialize mssql-python Row objects into plain tuples. +def materialize_tuple_rows(fetched: "Sequence[Any] | None") -> "list[tuple[Any, ...]]": + """Materialize mssql-python ``Row`` objects into plain tuples. - Accesses row._values directly when available, bypassing Python's __iter__ - protocol for significantly higher throughput on large result sets. + ``mssql-python`` returns ``mssql_python.Row`` objects that are iterable and + indexable but are not ``tuple`` subclasses. The driver reports + ``row_format="tuple"``, so fetched rows are converted to real tuples to keep + that contract accurate when results are materialized. """ if not fetched: return [] - first = fetched[0] - if isinstance(first, tuple): - return list(fetched) if not isinstance(fetched, list) else fetched - if hasattr(first, "_values"): - return [tuple(row._values) if not isinstance(row._values, tuple) else row._values for row in fetched] return [tuple(row) for row in fetched] def apply_driver_features( - statement_config: StatementConfig, driver_features: Mapping[str, Any] | None -) -> tuple[StatementConfig, dict[str, Any]]: + statement_config: "StatementConfig", driver_features: "Mapping[str, Any] | None" +) -> "tuple[StatementConfig, dict[str, Any]]": """Merge mssql-python driver-feature defaults with caller overrides.""" defaults: dict[str, Any] = {"use_pool": True, "json_serializer": to_json, "json_deserializer": from_json} defaults.update(driver_features or {}) @@ -205,7 +201,7 @@ def apply_driver_features( def build_connection_config(params: dict[str, Any]) -> tuple[str, dict[str, Any]]: """Build an ODBC connection string and mssql-python connect kwargs. - When both connection_string and discrete connection fields are provided, + When both ``connection_string`` and discrete connection fields are provided, discrete fields take precedence and override matching keys in the connection string. Key names are normalized case-insensitively to prevent duplicate keywords, satisfying mssql-python driver requirements. @@ -291,7 +287,7 @@ def build_connection_config(params: dict[str, Any]) -> tuple[str, dict[str, Any] return ";".join(parts) + ";", connect_kwargs -def build_profile() -> DriverParameterProfile: +def build_profile() -> "DriverParameterProfile": """Create the mssql-python driver parameter profile.""" return DriverParameterProfile( name="mssql_python", @@ -310,14 +306,14 @@ def build_profile() -> DriverParameterProfile: ) -def build_statement_config(*, json_serializer: Callable[[Any], str] | None = None) -> StatementConfig: +def build_statement_config(*, json_serializer: "Callable[[Any], str] | None" = None) -> "StatementConfig": """Construct the mssql-python statement configuration.""" return build_statement_config_from_profile( driver_profile, statement_overrides={"dialect": "tsql"}, json_serializer=json_serializer or to_json ) -def _constraint_exception_from_message(error: Exception) -> SQLSpecError | None: +def _constraint_exception_from_message(error: Exception) -> "SQLSpecError | None": """Classify SQL Server constraint messages when a driver omits the native error number.""" message = str(error) normalized = message.lower() @@ -332,7 +328,7 @@ def _constraint_exception_from_message(error: Exception) -> SQLSpecError | None: return None -def _custom_type_coercions() -> dict[type, Callable[[Any], Any]]: +def _custom_type_coercions() -> "dict[type, Callable[[Any], Any]]": """Return custom type coercions for mssql-python.""" return {bool: _identity, int: _identity, float: _identity, bytes: _identity, **build_uuid_coercions(native=True)} diff --git a/sqlspec/adapters/mssql_python/data_dictionary.py b/sqlspec/adapters/mssql_python/data_dictionary.py index f4a2bd50e..8e2e13d6d 100644 --- a/sqlspec/adapters/mssql_python/data_dictionary.py +++ b/sqlspec/adapters/mssql_python/data_dictionary.py @@ -14,21 +14,23 @@ SystemMetadataRequest, SystemMetadataResult, TableMetadata, + VersionInfo, ensure_system_metadata_request, get_data_dictionary_loader, get_dialect_config, system_metadata_gated_result, ) from sqlspec.data_dictionary.dialects.mssql import ( - MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, build_mssql_table_ddl_result, extract_mssql_version_value, get_mssql_data_dictionary_options, + is_mssql_azure_sql, list_mssql_available_features, merge_mssql_table_lists, + mssql_supports_native_json, parse_mssql_engine_edition, parse_mssql_version_components, resolve_mssql_feature_flag, @@ -42,13 +44,51 @@ from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver from sqlspec.core import SQL - from sqlspec.data_dictionary import DialectConfig, MetadataCapabilityProfile + from sqlspec.data_dictionary._types import DialectConfig, MetadataCapabilityProfile __all__ = ("MssqlPythonSyncDataDictionary", "MssqlVersionInfo") logger = get_logger("sqlspec.adapters.mssql_python.data_dictionary") +class MssqlVersionInfo(VersionInfo): + """MSSQL database version info with build, revision, and Azure SQL detection.""" + + def __init__( + self, + major: int, + minor: int = 0, + build: int = 0, + revision: int = 0, + edition: str | None = None, + engine_edition: int | None = None, + ) -> None: + super().__init__(major, minor, 0) + self.build = build + self.revision = revision + self.edition = edition + self.engine_edition = engine_edition + self.is_azure_sql = is_mssql_azure_sql(engine_edition) + + def supports_native_json(self) -> bool: + """Return whether this server supports the native JSON type.""" + return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) + + @property + def version_tuple(self) -> "tuple[int, int, int]": + """Get version tuple using the MSSQL build number as the third component.""" + return (self.major, self.minor, self.build) + + def __str__(self) -> str: + """String representation of version info.""" + version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" + if self.edition: + version_str += f" ({self.edition})" + if self.is_azure_sql: + version_str += " [Azure]" + return version_str + + class _MssqlDataDictionaryMixin: """Shared helpers for MSSQL data dictionaries.""" @@ -92,8 +132,6 @@ def _build_version_info( def _get_optimal_type_from_version(self, version_info: MssqlVersionInfo | None, type_category: str) -> str: if type_category in {"json", "jsonb"} and version_info is not None and version_info.supports_native_json(): return "JSON" - if type_category == "vector" and version_info is not None and version_info.supports_vector(): - return "VECTOR" return self.get_dialect_config().get_optimal_type(type_category) diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index f12dba242..5eddec33a 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -1,7 +1,6 @@ """mssql-python sync and async drivers.""" import contextlib -from collections.abc import Iterable from typing import TYPE_CHECKING, Any, TypedDict, cast from typing_extensions import NotRequired @@ -21,9 +20,6 @@ ) from sqlspec.adapters.mssql_python.data_dictionary import MssqlPythonSyncDataDictionary from sqlspec.core import ( - SQL, - ArrowResult, - StatementConfig, build_arrow_result_from_reader, build_arrow_result_from_table, get_cache_config, @@ -31,25 +27,34 @@ ) from sqlspec.driver import ( BaseSyncExceptionHandler, - ExecutionResult, SyncDriverAdapterBase, SyncRowStream, rows_to_dicts, validate_savepoint_name, ) from sqlspec.exceptions import SQLSpecError -from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.utils.arrow_helpers import arrow_reader_with_deferred_close from sqlspec.utils.logging import get_logger from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.text import split_qualified_identifier if TYPE_CHECKING: + from collections.abc import Iterable + from sqlspec.builder import QueryBuilder - from sqlspec.core import Statement, StatementFilter + from sqlspec.core import SQL, ArrowResult, Statement, StatementConfig, StatementFilter + from sqlspec.driver import ExecutionResult + from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.typing import ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters -__all__ = ("MssqlPythonCursor", "MssqlPythonDriver", "MssqlPythonExceptionHandler", "MssqlPythonSessionContext") + +__all__ = ( + "MssqlPythonBulkCopyResult", + "MssqlPythonCursor", + "MssqlPythonDriver", + "MssqlPythonExceptionHandler", + "MssqlPythonSessionContext", +) logger = get_logger("sqlspec.adapters.mssql_python") _COLUMN_CACHE_MAX_SIZE = 256 @@ -61,7 +66,6 @@ class MssqlPythonBulkCopyResult(TypedDict): rows_copied: int batch_count: NotRequired[int] elapsed_time: NotRequired[float] - rows_per_second: NotRequired[float] class MssqlPythonExceptionHandler(BaseSyncExceptionHandler): @@ -69,7 +73,7 @@ class MssqlPythonExceptionHandler(BaseSyncExceptionHandler): __slots__ = () - def _handle_exception(self, exc_type: type[BaseException] | None, exc_val: BaseException) -> bool: + def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool: if exc_type is None: return False if isinstance(exc_val, MssqlPythonError): @@ -105,7 +109,7 @@ def start(self) -> None: raise self._cursor_manager = cursor_manager - def fetch_chunk(self) -> list[dict[str, Any]]: + def fetch_chunk(self) -> "list[dict[str, Any]]": cursor_manager = self._cursor_manager if cursor_manager is None or cursor_manager.cursor is None: return [] @@ -145,9 +149,9 @@ class MssqlPythonDriver(SyncDriverAdapterBase): def __init__( self, - connection: MssqlPythonConnection, - statement_config: StatementConfig | None = None, - driver_features: dict[str, Any] | None = None, + connection: "MssqlPythonConnection", + statement_config: "StatementConfig | None" = None, + driver_features: "dict[str, Any] | None" = None, ) -> None: if statement_config is None: statement_config = default_statement_config.replace( @@ -161,12 +165,12 @@ def __init__( self._transaction_active = False @property - def data_dictionary(self) -> MssqlPythonSyncDataDictionary: + def data_dictionary(self) -> "MssqlPythonSyncDataDictionary": if self._data_dictionary is None: self._data_dictionary = MssqlPythonSyncDataDictionary() return self._data_dictionary - def dispatch_execute(self, cursor: MssqlPythonRawCursor, statement: SQL) -> ExecutionResult: + def dispatch_execute(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> "ExecutionResult": sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) _execute_cursor(cursor, sql, prepared_parameters) @@ -184,28 +188,28 @@ def dispatch_execute(self, cursor: MssqlPythonRawCursor, statement: SQL) -> Exec return self.create_execution_result(cursor, rowcount_override=_cursor_rowcount(cursor)) - def dispatch_execute_many(self, cursor: MssqlPythonRawCursor, statement: SQL) -> ExecutionResult: + def dispatch_execute_many(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> "ExecutionResult": sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) cursor.executemany(sql, cast("Any", prepared_parameters)) return self.create_execution_result(cursor, rowcount_override=_cursor_rowcount(cursor), is_many_result=True) - def dispatch_execute_script(self, cursor: MssqlPythonRawCursor, statement: SQL) -> ExecutionResult: + def dispatch_execute_script(self, cursor: "MssqlPythonRawCursor", statement: "SQL") -> "ExecutionResult": sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True) successful_count = 0 for stmt in statements: - _execute_cursor(cursor, stmt, prepared_parameters, use_prepare=False) + _execute_cursor(cursor, stmt, prepared_parameters) successful_count += 1 return self.create_execution_result( cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) - def collect_rows(self, cursor: MssqlPythonRawCursor, fetched: list[Any]) -> tuple[list[Any], list[str], int]: + def collect_rows(self, cursor: "MssqlPythonRawCursor", fetched: "list[Any]") -> "tuple[list[Any], list[str], int]": column_names = _resolve_column_names(cursor.description, self._column_name_cache) rows = materialize_tuple_rows(fetched) return rows, column_names, len(rows) - def resolve_rowcount(self, cursor: MssqlPythonRawCursor) -> int: + def resolve_rowcount(self, cursor: "MssqlPythonRawCursor") -> int: return _cursor_rowcount(cursor) def begin(self) -> None: @@ -239,13 +243,13 @@ def rollback(self) -> None: self._transaction_active = False self._restore_connection_autocommit() - def with_cursor(self, connection: MssqlPythonConnection) -> MssqlPythonCursor: + def with_cursor(self, connection: "MssqlPythonConnection") -> "MssqlPythonCursor": return MssqlPythonCursor(connection) - def handle_database_exceptions(self) -> MssqlPythonExceptionHandler: + def handle_database_exceptions(self) -> "MssqlPythonExceptionHandler": return MssqlPythonExceptionHandler() - def dispatch_select_stream(self, statement: SQL, chunk_size: int) -> SyncRowStream[dict[str, Any]] | None: + def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "SyncRowStream[dict[str, Any]] | None": """Return a native mssql-python row stream backed by ``fetchmany()``.""" if not statement.returns_rows(): return None @@ -265,17 +269,13 @@ def set_migration_session_schema(self, schema: str) -> None: """Point the database user's default schema at the migration schema, remembering the prior one.""" with self.with_cursor(self.connection) as cursor: if self._migration_schema_restore is None: - _execute_cursor( - cursor, "SELECT USER_NAME() AS user_name, SCHEMA_NAME() AS schema_name;", None, use_prepare=False - ) + _execute_cursor(cursor, "SELECT USER_NAME() AS user_name, SCHEMA_NAME() AS schema_name;", None) row: Any = cursor.fetchone() user_name, current_schema = row[0], row[1] - _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None, use_prepare=False) + _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None) self._migration_schema_restore = (str(user_name), str(current_schema)) return - _execute_cursor( - cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None, use_prepare=False - ) + _execute_cursor(cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None) def reset_migration_session_schema(self) -> None: """Restore the user's default schema captured by set_migration_session_schema and commit it.""" @@ -283,7 +283,7 @@ def reset_migration_session_schema(self) -> None: return user_name, previous_schema = self._migration_schema_restore with self.with_cursor(self.connection) as cursor: - _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None, use_prepare=False) + _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None) self.connection.commit() self._migration_schema_restore = None @@ -298,13 +298,13 @@ def select_to_arrow( statement: "Statement | QueryBuilder", /, *parameters: "StatementParameters | StatementFilter", - statement_config: StatementConfig | None = None, + statement_config: "StatementConfig | None" = None, return_format: "ArrowReturnFormat" = "table", native_only: bool = False, batch_size: int | None = None, arrow_schema: Any = None, **kwargs: Any, - ) -> ArrowResult: + ) -> "ArrowResult": """Execute a query and return native mssql-python Arrow results.""" ensure_pyarrow() config = statement_config or self.statement_config @@ -314,7 +314,7 @@ def select_to_arrow( arrow_kwargs: dict[str, int] = {"batch_size": batch_size} if batch_size is not None else {} table: Any | None = None - if return_format in ("reader", "batches"): + if return_format == "reader": cursor_manager = self.with_cursor(self.connection) cursor = None reader: object | None = None @@ -355,6 +355,16 @@ def select_to_arrow( exc_handler = self.handle_database_exceptions() with exc_handler, self.with_cursor(self.connection) as cursor: _execute_cursor(cursor, sql, prepared_parameters) + if return_format == "batches": + reader = _cursor_arrow_reader(cursor, arrow_kwargs) + if reader is not None: + return build_arrow_result_from_reader( + prepared_statement, + reader, + return_format=return_format, + batch_size=batch_size, + arrow_schema=arrow_schema, + ) table = cursor.arrow(**arrow_kwargs) self._check_pending_exception(exc_handler) @@ -365,10 +375,10 @@ def select_to_arrow( prepared_statement, table, return_format=return_format, batch_size=batch_size, arrow_schema=arrow_schema ) - def _bulk_copy( + def bulk_copy( self, target_table: str, - rows: Iterable[tuple[Any, ...]], + rows: "Iterable[tuple[Any, ...]]", *, batch_size: int = 0, timeout: int = 30, @@ -404,126 +414,26 @@ def _bulk_copy( def load_from_arrow( self, table: str, - source: ArrowResult | Any, + source: "ArrowResult | Any", *, - partitioner: dict[str, object] | None = None, + partitioner: "dict[str, object] | None" = None, overwrite: bool = False, - telemetry: StorageTelemetry | None = None, - batch_size: int = 0, - timeout: int = 30, - table_lock: bool = True, - check_constraints: bool = False, - fire_triggers: bool = False, - keep_identity: bool = False, - keep_nulls: bool = False, - use_internal_transaction: bool = False, - column_mappings: list[str] | list[tuple[int, str]] | None = None, - ) -> StorageBridgeJob: + telemetry: "StorageTelemetry | None" = None, + ) -> "StorageBridgeJob": """Load Arrow data into SQL Server via BulkCopy.""" self._require_capability("arrow_import_enabled") + arrow_table = self._coerce_arrow_table(source) if overwrite: - quoted_table = _quote_mssql_table(table) exc_handler = self.handle_database_exceptions() with exc_handler, self.with_cursor(self.connection) as cursor: - try: - _execute_cursor(cursor, f"TRUNCATE TABLE {quoted_table}", None, use_prepare=False) - except Exception as exc: - error_msg = str(exc) - if "4712" in error_msg or "foreign key" in error_msg.lower(): - _execute_cursor(cursor, f"DELETE FROM {quoted_table}", None, use_prepare=False) - else: - raise + cursor.execute(f"DELETE FROM {_quote_mssql_table(table)}") self._check_pending_exception(exc_handler) - - raw_result: Any = None - is_table_source = isinstance(source, ArrowResult) - is_reader = False - try: - import pyarrow as pa - - is_table_source = is_table_source or isinstance(source, pa.Table) - is_reader = isinstance(source, (pa.RecordBatchReader, pa.RecordBatch)) - except ImportError: - pass - is_stream = not is_table_source and (is_reader or hasattr(source, "__arrow_c_stream__")) - - if is_stream: - cols = column_mappings - source_schema = getattr(source, "schema", None) - if cols is None and source_schema is not None: - schema_names = getattr(source_schema, "names", None) - if schema_names is not None: - cols = list(schema_names) + if arrow_table.num_rows: exc_handler = self.handle_database_exceptions() with exc_handler, self.with_cursor(self.connection) as cursor: - raw_result = cursor.bulkcopy_arrow( - table, - source, - batch_size=batch_size, - timeout=timeout, - table_lock=table_lock, - check_constraints=check_constraints, - fire_triggers=fire_triggers, - keep_identity=keep_identity, - keep_nulls=keep_nulls, - use_internal_transaction=use_internal_transaction, - column_mappings=cols, - ) + cursor.bulkcopy_arrow(table, arrow_table, column_mappings=list(arrow_table.column_names)) self._check_pending_exception(exc_handler) - telemetry_payload = cast("StorageTelemetry", {"destination": table, "format": "arrow", "extra": {}}) - else: - arrow_table = self._coerce_arrow_table(source) - cols = column_mappings if column_mappings is not None else list(arrow_table.column_names) - if arrow_table.num_rows: - exc_handler = self.handle_database_exceptions() - use_fallback = False - with exc_handler, self.with_cursor(self.connection) as cursor: - if hasattr(cursor, "bulkcopy_arrow"): - raw_result = cursor.bulkcopy_arrow( - table, - arrow_table, - batch_size=batch_size, - timeout=timeout, - table_lock=table_lock, - check_constraints=check_constraints, - fire_triggers=fire_triggers, - keep_identity=keep_identity, - keep_nulls=keep_nulls, - use_internal_transaction=use_internal_transaction, - column_mappings=cols, - ) - else: - use_fallback = True - self._check_pending_exception(exc_handler) - if use_fallback: - _, records = self._arrow_table_to_rows(arrow_table) - raw_result = self._bulk_copy( - table, - records, - batch_size=batch_size, - timeout=timeout, - column_mappings=cols, - keep_identity=keep_identity, - check_constraints=check_constraints, - table_lock=table_lock, - keep_nulls=keep_nulls, - fire_triggers=fire_triggers, - use_internal_transaction=use_internal_transaction, - ) - telemetry_payload = self._ingest_telemetry(arrow_table) - - extra = telemetry_payload.setdefault("extra", {}) - if isinstance(raw_result, dict): - if "rows_copied" in raw_result: - telemetry_payload["rows_processed"] = raw_result["rows_copied"] - extra["rows_ingested"] = raw_result["rows_copied"] - if "elapsed_time" in raw_result: - extra["elapsed_time"] = raw_result["elapsed_time"] - if "rows_per_second" in raw_result: - extra["rows_per_second"] = raw_result["rows_per_second"] - if "batch_count" in raw_result: - extra["batch_count"] = raw_result["batch_count"] - + telemetry_payload = self._ingest_telemetry(arrow_table) telemetry_payload["destination"] = table self._attach_partition_telemetry(telemetry_payload, partitioner) return self._storage_job(telemetry_payload, telemetry) @@ -531,12 +441,12 @@ def load_from_arrow( def load_from_storage( self, table: str, - source: StorageDestination, + source: "StorageDestination", *, - file_format: StorageFormat, - partitioner: dict[str, object] | None = None, + file_format: "StorageFormat", + partitioner: "dict[str, object] | None" = None, overwrite: bool = False, - ) -> StorageBridgeJob: + ) -> "StorageBridgeJob": """Load staged artifacts from storage into SQL Server via BulkCopy.""" arrow_table, inbound = self._read_storage_arrow(source, file_format=file_format) return self.load_from_arrow(table, arrow_table, partitioner=partitioner, overwrite=overwrite, telemetry=inbound) @@ -566,27 +476,19 @@ def _quote_mssql_table(table: str) -> str: return ".".join(_quote_tsql_identifier(part) for part in split_qualified_identifier(table)) -def _execute_cursor(cursor: MssqlPythonRawCursor, sql: str, parameters: Any, *, use_prepare: bool = True) -> None: - if use_prepare or parameters: - if parameters is None: - cursor.execute(sql) - else: - cursor.execute(sql, parameters) - return - try: - cursor.execute(sql, use_prepare=False) - except TypeError as exc: - if "use_prepare" not in str(exc): - raise +def _execute_cursor(cursor: "MssqlPythonRawCursor", sql: str, parameters: Any) -> None: + if parameters is None: cursor.execute(sql) + else: + cursor.execute(sql, parameters) -def _cursor_rowcount(cursor: MssqlPythonRawCursor) -> int: +def _cursor_rowcount(cursor: "MssqlPythonRawCursor") -> int: rowcount = getattr(cursor, "rowcount", 0) return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0 -def _resolve_column_names(description: Any, cache: dict[int, tuple[Any, list[str]]]) -> list[str]: +def _resolve_column_names(description: Any, cache: "dict[int, tuple[Any, list[str]]]") -> list[str]: if not description: return [] cache_key = id(description) @@ -600,14 +502,16 @@ def _resolve_column_names(description: Any, cache: dict[int, tuple[Any, list[str return column_names -def _cursor_arrow_reader(cursor: MssqlPythonRawCursor, arrow_kwargs: dict[str, int]) -> "ArrowRecordBatchReader | None": +def _cursor_arrow_reader( + cursor: "MssqlPythonRawCursor", arrow_kwargs: "dict[str, int]" +) -> "ArrowRecordBatchReader | None": arrow_reader = getattr(cursor, "arrow_reader", None) if not callable(arrow_reader): return None return cast("ArrowRecordBatchReader", arrow_reader(**arrow_kwargs)) -def _coerce_bulk_copy_result(result: Any, cursor: MssqlPythonRawCursor) -> MssqlPythonBulkCopyResult: +def _coerce_bulk_copy_result(result: Any, cursor: "MssqlPythonRawCursor") -> MssqlPythonBulkCopyResult: if isinstance(result, dict): return cast("MssqlPythonBulkCopyResult", dict(result)) return {"rows_copied": _cursor_rowcount(cursor)} diff --git a/sqlspec/adapters/mssql_python/litestar/store.py b/sqlspec/adapters/mssql_python/litestar/store.py index 386ea45f4..d48001e5b 100644 --- a/sqlspec/adapters/mssql_python/litestar/store.py +++ b/sqlspec/adapters/mssql_python/litestar/store.py @@ -1,13 +1,15 @@ """mssql-python Litestar Store implementation.""" from datetime import datetime, timedelta, timezone -from typing import Any +from typing import TYPE_CHECKING, Any from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor -from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ +if TYPE_CHECKING: + from sqlspec.adapters.mssql_python.config import MssqlPythonConfig + __all__ = ("MssqlPythonStore",) @@ -16,7 +18,7 @@ class MssqlPythonStore(BaseSQLSpecStore["MssqlPythonConfig"]): __slots__ = () - def __init__(self, config: MssqlPythonConfig) -> None: + def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) async def create_table(self) -> None: @@ -27,11 +29,11 @@ async def create_table(self) -> None: await async_(self._create_table)() await self.reconcile_schema(assume_existing=True) - async def get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: + async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": """Get a session value by key.""" return await async_(self._get)(key, renew_for) - async def set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: + async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: """Store a session value.""" await async_(self._set)(key, value, expires_in) @@ -47,7 +49,7 @@ async def exists(self, key: str) -> bool: """Check if a session key exists and is not expired.""" return await async_(self._exists)(key) - async def expires_in(self, key: str) -> int | None: + async def expires_in(self, key: str) -> "int | None": """Get the time in seconds until the session expires.""" return await async_(self._expires_in)(key) @@ -79,7 +81,7 @@ def _table_ddl(self) -> str: END; """ - def _drop_table_sql(self) -> list[str]: + def _drop_table_sql(self) -> "list[str]": """Get SQL Server DROP TABLE statements.""" return [f"IF OBJECT_ID(N'dbo.{self._table_name}', N'U') IS NOT NULL DROP TABLE dbo.{self._table_name};"] @@ -89,7 +91,7 @@ def _create_table(self) -> None: driver.commit() self._log_table_created() - def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: + def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = ? @@ -120,7 +122,7 @@ def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | No return _coerce_bytes(_row_value(row, "data", 0)) - def _set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: + def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) sql = f""" @@ -162,7 +164,7 @@ def _exists(self, key: str) -> bool: cursor.execute(sql, (key,)) return cursor.fetchone() is not None - def _expires_in(self, key: str) -> int | None: + def _expires_in(self, key: str) -> "int | None": with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor: cursor.execute(f"SELECT expires_at FROM {self._table_name} WHERE session_id = ?", (key,)) row = cursor.fetchone() @@ -204,7 +206,7 @@ def _row_value(row: object, key: str, index: int) -> Any: return getattr(row, key, None) -def _normalize_utc(value: Any) -> datetime | None: +def _normalize_utc(value: Any) -> "datetime | None": if value is None: return None if not isinstance(value, datetime): diff --git a/sqlspec/adapters/mssql_python/pool.py b/sqlspec/adapters/mssql_python/pool.py index a9527e7f5..9d9a15faa 100644 --- a/sqlspec/adapters/mssql_python/pool.py +++ b/sqlspec/adapters/mssql_python/pool.py @@ -1,15 +1,16 @@ """mssql-python pool facade.""" -import contextlib import warnings -from collections.abc import Callable -from typing import Any +from typing import TYPE_CHECKING, Any, cast -from sqlspec.adapters.mssql_python._typing import MssqlPythonConnection, mssql_python_module +from sqlspec.adapters.mssql_python._typing import MSSQL_PYTHON_MODULE, MssqlPythonConnection + +if TYPE_CHECKING: + from collections.abc import Callable __all__ = ("MssqlPythonConnectionPool",) -_POOLING_PARAMS: tuple[int, int, bool] | None = None +_POOLING_PARAMS: "tuple[int, int, bool] | None" = None class MssqlPythonConnectionPool: @@ -29,11 +30,11 @@ def __init__( self, *, connection_string: str, - connect_kwargs: dict[str, Any] | None = None, + connect_kwargs: "dict[str, Any] | None" = None, max_size: int = 100, idle_timeout: int = 600, enabled: bool = True, - on_connection_create: Callable[[MssqlPythonConnection], None] | None = None, + on_connection_create: "Callable[[MssqlPythonConnection], None] | None" = None, ) -> None: self.connection_string = connection_string self.connect_kwargs = connect_kwargs or {} @@ -51,27 +52,22 @@ def __init__( stacklevel=2, ) if _POOLING_PARAMS is None or new_params != _POOLING_PARAMS: - mssql_python_module.pooling(max_size=max_size, idle_timeout=idle_timeout, enabled=enabled) + MSSQL_PYTHON_MODULE.pooling(max_size=max_size, idle_timeout=idle_timeout, enabled=enabled) _POOLING_PARAMS = new_params - def acquire(self) -> MssqlPythonConnection: + def acquire(self) -> "MssqlPythonConnection": if self._closed: msg = "Cannot acquire a connection from a closed mssql-python pool." raise RuntimeError(msg) - connection = mssql_python_module.connect(self.connection_string, **self.connect_kwargs) + connection = cast( + "MssqlPythonConnection", MSSQL_PYTHON_MODULE.connect(self.connection_string, **self.connect_kwargs) + ) if self.on_connection_create is not None: self.on_connection_create(connection) return connection - def release(self, connection: MssqlPythonConnection) -> None: + def release(self, connection: "MssqlPythonConnection") -> None: connection.close() - def close(self, *, close_driver_pooling: bool = False) -> None: + def close(self) -> None: self._closed = True - if close_driver_pooling: - global _POOLING_PARAMS - _POOLING_PARAMS = None - with contextlib.suppress(Exception): - ddbc = getattr(mssql_python_module, "ddbc_bindings", None) - if ddbc is not None and hasattr(ddbc, "close_pooling"): - ddbc.close_pooling() diff --git a/sqlspec/adapters/mssql_python/type_converter.py b/sqlspec/adapters/mssql_python/type_converter.py index cf6ced2fe..7099e801a 100644 --- a/sqlspec/adapters/mssql_python/type_converter.py +++ b/sqlspec/adapters/mssql_python/type_converter.py @@ -1,13 +1,14 @@ """Type converters for mssql-python parameter binding.""" -from collections.abc import Callable from typing import TYPE_CHECKING, Any, Final, cast from uuid import UUID -from sqlspec.utils.module_loader import ensure_pyarrow, import_optional +from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: + from collections.abc import Callable + import pyarrow as pa __all__ = ("MssqlPythonTypeConverter", "mssql_type_to_arrow") @@ -41,7 +42,6 @@ "nvarchar": ("string", (), {}), "text": ("string", (), {}), "ntext": ("string", (), {}), - "json": ("string", (), {}), } @@ -56,12 +56,12 @@ class MssqlPythonTypeConverter: __slots__ = ("_json_deserializer", "_json_serializer") def __init__( - self, json_serializer: Callable[[Any], str] = to_json, json_deserializer: Callable[[str], Any] = from_json + self, json_serializer: "Callable[[Any], str]" = to_json, json_deserializer: "Callable[[str], Any]" = from_json ) -> None: self._json_serializer = json_serializer self._json_deserializer = json_deserializer - def coerce_bind_value(self, value: Any) -> Any: + def coerce_bind_value(self, value: "Any") -> "Any": """Coerce Python values before mssql-python parameter binding.""" if isinstance(value, (dict, list)): return self._json_serializer(value) @@ -69,7 +69,7 @@ def coerce_bind_value(self, value: Any) -> Any: return value return value - def coerce_read_value(self, value: Any) -> Any: + def coerce_read_value(self, value: "Any") -> "Any": """Coerce mssql-python result values after fetching.""" return value @@ -77,10 +77,6 @@ def coerce_read_value(self, value: Any) -> Any: def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: int | None = None) -> "pa.DataType": """Resolve a T-SQL type name to an Arrow data type.""" normalized_type = sql_type.lower().split("(", 1)[0].strip() - if normalized_type == "vector": - ensure_pyarrow() - pyarrow_mod = cast("Any", import_optional("pyarrow")) - return cast("pa.DataType", pyarrow_mod.list_(pyarrow_mod.float32())) if normalized_type in {"decimal", "numeric"} and precision is not None and scale is not None: return _arrow_type("decimal128", (precision, scale)) spec = _MSSQL_ARROW_TYPE_SPECS.get(normalized_type) @@ -92,5 +88,6 @@ def mssql_type_to_arrow(sql_type: str, *, precision: int | None = None, scale: i def _arrow_type(name: str, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None) -> "pa.DataType": ensure_pyarrow() - pyarrow_mod = cast("Any", import_optional("pyarrow")) - return cast("pa.DataType", getattr(pyarrow_mod, name)(*args, **(kwargs or {}))) + import pyarrow as pa + + return cast("pa.DataType", getattr(pa, name)(*args, **(kwargs or {}))) diff --git a/sqlspec/adapters/pymssql/_typing.py b/sqlspec/adapters/pymssql/_typing.py index 9c6cea9cd..3ae4e9fab 100644 --- a/sqlspec/adapters/pymssql/_typing.py +++ b/sqlspec/adapters/pymssql/_typing.py @@ -7,25 +7,39 @@ import contextlib from typing import TYPE_CHECKING, Any -import pymssql as pymssql_module -from pymssql import Connection as PymssqlConnection -from pymssql import Cursor as PymssqlRawCursor +import pymssql as _pymssql # pyright: ignore[reportMissingTypeStubs] +from pymssql import Connection as _PymssqlConnection # pyright: ignore[reportMissingTypeStubs] +from pymssql import Cursor as _PymssqlRawCursor # pyright: ignore[reportMissingTypeStubs] from pymssql import Error as PymssqlError +PYMSSQL_MODULE = _pymssql + if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType + from typing import TypeAlias + + from pymssql._pymssql import QueryParams as PymssqlQueryParams from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.core import StatementConfig + PymssqlConnection: TypeAlias = _PymssqlConnection + PymssqlRawCursor: TypeAlias = _PymssqlRawCursor + +if not TYPE_CHECKING: + PymssqlQueryParams = Any + PymssqlConnection = _PymssqlConnection + PymssqlRawCursor = _PymssqlRawCursor + __all__ = ( + "PYMSSQL_MODULE", "PymssqlConnection", "PymssqlCursor", "PymssqlError", + "PymssqlQueryParams", "PymssqlRawCursor", "PymssqlSessionContext", - "pymssql_module", ) diff --git a/sqlspec/adapters/pymssql/adk/store.py b/sqlspec/adapters/pymssql/adk/store.py index b326ab9fc..7d1737f93 100644 --- a/sqlspec/adapters/pymssql/adk/store.py +++ b/sqlspec/adapters/pymssql/adk/store.py @@ -1,28 +1,27 @@ """pymssql ADK stores for Google Agent Development Kit session storage.""" -from collections.abc import Sequence -from datetime import datetime, timedelta -from typing import Any, ClassVar, Final, Literal, cast +from datetime import datetime +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast from typing_extensions import NotRequired from sqlspec.adapters.pymssql._typing import PymssqlCursor, PymssqlError -from sqlspec.adapters.pymssql.config import PymssqlConfig -from sqlspec.adapters.pymssql.core import extract_error_number, quote_tsql_identifier, resolve_rowcount +from sqlspec.adapters.pymssql.core import extract_error_number, resolve_rowcount from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo -from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.config import ADKConfig -from sqlspec.extensions.adk import ( - BaseSyncADKMemoryStore, - BaseSyncADKStore, - SessionOrderBy, - StoredEvent, - StoredMemory, - StoredSession, - normalize_session_list_options, -) +from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options +from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore from sqlspec.utils.serializers import from_json, to_json +if TYPE_CHECKING: + from collections.abc import Sequence + from datetime import timedelta + + from sqlspec.adapters.pymssql.config import PymssqlConfig + from sqlspec.adapters.pymssql.driver import PymssqlDriver + from sqlspec.extensions.adk import SessionOrderBy + from sqlspec.extensions.adk.memory._types import StoredMemory + __all__ = ("PymssqlADKConfig", "PymssqlADKMemoryStore", "PymssqlADKStore") MSSQL_TABLE_NOT_FOUND_ERROR: Final[int] = 208 @@ -46,7 +45,7 @@ class PymssqlADKStore(BaseSyncADKStore["PymssqlConfig"]): connector_name: ClassVar[str] = "pymssql" __slots__ = ("_json_column_type", "_native_json") - def __init__(self, config: PymssqlConfig) -> None: + def __init__(self, config: "PymssqlConfig") -> None: super().__init__(config) adk_config = _adk_config(config) native_json = adk_config.get("native_json") @@ -74,7 +73,7 @@ def create_tables(self) -> None: driver.commit() def create_session( - self, session_id: str, app_name: str, user_id: str, state: dict[str, Any], owner_id: Any | None = None + self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None ) -> StoredSession: """Create a new ADK session.""" owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" @@ -98,8 +97,8 @@ def create_session( return _session_record_from_row(row) def get_session( - self, app_name: str, user_id: str, session_id: str, *, renew_for: int | timedelta | None = None - ) -> StoredSession | None: + self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None + ) -> "StoredSession | None": """Return a scoped session or ``None`` if absent.""" try: if renew_for is not None and self._calculate_expires_at(renew_for) is not None: @@ -126,7 +125,7 @@ def get_session( raise return _session_record_from_row(row) if row is not None else None - def update_session_state(self, app_name: str, user_id: str, session_id: str, state: dict[str, Any]) -> None: + def update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None: """Replace a session's durable state.""" self._execute( f""" @@ -141,13 +140,13 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta def list_sessions( self, app_name: str, - user_id: str | None = None, + user_id: "str | None" = None, *, - order_by: SessionOrderBy = "update_time", + order_by: "SessionOrderBy" = "update_time", descending: bool = True, - limit: int | None = None, - offset: int | None = None, - ) -> list[StoredSession]: + limit: "int | None" = None, + offset: "int | None" = None, + ) -> "list[StoredSession]": """List ADK sessions for an application, optionally scoped to a user.""" column, direction, page_limit, page_offset = normalize_session_list_options(order_by, descending, limit, offset) if page_limit == 0: @@ -182,10 +181,10 @@ def append_event_and_update_state( app_name: str, user_id: str, session_id: str, - state: dict[str, Any], + state: "dict[str, Any]", *, - app_state: dict[str, Any] | None = None, - user_state: dict[str, Any] | None = None, + app_state: "dict[str, Any] | None" = None, + user_state: "dict[str, Any] | None" = None, ) -> StoredSession: """Atomically append an event and update durable session/scoped state.""" update_sql = f""" @@ -216,9 +215,9 @@ def get_events( app_name: str, user_id: str, session_id: str, - after_timestamp: datetime | None = None, - limit: int | None = None, - ) -> list[StoredEvent]: + after_timestamp: "datetime | None" = None, + limit: "int | None" = None, + ) -> "list[StoredEvent]": """Return events for a scoped session ordered by event timestamp.""" if limit == 0: return [] @@ -231,7 +230,7 @@ def get_events( raise return [_event_record_from_row(row) for row in rows] - def delete_expired_events(self, before: datetime, app_name: str | None = None) -> int: + def delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: """Delete events older than ``before``.""" sql = f"DELETE FROM {_table_ref(self._events_table)} WHERE timestamp < %s" params: list[Any] = [before] @@ -245,7 +244,7 @@ def delete_expired_events(self, before: datetime, app_name: str | None = None) - return 0 raise - def delete_idle_sessions(self, updated_before: datetime, app_name: str | None = None) -> int: + def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: """Delete sessions whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._session_table)} WHERE update_time < %s" params: list[Any] = [updated_before] @@ -259,7 +258,7 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: str | None = return 0 raise - def delete_idle_user_states(self, updated_before: datetime, app_name: str | None = None) -> int: + def delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: """Delete user state rows whose update_time is older than ``updated_before``.""" sql = f"DELETE FROM {_table_ref(self._user_state_table)} WHERE update_time < %s" params: list[Any] = [updated_before] @@ -273,7 +272,7 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: str | None return 0 raise - def get_app_state(self, app_name: str) -> dict[str, Any] | None: + def get_app_state(self, app_name: str) -> "dict[str, Any] | None": """Return app-scoped state.""" try: row = self._execute_fetchone( @@ -285,7 +284,7 @@ def get_app_state(self, app_name: str) -> dict[str, Any] | None: raise return _json_dict(row[0]) if row is not None else None - def get_user_state(self, app_name: str, user_id: str) -> dict[str, Any] | None: + def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None": """Return user-scoped state.""" try: row = self._execute_fetchone( @@ -302,15 +301,15 @@ def get_user_state(self, app_name: str, user_id: str) -> dict[str, Any] | None: raise return _json_dict(row[0]) if row is not None else None - def upsert_app_state(self, app_name: str, state: dict[str, Any]) -> None: + def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """Insert or replace app-scoped state.""" self._execute(self._upsert_app_state_sql(), (app_name, to_json(state)), commit=True) - def upsert_user_state(self, app_name: str, user_id: str, state: dict[str, Any]) -> None: + def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: """Insert or replace user-scoped state.""" self._execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(state)), commit=True) - def get_metadata(self, key: str) -> str | None: + def get_metadata(self, key: str) -> "str | None": """Return an ADK metadata value.""" try: row = self._execute_fetchone( @@ -326,7 +325,7 @@ def set_metadata(self, key: str, value: str) -> None: """Set an ADK metadata value.""" self._execute(_upsert_metadata_sql(self._metadata_table), (key, value), commit=True) - def _index_specs(self) -> list[tuple[str, str, str]]: + def _index_specs(self) -> "list[tuple[str, str, str]]": """Return ``(index_name, table, columns)`` specs for session and event indexes.""" return [*_sessions_index_specs(self._session_table), *_events_index_specs(self._events_table)] @@ -359,7 +358,7 @@ def _drop_user_states_table_sql(self) -> str: def _drop_metadata_table_sql(self) -> str: return f"DROP TABLE IF EXISTS {_table_ref(self._metadata_table)}" - def _drop_tables_sql(self) -> list[str]: + def _drop_tables_sql(self) -> "list[str]": return [ self._drop_metadata_table_sql(), self._drop_user_states_table_sql(), @@ -379,9 +378,9 @@ def _events_query( app_name: str, user_id: str, session_id: str, - after_timestamp: datetime | None = None, - limit: int | None = None, - ) -> tuple[str, tuple[Any, ...]]: + after_timestamp: "datetime | None" = None, + limit: "int | None" = None, + ) -> "tuple[str, tuple[Any, ...]]": return _events_query(self._events_table, app_name, user_id, session_id, after_timestamp, limit) def _json_column_type_sync(self) -> str: @@ -395,7 +394,7 @@ def _json_column_type_sync(self) -> str: self._json_column_type = _json_column_type_from_sync_driver(driver) return self._json_column_type - def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> Any | None: + def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None": with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: cursor.execute(sql, params) row = cursor.fetchone() @@ -403,12 +402,12 @@ def _execute_fetchone(self, sql: str, params: tuple[Any, ...] = (), *, commit: b conn.commit() return row - def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: cursor.execute(sql, params) return list(cursor.fetchall()) - def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: cursor.execute(sql, params) rowcount = resolve_rowcount(cursor) @@ -422,7 +421,7 @@ class PymssqlADKMemoryStore(BaseSyncADKMemoryStore["PymssqlConfig"]): __slots__ = () - def __init__(self, config: PymssqlConfig) -> None: + def __init__(self, config: "PymssqlConfig") -> None: super().__init__(config) def create_tables(self) -> None: @@ -443,7 +442,7 @@ def create_tables(self) -> None: driver.execute(_create_index_sql(index_table, index_name, columns)) driver.commit() - def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | None = None) -> int: + def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: """Bulk insert memory entries with event-id deduplication.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -453,6 +452,7 @@ def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" owner_value = ", %s" if self._owner_id_column_name else "" + # Keep the key-range lock and insertion in one statement, including autocommit. sql = f""" INSERT INTO {_table_ref(self._memory_table)} ( id, session_id, app_name, user_id, scope, event_id, author, timestamp, @@ -492,10 +492,10 @@ def search_entries( query: str, app_name: str, user_id: str, - limit: int | None = None, + limit: "int | None" = None, scope_filter: Literal["all", "user", "app"] = "all", - embedding: Sequence[float] | None = None, - ) -> list[StoredMemory]: + embedding: "Sequence[float] | None" = None, + ) -> "list[StoredMemory]": """Search memory entries by text query.""" if not self._enabled: msg = "ADK memory store is disabled" @@ -519,7 +519,7 @@ def delete_entries_by_session(self, session_id: str) -> int: f"DELETE FROM {_table_ref(self._memory_table)} WHERE session_id = %s", (session_id,), commit=True ) - def delete_entries_older_than(self, days: int, app_name: str | None = None, scope: str | None = None) -> int: + def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: """Delete memory entries older than the retention window.""" clauses = ["inserted_at < DATEADD(day, -%s, SYSUTCDATETIME())"] params: list[Any] = [days] @@ -560,7 +560,7 @@ def _memory_table_ddl(self) -> str: END; """ - def _memory_index_specs(self) -> list[tuple[str, str, str]]: + def _memory_index_specs(self) -> "list[tuple[str, str, str]]": """Return ``(index_name, table, columns)`` specs for memory-table indexes.""" return [ ( @@ -573,15 +573,15 @@ def _memory_index_specs(self) -> list[tuple[str, str, str]]: (f"idx_{self._memory_table}_timestamp", self._memory_table, "timestamp DESC"), ] - def _drop_memory_table_sql(self) -> list[str]: + def _drop_memory_table_sql(self) -> "list[str]": return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"] - def _execute_fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[Any]: + def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]": with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: cursor.execute(sql, params) return list(cursor.fetchall()) - def _execute(self, sql: str, params: tuple[Any, ...] = (), *, commit: bool = False) -> int: + def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int: with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: cursor.execute(sql, params) rowcount = resolve_rowcount(cursor) @@ -600,20 +600,20 @@ def _adk_config(config: Any) -> PymssqlADKConfig: return cast("PymssqlADKConfig", adk_config) -def _configured_json_column_type(native_json: bool | None) -> str | None: +def _configured_json_column_type(native_json: "bool | None") -> "str | None": if native_json is True: return JSON_NATIVE_COLUMN_TYPE return JSON_FALLBACK_COLUMN_TYPE -def _json_column_type_from_sync_driver(driver: PymssqlDriver) -> str: +def _json_column_type_from_sync_driver(driver: "PymssqlDriver") -> str: version_info = driver.data_dictionary.get_version(driver) if isinstance(version_info, MssqlVersionInfo) and version_info.supports_native_json(): return JSON_NATIVE_COLUMN_TYPE return JSON_FALLBACK_COLUMN_TYPE -def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: str | None) -> str: +def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: "str | None") -> str: owner_line = f",\n {owner_id_column_ddl}" if owner_id_column_ddl else "" return f""" IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo')) @@ -633,7 +633,7 @@ def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: """ -def _sessions_index_specs(table: str) -> list[tuple[str, str, str]]: +def _sessions_index_specs(table: str) -> "list[tuple[str, str, str]]": return [ (f"idx_{table}_app_user", table, "app_name, user_id"), (f"idx_{table}_update_time", table, "update_time DESC"), @@ -662,7 +662,7 @@ def _events_table_ddl(table: str, session_table: str, json_column_type: str) -> """ -def _events_index_specs(table: str) -> list[tuple[str, str, str]]: +def _events_index_specs(table: str) -> "list[tuple[str, str, str]]": return [ (f"idx_{table}_scope", table, "app_name, user_id, session_id, timestamp ASC"), (f"idx_{table}_session", table, "session_id, timestamp ASC"), @@ -718,7 +718,7 @@ def _create_index_sql(table: str, index_name: str, columns: str) -> str: return f"CREATE INDEX {_quote_identifier(index_name)} ON {_table_ref(table)} ({columns})" -def _casefold_names(rows: list[Any], key: str) -> set[str]: +def _casefold_names(rows: "list[Any]", key: str) -> "set[str]": """Collapse data-dictionary rows into a case-folded, schema-stripped name set.""" return {str(row.get(key, "")).rsplit(".", 1)[-1].casefold() for row in rows} @@ -737,7 +737,7 @@ def _insert_event_sql(table: str) -> str: """ -def _upsert_state_sql(table: str, key_columns: tuple[str, ...], key_params: tuple[str, ...]) -> str: +def _upsert_state_sql(table: str, key_columns: "tuple[str, ...]", key_params: "tuple[str, ...]") -> str: source_columns = ", ".join( f"{param} AS {_quote_identifier(column)}" for column, param in zip(key_columns, key_params, strict=False) ) @@ -773,8 +773,8 @@ def _upsert_metadata_sql(table: str) -> str: def _events_query( - table: str, app_name: str, user_id: str, session_id: str, after_timestamp: datetime | None, limit: int | None -) -> tuple[str, tuple[Any, ...]]: + table: str, app_name: str, user_id: str, session_id: str, after_timestamp: "datetime | None", limit: "int | None" +) -> "tuple[str, tuple[Any, ...]]": top_clause = "TOP (%s) " if limit is not None else "" params: list[Any] = [limit] if limit is not None else [] params.extend([app_name, user_id, session_id]) @@ -791,7 +791,7 @@ def _events_query( return sql, tuple(params) -def _event_insert_params(event_record: StoredEvent) -> tuple[Any, ...]: +def _event_insert_params(event_record: StoredEvent) -> "tuple[Any, ...]": return ( event_record["id"], event_record["app_name"], @@ -821,7 +821,7 @@ def _event_record_from_row(row: Any) -> StoredEvent: ) -def _memory_record_from_row(row: Any) -> StoredMemory: +def _memory_record_from_row(row: Any) -> "StoredMemory": return cast( "StoredMemory", { @@ -842,7 +842,7 @@ def _memory_record_from_row(row: Any) -> StoredMemory: ) -def _json_dict(value: Any) -> dict[str, Any]: +def _json_dict(value: Any) -> "dict[str, Any]": if value is None: return {} if isinstance(value, dict): @@ -860,7 +860,7 @@ def _is_mssql_table_missing(exc: BaseException) -> bool: def _quote_identifier(identifier: str) -> str: - return quote_tsql_identifier(identifier) + return f"[{identifier.replace(']', ']]')}]" def _table_ref(table: str) -> str: @@ -891,8 +891,14 @@ def _build_mssql_scope_where( def _session_list_query( - session_table: str, app_name: str, user_id: str | None, column: str, direction: str, limit: int | None, offset: int -) -> tuple[str, tuple[Any, ...]]: + session_table: str, + app_name: str, + user_id: "str | None", + column: str, + direction: str, + limit: "int | None", + offset: int, +) -> "tuple[str, tuple[Any, ...]]": """Return the bounded session-list query and its bound values.""" params: list[Any] = [app_name] where_clause = "app_name = %s" diff --git a/sqlspec/adapters/pymssql/config.py b/sqlspec/adapters/pymssql/config.py index c35fa66ce..c00cb661b 100644 --- a/sqlspec/adapters/pymssql/config.py +++ b/sqlspec/adapters/pymssql/config.py @@ -1,7 +1,7 @@ """pymssql database configuration.""" from collections.abc import Callable, Mapping -from typing import Any, ClassVar, Literal, TypedDict, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast from typing_extensions import NotRequired @@ -11,12 +11,15 @@ from sqlspec.adapters.pymssql.migrations import PymssqlSyncMigrationTracker from sqlspec.adapters.pymssql.pool import PymssqlConnectionPool from sqlspec.config import ExtensionConfigs, SyncDatabaseConfig -from sqlspec.core import StatementConfig, TypeCoercionCapabilities +from sqlspec.core import TypeCoercionCapabilities from sqlspec.driver import SyncPoolConnectionContext, SyncPoolSessionFactory from sqlspec.extensions.events import EventRuntimeHints -from sqlspec.observability import ObservabilityConfig from sqlspec.utils.config_tools import normalize_connection_config +if TYPE_CHECKING: + from sqlspec.core import StatementConfig + from sqlspec.observability import ObservabilityConfig + __all__ = ("PymssqlConfig", "PymssqlConnectionParams", "PymssqlDriverFeatures", "PymssqlPoolParams", "PymssqlTimeout") PymssqlTimeout = int | float @@ -39,14 +42,13 @@ class PymssqlConnectionParams(TypedDict): conn_properties: NotRequired[str] autocommit: NotRequired[bool] tds_version: NotRequired[str] - encryption: NotRequired[Literal["default", "off", "request", "require"]] use_datetime2: NotRequired[bool] arraysize: NotRequired[int] conv: NotRequired[Mapping[int | type[Any], Callable[..., Any]]] read_only: NotRequired[bool] pool_recycle_seconds: NotRequired[int] health_check_interval: NotRequired[float] - extra: NotRequired[dict[str, Any]] + extra: NotRequired["dict[str, Any]"] class PymssqlPoolParams(PymssqlConnectionParams): @@ -67,9 +69,9 @@ class PymssqlDriverFeatures(TypedDict): events_backend: Event channel backend selection. """ - json_serializer: NotRequired[Callable[[Any], str]] - json_deserializer: NotRequired[Callable[[str], Any]] - on_connection_create: NotRequired[Callable[[PymssqlConnection], None]] + json_serializer: NotRequired["Callable[[Any], str]"] + json_deserializer: NotRequired["Callable[[str], Any]"] + on_connection_create: "NotRequired[Callable[[PymssqlConnection], None]]" enable_events: NotRequired[bool] events_backend: NotRequired[Literal["poll_queue"]] @@ -89,35 +91,35 @@ class PymssqlConfig(SyncDatabaseConfig[PymssqlConnection, PymssqlConnectionPool, __slots__ = ("_user_connection_hook",) - driver_type: ClassVar[type[PymssqlDriver]] = PymssqlDriver - connection_type: ClassVar[type[PymssqlConnection]] = cast("type[PymssqlConnection]", PymssqlConnection) - migration_tracker_type: ClassVar[type[PymssqlSyncMigrationTracker]] = PymssqlSyncMigrationTracker - supports_transactional_ddl: ClassVar[bool] = True - supports_migration_schemas: ClassVar[bool] = True - supports_native_arrow_export: ClassVar[bool] = False - supports_native_arrow_import: ClassVar[bool] = True - supports_native_parquet_export: ClassVar[bool] = False - supports_native_parquet_import: ClassVar[bool] = False - supports_native_row_streaming: ClassVar[bool] = True - type_coercion_capabilities: ClassVar[TypeCoercionCapabilities] = TypeCoercionCapabilities( + driver_type: "ClassVar[type[PymssqlDriver]]" = PymssqlDriver + connection_type: "ClassVar[type[PymssqlConnection]]" = cast("type[PymssqlConnection]", PymssqlConnection) + migration_tracker_type: "ClassVar[type[PymssqlSyncMigrationTracker]]" = PymssqlSyncMigrationTracker + supports_transactional_ddl: "ClassVar[bool]" = True + supports_migration_schemas: "ClassVar[bool]" = True + supports_native_arrow_export: "ClassVar[bool]" = False + supports_native_arrow_import: "ClassVar[bool]" = False + supports_native_parquet_export: "ClassVar[bool]" = False + supports_native_parquet_import: "ClassVar[bool]" = False + supports_native_row_streaming: "ClassVar[bool]" = True + type_coercion_capabilities: "ClassVar[TypeCoercionCapabilities]" = TypeCoercionCapabilities( datetime_binding="native", timestamp_precision="microsecond", json_columns_decoded=False, uuid_binding="text" ) - _connection_context_class: ClassVar[type[PymssqlConnectionContext]] = PymssqlConnectionContext - _session_factory_class: ClassVar[type[_PymssqlSessionConnectionHandler]] = _PymssqlSessionConnectionHandler - _session_context_class: ClassVar[type[PymssqlSessionContext]] = PymssqlSessionContext + _connection_context_class: "ClassVar[type[PymssqlConnectionContext]]" = PymssqlConnectionContext + _session_factory_class: "ClassVar[type[_PymssqlSessionConnectionHandler]]" = _PymssqlSessionConnectionHandler + _session_context_class: "ClassVar[type[PymssqlSessionContext]]" = PymssqlSessionContext _default_statement_config = default_statement_config def __init__( self, *, - connection_config: PymssqlPoolParams | dict[str, Any] | None = None, - connection_instance: PymssqlConnectionPool | None = None, - migration_config: dict[str, Any] | None = None, - statement_config: StatementConfig | None = None, - driver_features: PymssqlDriverFeatures | dict[str, Any] | None = None, - bind_key: str | None = None, - extension_config: ExtensionConfigs | None = None, - observability_config: ObservabilityConfig | None = None, + connection_config: "PymssqlPoolParams | dict[str, Any] | None" = None, + connection_instance: "PymssqlConnectionPool | None" = None, + migration_config: "dict[str, Any] | None" = None, + statement_config: "StatementConfig | None" = None, + driver_features: "PymssqlDriverFeatures | dict[str, Any] | None" = None, + bind_key: "str | None" = None, + extension_config: "ExtensionConfigs | None" = None, + observability_config: "ObservabilityConfig | None" = None, **kwargs: Any, ) -> None: connection_config = build_connection_config(normalize_connection_config(connection_config)) @@ -142,7 +144,7 @@ def __init__( **kwargs, ) - def _create_pool(self) -> PymssqlConnectionPool: + def _create_pool(self) -> "PymssqlConnectionPool": config = dict(self.connection_config) pool_recycle = config.pop("pool_recycle_seconds", 86400) health_check = config.pop("health_check_interval", 30.0) @@ -158,7 +160,7 @@ def _close_pool(self) -> None: self.connection_instance.close() self.connection_instance = None - def create_connection(self) -> PymssqlConnection: + def create_connection(self) -> "PymssqlConnection": """Open a standalone connection owned by the caller. The connection carries the same parameters and creation hook the pool @@ -170,7 +172,7 @@ def create_connection(self) -> PymssqlConnection: """ return self.provide_pool().new_connection() - def get_signature_namespace(self) -> dict[str, Any]: + def get_signature_namespace(self) -> "dict[str, Any]": namespace = super().get_signature_namespace() namespace.update({ "PymssqlConfig": PymssqlConfig, @@ -188,6 +190,6 @@ def get_signature_namespace(self) -> dict[str, Any]: }) return namespace - def get_event_runtime_hints(self) -> EventRuntimeHints: + def get_event_runtime_hints(self) -> "EventRuntimeHints": """Return runtime hints for pymssql event channels.""" return EventRuntimeHints(poll_interval=0.25, lease_seconds=5) diff --git a/sqlspec/adapters/pymssql/core.py b/sqlspec/adapters/pymssql/core.py index aae8a71b2..6be2f5796 100644 --- a/sqlspec/adapters/pymssql/core.py +++ b/sqlspec/adapters/pymssql/core.py @@ -1,11 +1,8 @@ """pymssql adapter compiled helpers.""" import re -from collections.abc import Callable, Mapping, Sequence, Sized -from logging import Logger -from typing import Any, Final, Literal - -from sqlglot import exp +from collections.abc import Callable, Sized +from typing import TYPE_CHECKING, Any, Final, Literal from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile from sqlspec.exceptions import ( @@ -27,11 +24,14 @@ from sqlspec.utils.type_converters import build_uuid_coercions from sqlspec.utils.type_guards import has_rowcount +if TYPE_CHECKING: + from collections.abc import Mapping, Sequence + from logging import Logger + __all__ = ( "apply_driver_features", "build_connection_config", "build_insert_statement", - "build_multi_row_insert", "build_profile", "build_statement_config", "collect_rows", @@ -40,10 +40,8 @@ "driver_profile", "extract_error_number", "format_identifier", - "is_plain_values_insert", "normalize_execute_many_parameters", "normalize_execute_parameters", - "quote_tsql_identifier", "resolve_column_names", "resolve_many_rowcount", "resolve_rowcount", @@ -66,7 +64,7 @@ } -def quote_tsql_identifier(identifier: str) -> str: +def _quote_bracket_identifier(identifier: str) -> str: """Quote a T-SQL identifier with square brackets.""" cleaned = identifier.strip() if cleaned.startswith("[") and cleaned.endswith("]"): @@ -81,61 +79,20 @@ def format_identifier(identifier: str) -> str: msg = "Table name must not be empty" raise SQLSpecError(msg) parts = split_qualified_identifier(cleaned, quote_chars='"', allow_bracket_quotes=True) - return ".".join(quote_tsql_identifier(part) for part in parts) + return ".".join(_quote_bracket_identifier(part) for part in parts) -def build_insert_statement(table: str, columns: list[str]) -> str: +def build_insert_statement(table: str, columns: "list[str]") -> str: """Build a pymssql-compatible INSERT statement.""" - column_clause = ", ".join(quote_tsql_identifier(column) for column in columns) + column_clause = ", ".join(_quote_bracket_identifier(column) for column in columns) placeholders = ", ".join("%s" for _ in columns) return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES ({placeholders})" -def build_multi_row_insert(table: str, columns: Sequence[str], num_rows: int, *, num_columns: int | None = None) -> str: - """Build a multi-row VALUES (...), (...) batch INSERT statement. - - Args: - table: Target table name. - columns: Column names to insert (empty when inserting into all table columns). - num_rows: Number of row tuples in the VALUES clause (up to 1,000). - num_columns: Explicit column count when ``columns`` is empty. - - Returns: - Parameterized T-SQL INSERT statement. - """ - col_count = len(columns) if columns else (num_columns or 0) - single_row = f"({', '.join('%s' for _ in range(col_count))})" - values_clause = ", ".join(single_row for _ in range(num_rows)) - if columns: - column_clause = ", ".join(quote_tsql_identifier(column) for column in columns) - return f"INSERT INTO {format_identifier(table)} ({column_clause}) VALUES {values_clause}" - return f"INSERT INTO {format_identifier(table)} VALUES {values_clause}" - - -def is_plain_values_insert(expression: Any, expected_columns: int) -> bool: - """Return whether a parsed INSERT expression is a single-row plain VALUES insert without OUTPUT/RETURNING.""" - if not isinstance(expression, exp.Insert): - return False - if expression.args.get("output") or expression.args.get("returning"): - return False - values = expression.expression - if not isinstance(values, exp.Values): - return False - rows = values.expressions - if len(rows) != 1: - return False - row = rows[0] - if not isinstance(row, exp.Tuple): - return False - return len(row.expressions) == expected_columns - - def normalize_execute_parameters(parameters: Any) -> Any: """Normalize parameters for pymssql execute calls.""" if parameters is None: return None - if isinstance(parameters, tuple): - return parameters if isinstance(parameters, list): return tuple(parameters) return parameters @@ -146,7 +103,7 @@ def normalize_execute_many_parameters(parameters: Any) -> Any: return parameters -def build_profile() -> DriverParameterProfile: +def build_profile() -> "DriverParameterProfile": """Create the pymssql driver parameter profile.""" return DriverParameterProfile( name="pymssql", @@ -166,8 +123,8 @@ def build_profile() -> DriverParameterProfile: def build_statement_config( - *, json_serializer: Callable[[Any], str] | None = None, json_deserializer: Callable[[str], Any] | None = None -) -> StatementConfig: + *, json_serializer: "Callable[[Any], str] | None" = None, json_deserializer: "Callable[[str], Any] | None" = None +) -> "StatementConfig": """Construct the pymssql statement configuration.""" return build_statement_config_from_profile( driver_profile, @@ -178,8 +135,8 @@ def build_statement_config( def apply_driver_features( - statement_config: StatementConfig, driver_features: Mapping[str, Any] | None -) -> tuple[StatementConfig, dict[str, Any]]: + statement_config: "StatementConfig", driver_features: "Mapping[str, Any] | None" +) -> "tuple[StatementConfig, dict[str, Any]]": """Apply pymssql driver feature defaults to statement config.""" features: dict[str, Any] = dict(driver_features) if driver_features else {} json_serializer = features.setdefault("json_serializer", to_json) @@ -194,7 +151,7 @@ def apply_driver_features( return statement_config, features -def create_mapped_exception(error: Exception, *, logger: Logger | None = None) -> SQLSpecError: +def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) -> SQLSpecError: """Map a pymssql exception to SQLSpec's exception hierarchy.""" error_number = extract_error_number(error) if error_number == _MSSQL_CONSTRAINT_547: @@ -227,8 +184,8 @@ def create_mapped_exception(error: Exception, *, logger: Logger | None = None) - def resolve_column_names( - description: Sequence[Any] | None, column_name_cache: dict[int, tuple[Any, list[str]]] | None = None -) -> list[str]: + description: "Sequence[Any] | None", column_name_cache: "dict[int, tuple[Any, list[str]]] | None" = None +) -> "list[str]": """Resolve ordered column names from cursor metadata.""" if not description: return [] @@ -246,10 +203,10 @@ def resolve_column_names( def collect_rows( - fetched_data: Sequence[Any] | None, - description: Sequence[Any] | None, - column_name_cache: dict[int, tuple[Any, list[str]]] | None = None, -) -> tuple[list[Any], list[str], Literal["dict", "tuple", "record"]]: + fetched_data: "Sequence[Any] | None", + description: "Sequence[Any] | None", + column_name_cache: "dict[int, tuple[Any, list[str]]] | None" = None, +) -> "tuple[list[Any], list[str], Literal['dict', 'tuple', 'record']]": """Collect pymssql rows, preserving dictionary or tuple row shape.""" column_names = resolve_column_names(description, column_name_cache) if not fetched_data: @@ -270,7 +227,7 @@ def resolve_rowcount(cursor: Any) -> int: return 0 -def resolve_many_rowcount(cursor: Any, parameters: Any, *, fallback_count: int | None = None) -> int: +def resolve_many_rowcount(cursor: Any, parameters: Any, *, fallback_count: "int | None" = None) -> int: """Resolve executemany rowcount using cursor metadata with payload fallback.""" rowcount = resolve_rowcount(cursor) if rowcount > 0: @@ -286,7 +243,7 @@ def _bool_to_int(value: bool) -> int: return int(value) -def _constraint_exception_from_message(error: Exception) -> SQLSpecError | None: +def _constraint_exception_from_message(error: Exception) -> "SQLSpecError | None": """Classify SQL Server constraint messages when a driver omits the native error number.""" message = str(error) normalized = message.lower() @@ -319,26 +276,22 @@ def extract_error_number(exc: BaseException | None) -> int | None: val = getattr(exc, attr, None) if isinstance(val, int) and not isinstance(val, bool) and val != 0: return val - if hasattr(exc, "args") and exc.args: + if exc.args: first = exc.args[0] if isinstance(first, int) and not isinstance(first, bool): return first matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) - if matches: - last_match = matches[-1] - raw_num = last_match[0] or last_match[1] if isinstance(last_match, tuple) else last_match - try: - return int(raw_num) - except ValueError: - pass - return None + if not matches: + return None + last_match = matches[-1] + return int(last_match[0] or last_match[1]) driver_profile = build_profile() default_statement_config = build_statement_config() -def build_connection_config(connection_config: Mapping[str, Any]) -> dict[str, Any]: +def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str, Any]": """Build a normalized connection configuration dictionary. Args: diff --git a/sqlspec/adapters/pymssql/data_dictionary.py b/sqlspec/adapters/pymssql/data_dictionary.py index b257e4f60..dc39b562e 100644 --- a/sqlspec/adapters/pymssql/data_dictionary.py +++ b/sqlspec/adapters/pymssql/data_dictionary.py @@ -14,21 +14,23 @@ SystemMetadataRequest, SystemMetadataResult, TableMetadata, + VersionInfo, ensure_system_metadata_request, get_data_dictionary_loader, get_dialect_config, system_metadata_gated_result, ) from sqlspec.data_dictionary.dialects.mssql import ( - MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, build_mssql_table_ddl_result, extract_mssql_version_value, get_mssql_data_dictionary_options, + is_mssql_azure_sql, list_mssql_available_features, merge_mssql_table_lists, + mssql_supports_native_json, parse_mssql_engine_edition, parse_mssql_version_components, resolve_mssql_feature_flag, @@ -42,13 +44,51 @@ from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.core import SQL - from sqlspec.data_dictionary import DialectConfig, MetadataCapabilityProfile + from sqlspec.data_dictionary._types import DialectConfig, MetadataCapabilityProfile __all__ = ("MssqlVersionInfo", "PymssqlSyncDataDictionary") logger = get_logger("sqlspec.adapters.pymssql.data_dictionary") +class MssqlVersionInfo(VersionInfo): + """MSSQL database version info with build, revision, and Azure SQL detection.""" + + def __init__( + self, + major: int, + minor: int = 0, + build: int = 0, + revision: int = 0, + edition: str | None = None, + engine_edition: int | None = None, + ) -> None: + super().__init__(major, minor, 0) + self.build = build + self.revision = revision + self.edition = edition + self.engine_edition = engine_edition + self.is_azure_sql = is_mssql_azure_sql(engine_edition) + + def supports_native_json(self) -> bool: + """Return whether this server supports the native JSON type.""" + return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) + + @property + def version_tuple(self) -> "tuple[int, int, int]": + """Get version tuple using the MSSQL build number as the third component.""" + return (self.major, self.minor, self.build) + + def __str__(self) -> str: + """String representation of version info.""" + version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" + if self.edition: + version_str += f" ({self.edition})" + if self.is_azure_sql: + version_str += " [Azure]" + return version_str + + class _MssqlDataDictionaryMixin: """Shared helpers for MSSQL data dictionaries.""" @@ -92,8 +132,6 @@ def _build_version_info( def _get_optimal_type_from_version(self, version_info: MssqlVersionInfo | None, type_category: str) -> str: if type_category in {"json", "jsonb"} and version_info is not None and version_info.supports_native_json(): return "JSON" - if type_category == "vector" and version_info is not None and version_info.supports_vector(): - return "VECTOR" return self.get_dialect_config().get_optimal_type(type_category) diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 00c57827b..0655b8052 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -1,12 +1,9 @@ """pymssql SQL Server driver implementation.""" import contextlib -from collections.abc import Iterable, Sequence, Sized +from collections.abc import Sized from typing import TYPE_CHECKING, Any, cast -import sqlglot -from sqlglot import exp - from sqlspec.adapters.pymssql._typing import ( PymssqlConnection, PymssqlCursor, @@ -15,22 +12,18 @@ PymssqlSessionContext, ) from sqlspec.adapters.pymssql.core import ( - build_multi_row_insert, collect_rows, create_mapped_exception, default_statement_config, driver_profile, - format_identifier, - is_plain_values_insert, normalize_execute_many_parameters, normalize_execute_parameters, - quote_tsql_identifier, resolve_column_names, resolve_many_rowcount, resolve_rowcount, ) from sqlspec.adapters.pymssql.data_dictionary import PymssqlSyncDataDictionary -from sqlspec.core import SQL, ArrowResult, StatementConfig, get_cache_config, register_driver_profile +from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile from sqlspec.driver import ( BaseSyncExceptionHandler, ExecutionResult, @@ -40,11 +33,12 @@ validate_savepoint_name, ) from sqlspec.exceptions import SQLSpecError -from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.utils.logging import get_logger if TYPE_CHECKING: - from pymssql._pymssql import QueryParams + from collections.abc import Sequence + + from sqlspec.adapters.pymssql._typing import PymssqlQueryParams as QueryParams __all__ = ("PymssqlCursor", "PymssqlDriver", "PymssqlExceptionHandler", "PymssqlSessionContext") @@ -56,7 +50,7 @@ class PymssqlExceptionHandler(BaseSyncExceptionHandler): __slots__ = () - def _handle_exception(self, exc_type: type[BaseException] | None, exc_val: BaseException) -> bool: + def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool: if exc_type is None: return False if isinstance(exc_val, PymssqlError): @@ -92,7 +86,7 @@ def start(self) -> None: raise self._cursor_manager = cursor_manager - def fetch_chunk(self) -> list[dict[str, Any]]: + def fetch_chunk(self) -> "list[dict[str, Any]]": cursor_manager = self._cursor_manager if cursor_manager is None or cursor_manager.cursor is None: return [] @@ -132,9 +126,9 @@ class PymssqlDriver(SyncDriverAdapterBase): def __init__( self, - connection: PymssqlConnection, - statement_config: StatementConfig | None = None, - driver_features: dict[str, Any] | None = None, + connection: "PymssqlConnection", + statement_config: "StatementConfig | None" = None, + driver_features: "dict[str, Any] | None" = None, ) -> None: if statement_config is None: statement_config = default_statement_config.replace( @@ -148,7 +142,7 @@ def __init__( self._transaction_active = False self._explicit_transaction = False - def dispatch_execute(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: + def dispatch_execute(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "ExecutionResult": sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) cursor.execute(sql, normalize_execute_parameters(prepared_parameters)) @@ -167,17 +161,8 @@ def dispatch_execute(self, cursor: PymssqlRawCursor, statement: SQL) -> Executio return self.create_execution_result(cursor, rowcount_override=resolve_rowcount(cursor)) - def dispatch_execute_many(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: - cached_statement, prepared_parameters = self._compiled_statement(statement, self.statement_config) - sql = cached_statement.compiled_sql - parsed_expression = cached_statement.expression - if parsed_expression is None and statement.raw_sql.lstrip().upper().startswith("INSERT"): - with contextlib.suppress(Exception): - parsed_expression = sqlglot.parse_one(statement.raw_sql, read="tsql") - if isinstance(parsed_expression, exp.Insert): - bulk_result = self._execute_bulk_insert_many(cursor, parsed_expression, prepared_parameters) - if bulk_result is not None: - return bulk_result + def dispatch_execute_many(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "ExecutionResult": + sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) prepared_parameters = normalize_execute_many_parameters(prepared_parameters) parameter_count = len(prepared_parameters) if isinstance(prepared_parameters, Sized) else None @@ -186,7 +171,7 @@ def dispatch_execute_many(self, cursor: PymssqlRawCursor, statement: SQL) -> Exe affected_rows = resolve_many_rowcount(cursor, prepared_parameters, fallback_count=parameter_count) return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True) - def dispatch_execute_script(self, cursor: PymssqlRawCursor, statement: SQL) -> ExecutionResult: + def dispatch_execute_script(self, cursor: "PymssqlRawCursor", statement: "SQL") -> "ExecutionResult": sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True) @@ -198,13 +183,6 @@ def dispatch_execute_script(self, cursor: PymssqlRawCursor, statement: SQL) -> E cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) - def collect_rows(self, cursor: PymssqlRawCursor, fetched: list[Any]) -> tuple[list[Any], list[str], int]: - rows, column_names, _ = collect_rows(fetched, cursor.description or None, self._column_name_cache) - return rows, column_names, len(rows) - - def resolve_rowcount(self, cursor: PymssqlRawCursor) -> int: - return resolve_rowcount(cursor) - def begin(self) -> None: """Begin a transaction on the connection. @@ -250,13 +228,13 @@ def rollback(self) -> None: msg = f"Failed to rollback SQL Server transaction: {exc}" raise SQLSpecError(msg) from exc - def with_cursor(self, connection: PymssqlConnection) -> PymssqlCursor: + def with_cursor(self, connection: "PymssqlConnection") -> "PymssqlCursor": return PymssqlCursor(connection) - def handle_database_exceptions(self) -> PymssqlExceptionHandler: + def handle_database_exceptions(self) -> "PymssqlExceptionHandler": return PymssqlExceptionHandler() - def dispatch_select_stream(self, statement: SQL, chunk_size: int) -> SyncRowStream[dict[str, Any]] | None: + def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "SyncRowStream[dict[str, Any]] | None": """Return a native pymssql row stream backed by ``fetchmany()``.""" if not statement.returns_rows(): return None @@ -303,171 +281,30 @@ def has_schema(self, schema: str) -> bool: return cursor.fetchone() is not None @property - def data_dictionary(self) -> PymssqlSyncDataDictionary: + def data_dictionary(self) -> "PymssqlSyncDataDictionary": if self._data_dictionary is None: self._data_dictionary = PymssqlSyncDataDictionary() return self._data_dictionary - def _execute_bulk_insert_many( - self, cursor: PymssqlRawCursor, expression: exp.Insert, prepared_parameters: Any - ) -> ExecutionResult | None: - """Execute a batch INSERT via multi-row VALUES chunking up to 1,000 rows.""" - if not isinstance(prepared_parameters, (list, tuple)) or not prepared_parameters: - return None - first_row = prepared_parameters[0] - if not isinstance(first_row, (list, tuple)) or not first_row: - return None - - target = expression.this - if isinstance(target, exp.Schema): - table_expr = target.this - column_names = [column.name for column in target.expressions] - elif isinstance(target, exp.Table): - table_expr = target - column_names = [] - else: - return None - - if not isinstance(table_expr, exp.Table) or table_expr.alias: - return None - - expected_columns = len(column_names) if column_names else len(first_row) - if expected_columns <= 0 or not is_plain_values_insert(expression, expected_columns): - return None - - target_table = table_expr.sql(dialect="tsql") - rows = prepared_parameters - total_affected = 0 - chunk_size = max(1, min(1000, 2000 // expected_columns)) + def collect_rows(self, cursor: "PymssqlRawCursor", fetched: "list[Any]") -> "tuple[list[Any], list[str], int]": + column_names = resolve_column_names(cursor.description or None, self._column_name_cache) + return fetched, column_names, len(fetched) - for i in range(0, len(rows), chunk_size): - chunk = rows[i : i + chunk_size] - chunk_sql = build_multi_row_insert(target_table, column_names, len(chunk), num_columns=expected_columns) - flat_params: list[Any] = [] - for row in chunk: - flat_params.extend(row) - cursor.execute(chunk_sql, tuple(flat_params)) - rowcount = resolve_rowcount(cursor) - total_affected += rowcount if rowcount > 0 else len(chunk) - - return self.create_execution_result(cursor, rowcount_override=total_affected, is_many_result=True) - - def _bulk_copy( - self, - table_name: str, - rows: Sequence[Sequence[Any]] | Iterable[Sequence[Any]], - *, - column_ids: Sequence[int] | None = None, - batch_size: int = 1000, - tablock: bool = False, - check_constraints: bool = False, - fire_triggers: bool = False, - ) -> int: - """Perform high-performance bulk insert using FreeTDS BCP APIs. - - Args: - table_name: Target SQL Server table name. - rows: Sequence or iterable of row tuples/sequences. - column_ids: Optional 1-based column IDs mapping elements to table columns. - batch_size: Number of rows per batch commit. Defaults to 1000. - tablock: Apply TABLOCK hint for minimal logging. - check_constraints: Enforce table constraints during BCP. - fire_triggers: Execute insert triggers during BCP. - - Returns: - Number of rows ingested. - """ - row_list = list(rows) if not isinstance(rows, (list, tuple)) else rows - if not row_list: - return 0 - formatted_table = format_identifier(table_name) - handler = self.handle_database_exceptions() - with handler: - self.connection.bulk_copy( - formatted_table, - row_list, - column_ids=list(column_ids) if column_ids is not None else None, - batch_size=batch_size, - tablock=tablock, - check_constraints=check_constraints, - fire_triggers=fire_triggers, - ) - self._check_pending_exception(handler) - return len(row_list) - - def load_from_arrow( - self, - table: str, - source: ArrowResult | Any, - *, - partitioner: dict[str, object] | None = None, - overwrite: bool = False, - telemetry: StorageTelemetry | None = None, - batch_size: int = 1000, - tablock: bool = False, - check_constraints: bool = False, - fire_triggers: bool = False, - column_ids: Sequence[int] | None = None, - ) -> StorageBridgeJob: - """Load Arrow data into SQL Server via FreeTDS BCP bulk copy.""" - self._require_capability("arrow_import_enabled") - if overwrite: - quoted_table = format_identifier(table) - handler = self.handle_database_exceptions() - with handler, self.with_cursor(self.connection) as cursor: - try: - cursor.execute(f"TRUNCATE TABLE {quoted_table}") - except Exception as exc: - error_msg = str(exc) - if "4712" in error_msg or "foreign key" in error_msg.lower(): - cursor.execute(f"DELETE FROM {quoted_table}") - else: - raise - self._check_pending_exception(handler) - - arrow_table = self._coerce_arrow_table(source) - if arrow_table.num_rows > 0: - for batch in arrow_table.to_batches(): - pydict = batch.to_pydict() - rows = list(zip(*pydict.values(), strict=False)) - self._bulk_copy( - table, - rows, - column_ids=column_ids, - batch_size=batch_size, - tablock=tablock, - check_constraints=check_constraints, - fire_triggers=fire_triggers, - ) - - telemetry_payload = self._ingest_telemetry(arrow_table) - extra = telemetry_payload.setdefault("extra", {}) - extra["rows_ingested"] = arrow_table.num_rows - telemetry_payload["rows_processed"] = arrow_table.num_rows - telemetry_payload["destination"] = table - self._attach_partition_telemetry(telemetry_payload, partitioner) - return self._storage_job(telemetry_payload, telemetry) - - def load_from_storage( - self, - table: str, - source: StorageDestination, - *, - file_format: StorageFormat, - partitioner: dict[str, object] | None = None, - overwrite: bool = False, - ) -> StorageBridgeJob: - """Load staged artifacts from storage into SQL Server via BCP.""" - arrow_table, inbound = self._read_storage_arrow(source, file_format=file_format) - return self.load_from_arrow(table, arrow_table, partitioner=partitioner, overwrite=overwrite, telemetry=inbound) + def resolve_rowcount(self, cursor: "PymssqlRawCursor") -> int: + return resolve_rowcount(cursor) def _connection_in_transaction(self) -> bool: """Return whether a transaction opened by this driver remains active.""" return self._transaction_active +def _quote_tsql_identifier(identifier: str) -> str: + """Bracket-quote an identifier so the statement is valid regardless of the session's QUOTED_IDENTIFIER setting.""" + return f"[{identifier.replace(']', ']]')}]" + + def _alter_default_schema_sql(user_name: str, schema: str) -> str: - return f"ALTER USER {quote_tsql_identifier(user_name)} WITH DEFAULT_SCHEMA = {quote_tsql_identifier(schema)};" + return f"ALTER USER {_quote_tsql_identifier(user_name)} WITH DEFAULT_SCHEMA = {_quote_tsql_identifier(schema)};" register_driver_profile("pymssql", driver_profile) diff --git a/sqlspec/adapters/pymssql/events/store.py b/sqlspec/adapters/pymssql/events/store.py index 9a40f6b09..89c17cdc8 100644 --- a/sqlspec/adapters/pymssql/events/store.py +++ b/sqlspec/adapters/pymssql/events/store.py @@ -3,7 +3,6 @@ import re from sqlspec.adapters.pymssql.config import PymssqlConfig -from sqlspec.adapters.pymssql.core import quote_tsql_identifier from sqlspec.extensions.events import BaseEventQueueStore from sqlspec.utils.text import split_qualified_identifier @@ -70,4 +69,8 @@ def _split_table_name(table_name: str) -> tuple[str, str]: def _object_name(table_name: str) -> str: schema_name, bare_table_name = _split_table_name(table_name) - return f"{quote_tsql_identifier(schema_name)}.{quote_tsql_identifier(bare_table_name)}" + return f"{_quote_bracket_identifier(schema_name)}.{_quote_bracket_identifier(bare_table_name)}" + + +def _quote_bracket_identifier(identifier: str) -> str: + return f"[{identifier.replace(']', ']]')}]" diff --git a/sqlspec/adapters/pymssql/litestar/store.py b/sqlspec/adapters/pymssql/litestar/store.py index 7fd9f1ec8..6661d8b2a 100644 --- a/sqlspec/adapters/pymssql/litestar/store.py +++ b/sqlspec/adapters/pymssql/litestar/store.py @@ -1,13 +1,15 @@ """pymssql Litestar Store implementation.""" from datetime import datetime, timedelta, timezone -from typing import Any +from typing import TYPE_CHECKING, Any from sqlspec.adapters.pymssql._typing import PymssqlCursor -from sqlspec.adapters.pymssql.config import PymssqlConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ +if TYPE_CHECKING: + from sqlspec.adapters.pymssql.config import PymssqlConfig + __all__ = ("PymssqlStore",) @@ -16,7 +18,7 @@ class PymssqlStore(BaseSQLSpecStore["PymssqlConfig"]): __slots__ = () - def __init__(self, config: PymssqlConfig) -> None: + def __init__(self, config: "PymssqlConfig") -> None: super().__init__(config) async def create_table(self) -> None: @@ -27,11 +29,11 @@ async def create_table(self) -> None: await async_(self._create_table)() await self.reconcile_schema(assume_existing=True) - async def get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: + async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": """Get a session value by key.""" return await async_(self._get)(key, renew_for) - async def set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: + async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: """Store a session value.""" await async_(self._set)(key, value, expires_in) @@ -47,7 +49,7 @@ async def exists(self, key: str) -> bool: """Check if a session key exists and is not expired.""" return await async_(self._exists)(key) - async def expires_in(self, key: str) -> int | None: + async def expires_in(self, key: str) -> "int | None": """Get the time in seconds until the session expires.""" return await async_(self._expires_in)(key) @@ -79,7 +81,7 @@ def _table_ddl(self) -> str: END; """ - def _drop_table_sql(self) -> list[str]: + def _drop_table_sql(self) -> "list[str]": """Get SQL Server DROP TABLE statements.""" return [f"IF OBJECT_ID(N'dbo.{self._table_name}', N'U') IS NOT NULL DROP TABLE dbo.{self._table_name};"] @@ -89,7 +91,7 @@ def _create_table(self) -> None: driver.commit() self._log_table_created() - def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | None: + def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = %s @@ -120,7 +122,7 @@ def _get(self, key: str, renew_for: int | timedelta | None = None) -> bytes | No return _coerce_bytes(_row_value(row, "data", 0)) - def _set(self, key: str, value: str | bytes, expires_in: int | timedelta | None = None) -> None: + def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) sql = f""" @@ -162,7 +164,7 @@ def _exists(self, key: str) -> bool: cursor.execute(sql, (key,)) return cursor.fetchone() is not None - def _expires_in(self, key: str) -> int | None: + def _expires_in(self, key: str) -> "int | None": with self._config.provide_connection() as conn, PymssqlCursor(conn) as cursor: cursor.execute(f"SELECT expires_at FROM {self._table_name} WHERE session_id = %s", (key,)) row = cursor.fetchone() @@ -204,7 +206,7 @@ def _row_value(row: object, key: str, index: int) -> Any: return getattr(row, key, None) -def _normalize_utc(value: Any) -> datetime | None: +def _normalize_utc(value: Any) -> "datetime | None": if value is None: return None if not isinstance(value, datetime): diff --git a/sqlspec/adapters/pymssql/pool.py b/sqlspec/adapters/pymssql/pool.py index 8c8ac29bc..a89b17e46 100644 --- a/sqlspec/adapters/pymssql/pool.py +++ b/sqlspec/adapters/pymssql/pool.py @@ -7,8 +7,7 @@ from contextlib import contextmanager from typing import TYPE_CHECKING, Any, cast -from sqlspec.adapters.pymssql._typing import PymssqlConnection -from sqlspec.adapters.pymssql._typing import pymssql_module as pymssql +from sqlspec.adapters.pymssql._typing import PYMSSQL_MODULE, PymssqlConnection from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -20,6 +19,7 @@ logger = get_logger(POOL_LOGGER_NAME) _ADAPTER_NAME = "pymssql" +pymssql = PYMSSQL_MODULE class PymssqlConnectionPool: diff --git a/sqlspec/data_dictionary/dialects/mssql/__init__.py b/sqlspec/data_dictionary/dialects/mssql/__init__.py index 85f31759a..43b098e76 100644 --- a/sqlspec/data_dictionary/dialects/mssql/__init__.py +++ b/sqlspec/data_dictionary/dialects/mssql/__init__.py @@ -4,7 +4,6 @@ MSSQL_CONFIG, MSSQL_PRODUCT_VERSION_PATTERN, MSSQL_VERSION_PATTERN, - MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, @@ -18,7 +17,6 @@ mssql_supports_json_functions, mssql_supports_native_json, mssql_supports_string_agg, - mssql_supports_vector, mssql_system_metadata_denied, parse_mssql_engine_edition, parse_mssql_version_components, @@ -30,7 +28,6 @@ "MSSQL_CONFIG", "MSSQL_PRODUCT_VERSION_PATTERN", "MSSQL_VERSION_PATTERN", - "MssqlVersionInfo", "build_mssql_metadata_capability_profile", "build_mssql_system_metadata_capability", "build_mssql_system_metadata_result", @@ -44,7 +41,6 @@ "mssql_supports_json_functions", "mssql_supports_native_json", "mssql_supports_string_agg", - "mssql_supports_vector", "mssql_system_metadata_denied", "parse_mssql_engine_edition", "parse_mssql_version_components", diff --git a/sqlspec/data_dictionary/dialects/mssql/config.py b/sqlspec/data_dictionary/dialects/mssql/config.py index 3b8ae1834..622e3dea7 100644 --- a/sqlspec/data_dictionary/dialects/mssql/config.py +++ b/sqlspec/data_dictionary/dialects/mssql/config.py @@ -18,16 +18,14 @@ SystemMetadataRedactionPolicy, SystemMetadataRequest, SystemMetadataResult, - VersionInfo, register_dialect, system_metadata_gated_result, ) if TYPE_CHECKING: - from sqlspec.data_dictionary import TableMetadata + from sqlspec.data_dictionary import TableMetadata, VersionInfo __all__ = ( - "MssqlVersionInfo", "build_mssql_metadata_capability_profile", "build_mssql_system_metadata_capability", "build_mssql_system_metadata_result", @@ -41,7 +39,6 @@ "mssql_supports_json_functions", "mssql_supports_native_json", "mssql_supports_string_agg", - "mssql_supports_vector", "mssql_system_metadata_denied", "parse_mssql_engine_edition", "parse_mssql_version_components", @@ -56,7 +53,6 @@ MSSQL_MIN_STRING_AGG_VERSION: Final[int] = 14 MSSQL_MIN_GREATEST_LEAST_VERSION: Final[int] = 16 MSSQL_MIN_NATIVE_JSON_VERSION: Final[int] = 17 -MSSQL_VECTOR_MIN_MAJOR: Final[int] = 17 MSSQL_ENGINE_EDITION_AZURE_SET: Final[frozenset[int]] = frozenset({5, 8, 11}) MSSQL_DYNAMIC_FEATURES: Final[tuple[str, ...]] = ( @@ -65,7 +61,6 @@ "supports_string_agg", "supports_greatest_least", "supports_native_json", - "supports_vector", ) MSSQL_REPLACEMENT_DOMAINS: Final[tuple[str, ...]] = ( @@ -138,7 +133,6 @@ "text": "NVARCHAR(MAX)", "json": "NVARCHAR(MAX)", "jsonb": "NVARCHAR(MAX)", - "vector": "VARBINARY(MAX)", "timestamp": "DATETIME2(6)", "timestamptz": "DATETIMEOFFSET(6)", "bytea": "VARBINARY(MAX)", @@ -162,48 +156,6 @@ register_dialect(MSSQL_CONFIG) -class MssqlVersionInfo(VersionInfo): - """MSSQL database version info with build, revision, and Azure SQL detection.""" - - def __init__( - self, - major: int, - minor: int = 0, - build: int = 0, - revision: int = 0, - edition: str | None = None, - engine_edition: int | None = None, - ) -> None: - super().__init__(major, minor, 0) - self.build = build - self.revision = revision - self.edition = edition - self.engine_edition = engine_edition - self.is_azure_sql = is_mssql_azure_sql(engine_edition) - - def supports_native_json(self) -> bool: - """Return whether this server supports the native JSON type.""" - return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) - - def supports_vector(self) -> bool: - """Return whether this server supports native VECTOR data types and functions.""" - return mssql_supports_vector(self.major, is_azure_sql=self.is_azure_sql) - - @property - def version_tuple(self) -> tuple[int, int, int]: - """Get version tuple using the MSSQL build number as the third component.""" - return (self.major, self.minor, self.build) - - def __str__(self) -> str: - """String representation of version info.""" - version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" - if self.edition: - version_str += f" ({self.edition})" - if self.is_azure_sql: - version_str += " [Azure]" - return version_str - - def extract_mssql_version_value(row: object) -> "str | None": """Extract a SQL Server version string from a row-like object.""" if isinstance(row, dict): @@ -270,11 +222,6 @@ def mssql_supports_native_json(major: int, is_azure_sql: bool = False) -> bool: return is_azure_sql or major >= MSSQL_MIN_NATIVE_JSON_VERSION -def mssql_supports_vector(major: int, is_azure_sql: bool = False) -> bool: - """Return whether the SQL Server version supports native VECTOR data types and functions.""" - return is_azure_sql or major >= MSSQL_VECTOR_MIN_MAJOR - - def resolve_mssql_feature_flag( feature: str, *, @@ -296,8 +243,6 @@ def resolve_mssql_feature_flag( return mssql_supports_greatest_least(major) if feature == "supports_native_json": return mssql_supports_native_json(major, is_azure_sql=is_azure_sql) - if feature == "supports_vector": - return mssql_supports_vector(major, is_azure_sql=is_azure_sql) dialect_config = config or MSSQL_CONFIG flag = dialect_config.get_feature_flag(feature) diff --git a/tests/unit/adapters/test_mssql_python/test_arrow.py b/tests/unit/adapters/test_mssql_python/test_arrow.py index 8ab601875..7a0f24a26 100644 --- a/tests/unit/adapters/test_mssql_python/test_arrow.py +++ b/tests/unit/adapters/test_mssql_python/test_arrow.py @@ -3,7 +3,6 @@ from collections.abc import Iterable from typing import TYPE_CHECKING, cast -import pyarrow as pa import pytest from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver @@ -37,9 +36,13 @@ def fetchmany(self, size: int) -> list[tuple[int, str]]: return chunk def arrow(self, batch_size: int = 8192) -> object: + import pyarrow as pa + return pa.table({"x": [1, 2, 3]}) def arrow_reader(self, batch_size: int = 8192) -> object: + import pyarrow as pa + table = pa.table({"x": [1, 2, 3]}) return pa.RecordBatchReader.from_batches(table.schema, table.to_batches(max_chunksize=batch_size)) @@ -155,7 +158,7 @@ def test_bulk_copy_forwards_options_to_cursor_bulkcopy() -> None: connection = ArrowConnection() driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) - result = driver._bulk_copy( + result = driver.bulk_copy( "dbo.target", [(1, "a"), (2, "b")], batch_size=1000, timeout=30, table_lock=True, keep_nulls=True ) @@ -175,7 +178,7 @@ def test_bulk_copy_defaults_match_mssql_python_runtime() -> None: connection = ArrowConnection() driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) - result = driver._bulk_copy("dbo.target", [(1,)]) + result = driver.bulk_copy("dbo.target", [(1,)]) _, _, options = connection.cursor_obj.bulkcopy_calls[0] assert result["rows_copied"] == 1 @@ -190,24 +193,6 @@ def test_bulk_copy_raises_mapped_driver_exception() -> None: driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) with pytest.raises(UniqueViolationError): - driver._bulk_copy("dbo.target", [(1,)]) + driver.bulk_copy("dbo.target", [(1,)]) assert connection.cursor_obj.closed is True - - -def test_load_from_arrow_falls_back_to_bulk_copy_when_bulkcopy_arrow_absent() -> None: - """load_from_arrow should fall back to _bulk_copy when cursor lacks bulkcopy_arrow.""" - connection = ArrowConnection() - driver = MssqlPythonDriver( - cast("MssqlPythonConnection", connection), - driver_features={"storage_capabilities": {"arrow_import_enabled": True}}, - ) - table = pa.table({"id": [1, 2], "name": ["Ada", "Grace"]}) - - job = driver.load_from_arrow("dbo.target", table, column_mappings=[]) - - assert job.telemetry["rows_processed"] == 2 - target_table, rows, options = connection.cursor_obj.bulkcopy_calls[0] - assert target_table == "dbo.target" - assert rows == [(1, "Ada"), (2, "Grace")] - assert options["column_mappings"] == [] diff --git a/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py b/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py index 511cef051..462f74757 100644 --- a/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py +++ b/tests/unit/adapters/test_mssql_python/test_bulk_copy_result.py @@ -36,7 +36,7 @@ def test_bulk_copy_defaults_match_upstream(driver_cls: Any) -> None: from mssql_python.cursor import Cursor upstream = inspect.signature(Cursor.bulkcopy).parameters - wrapper = inspect.signature(driver_cls._bulk_copy).parameters + wrapper = inspect.signature(driver_cls.bulk_copy).parameters for name in ( "batch_size", "timeout", diff --git a/tests/unit/adapters/test_mssql_python/test_config.py b/tests/unit/adapters/test_mssql_python/test_config.py index 9db77f7f7..4d35a8473 100644 --- a/tests/unit/adapters/test_mssql_python/test_config.py +++ b/tests/unit/adapters/test_mssql_python/test_config.py @@ -55,8 +55,8 @@ def fake_connect(connection_string: str, **kwargs: Any) -> DummyConnection: calls.append(("connect", (connection_string,), kwargs)) return connection - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", fake_pooling) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", fake_connect) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", fake_pooling) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", fake_connect) pool = MssqlPythonConnectionPool( connection_string="Server=localhost;", connect_kwargs={"timeout": 5}, max_size=7, idle_timeout=30, enabled=True @@ -79,7 +79,7 @@ def test_config_create_pool_splits_connection_and_pool_options(monkeypatch: pyte def fake_pooling(**kwargs: Any) -> None: pooling_calls.append(kwargs) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", fake_pooling) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", fake_pooling) config = MssqlPythonConfig( connection_config={ @@ -101,7 +101,7 @@ def fake_pooling(**kwargs: Any) -> None: def test_config_connection_string_with_discrete_override(monkeypatch: pytest.MonkeyPatch) -> None: """MssqlPythonConfig should merge discrete database overrides over connection_string.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **kw: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **kw: None) config = MssqlPythonConfig( connection_config={ @@ -161,7 +161,7 @@ def test_config_create_pool_normalizes_current_odbc_aliases(monkeypatch: pytest. def fake_pooling(**kwargs: Any) -> None: pooling_calls.append(kwargs) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", fake_pooling) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", fake_pooling) config = MssqlPythonConfig( connection_config={ @@ -224,9 +224,9 @@ def test_config_connection_hook_runs_for_session_connections(monkeypatch: pytest seen: list[DummyConnection] = [] monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **_: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) monkeypatch.setattr( - "sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", lambda *_args, **_kwargs: connection + "sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", lambda *_args, **_kwargs: connection ) config = MssqlPythonConfig( @@ -242,9 +242,9 @@ def test_config_connection_hook_runs_for_session_connections(monkeypatch: pytest def test_second_pool_warns_on_different_params(monkeypatch: pytest.MonkeyPatch) -> None: """A second pool with different process-wide pooling params emits one warning.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **_: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) monkeypatch.setattr( - "sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", lambda *_args, **_kwargs: DummyConnection() + "sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", lambda *_args, **_kwargs: DummyConnection() ) MssqlPythonConnectionPool(connection_string="Server=srv1;", max_size=10, idle_timeout=60, enabled=True) @@ -263,9 +263,9 @@ def test_second_pool_warns_on_different_params(monkeypatch: pytest.MonkeyPatch) def test_second_pool_same_params_no_warn(monkeypatch: pytest.MonkeyPatch) -> None: """A second pool with identical process-wide pooling params emits no warning.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **_: None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) monkeypatch.setattr( - "sqlspec.adapters.mssql_python.pool.mssql_python_module.connect", lambda *_args, **_kwargs: DummyConnection() + "sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.connect", lambda *_args, **_kwargs: DummyConnection() ) MssqlPythonConnectionPool(connection_string="Server=srv1;", max_size=10, idle_timeout=60, enabled=True) @@ -348,27 +348,14 @@ def commit(self) -> None: assert calls == ["rollback", "release"] -def test_pool_close_calls_ddbc_close_pooling(monkeypatch: pytest.MonkeyPatch) -> None: - """Connection pool close should call ddbc_bindings.close_pooling when requested.""" - closed_pooling: list[bool] = [] - - class FakeBindings: - @staticmethod - def close_pooling() -> None: - closed_pooling.append(True) - - monkeypatch.setattr(_mssql_pool.mssql_python_module, "ddbc_bindings", FakeBindings, raising=False) - pool = MssqlPythonConnectionPool(connection_string="Server=localhost;") - pool.close(close_driver_pooling=True) - assert closed_pooling == [True] - - -def test_pool_suppresses_warning_when_params_match(monkeypatch: pytest.MonkeyPatch) -> None: - """Pool reconfiguration should not warn if params are identical to previous.""" +def test_pool_does_not_reconfigure_when_params_match(monkeypatch: pytest.MonkeyPatch) -> None: + """Identical pool configuration should avoid touching the native process-wide pool.""" monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", (10, 60, True)) - monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.mssql_python_module.pooling", lambda **kw: None) + calls: list[dict[str, object]] = [] + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **kw: calls.append(kw)) with warnings.catch_warnings(record=True) as recorded: warnings.simplefilter("always") MssqlPythonConnectionPool(connection_string="Server=localhost;", max_size=10, idle_timeout=60, enabled=True) assert not any("Pooling configuration was already set" in str(w.message) for w in recorded) + assert calls == [] diff --git a/tests/unit/adapters/test_mssql_python/test_core.py b/tests/unit/adapters/test_mssql_python/test_core.py index ec8cdf955..d6a1a1f86 100644 --- a/tests/unit/adapters/test_mssql_python/test_core.py +++ b/tests/unit/adapters/test_mssql_python/test_core.py @@ -2,7 +2,7 @@ import pytest -from sqlspec.adapters.mssql_python._typing import mssql_python_module +from sqlspec.adapters.mssql_python._typing import MSSQL_PYTHON_MODULE from sqlspec.adapters.mssql_python.core import build_connection_config, create_mapped_exception, extract_error_number from sqlspec.exceptions import ( CheckViolationError, @@ -65,7 +65,7 @@ def test_build_connection_config_no_duplicate_pwd() -> None: def test_create_mapped_exception_extracts_sql_server_error_number() -> None: """SQL Server native error numbers should map to specific SQLSpec exceptions.""" - exc = mssql_python_module.IntegrityError( + exc = MSSQL_PYTHON_MODULE.IntegrityError( "23000", "[23000] [Microsoft][ODBC Driver 18 for SQL Server][SQL Server]Violation of UNIQUE KEY constraint. (2627)", ) @@ -78,7 +78,7 @@ def test_create_mapped_exception_extracts_sql_server_error_number() -> None: def test_create_mapped_exception_falls_back_for_connection_errors() -> None: """Known connection error numbers should map to DatabaseConnectionError.""" - exc = mssql_python_module.OperationalError( + exc = MSSQL_PYTHON_MODULE.OperationalError( "08001", "[08001] [Microsoft][ODBC Driver 18 for SQL Server]Named Pipes Provider: " "Could not open a connection to SQL Server (53)", @@ -140,7 +140,7 @@ def test_create_mapped_exception_classifies_constraint_messages_without_error_nu message: str, expected_type: type[Exception] ) -> None: """Constraint messages remain classifiable when the driver omits SQL Server error numbers.""" - mapped = create_mapped_exception(mssql_python_module.IntegrityError("23000", message)) + mapped = create_mapped_exception(MSSQL_PYTHON_MODULE.IntegrityError("23000", message)) assert isinstance(mapped, expected_type) diff --git a/tests/unit/adapters/test_mssql_python/test_data_dictionary.py b/tests/unit/adapters/test_mssql_python/test_data_dictionary.py index 7e00102a7..692eede3e 100644 --- a/tests/unit/adapters/test_mssql_python/test_data_dictionary.py +++ b/tests/unit/adapters/test_mssql_python/test_data_dictionary.py @@ -138,28 +138,3 @@ def test_sync_data_dictionary_explicit_schema_skips_connection_lookup() -> None: data_dictionary.get_tables(cast(Any, driver), schema="custom") assert driver.select_calls[0][1]["schema_name"] == "custom" assert len(driver.executed) == 0 - - -def test_mssql_version_info_supports_vector() -> None: - """Version 17+ or Azure SQL engine editions support vectors.""" - v16 = MssqlVersionInfo(16, 0, 0, engine_edition=3) - v17 = MssqlVersionInfo(17, 0, 0, engine_edition=3) - azure = MssqlVersionInfo(16, 0, 0, engine_edition=5) - - assert v16.supports_vector() is False - assert v17.supports_vector() is True - assert azure.supports_vector() is True - - -def test_data_dictionary_vector_feature_flag_and_optimal_type() -> None: - """Sync data dictionary resolves supports_vector and optimal type for vector.""" - - class VectorDriver: - def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: - return {"product_version": "17.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} - - data_dictionary = MssqlPythonSyncDataDictionary() - driver = VectorDriver() - - assert data_dictionary.get_feature_flag(cast(Any, driver), "supports_vector") is True - assert data_dictionary.get_optimal_type(cast(Any, driver), "vector") == "VECTOR" diff --git a/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py b/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py index 1d5b7f89d..29cf4ef3c 100644 --- a/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py +++ b/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py @@ -78,13 +78,13 @@ def test_sync_load_from_arrow_skips_an_empty_table() -> None: assert conn._cursor.bulkcopy_calls == [] -def test_sync_load_from_arrow_overwrite_truncates_first() -> None: +def test_sync_load_from_arrow_overwrite_deletes_first() -> None: conn = _FakeConnection() driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) driver.load_from_arrow("dbo.orders", pa.table({"id": [1]}), overwrite=True) - assert conn._cursor.execute_calls == ["TRUNCATE TABLE [dbo].[orders]"] + assert conn._cursor.execute_calls == ["DELETE FROM [dbo].[orders]"] assert conn._cursor.arrow_calls @@ -94,44 +94,5 @@ def test_sync_load_from_arrow_overwrite_preserves_quoted_dots() -> None: driver.load_from_arrow('"dbo.schema"."orders.table"', pa.table({"id": [1]}), overwrite=True) - assert conn._cursor.execute_calls == ["TRUNCATE TABLE [dbo.schema].[orders.table]"] + assert conn._cursor.execute_calls == ["DELETE FROM [dbo.schema].[orders.table]"] assert conn._cursor.arrow_calls - - -def test_sync_load_from_arrow_overwrite_falls_back_to_delete_on_fk_reference() -> None: - conn = _FakeConnection() - driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) - - class FkError(Exception): - number = 4712 - - original_execute = conn._cursor.execute - - def execute_with_fk(sql: str, *args: Any) -> None: - original_execute(sql, *args) - if sql.startswith("TRUNCATE"): - raise FkError("Cannot truncate table referenced by foreign key") - - conn._cursor.execute = cast("Any", execute_with_fk) - driver.load_from_arrow("dbo.orders", pa.table({"id": [1]}), overwrite=True) - - assert conn._cursor.execute_calls == ["TRUNCATE TABLE [dbo].[orders]", "DELETE FROM [dbo].[orders]"] - assert conn._cursor.arrow_calls - - -def test_sync_load_from_arrow_forwards_bulk_copy_options() -> None: - conn = _FakeConnection() - driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) - table = pa.table({"id": [1, 2], "name": ["a", "b"]}) - - job = driver.load_from_arrow( - "orders", table, batch_size=500, check_constraints=True, fire_triggers=True, keep_nulls=True, table_lock=True - ) - - assert job.telemetry["rows_processed"] == 2 - _, _, kwargs = conn._cursor.arrow_calls[0] - assert kwargs["batch_size"] == 500 - assert kwargs["check_constraints"] is True - assert kwargs["fire_triggers"] is True - assert kwargs["keep_nulls"] is True - assert kwargs["table_lock"] is True diff --git a/tests/unit/adapters/test_mssql_python/test_type_converter.py b/tests/unit/adapters/test_mssql_python/test_type_converter.py index 22540d70a..83b1dd3c3 100644 --- a/tests/unit/adapters/test_mssql_python/test_type_converter.py +++ b/tests/unit/adapters/test_mssql_python/test_type_converter.py @@ -64,9 +64,3 @@ def test_tsql_time_maps_to_arrow_time64() -> None: def test_tsql_timestamp_remains_the_rowversion_binary_type() -> None: """TIMESTAMP is a T-SQL rowversion alias and must stay binary.""" assert mssql_type_to_arrow("timestamp") == pa.binary() - - -def test_mssql_type_to_arrow_maps_json_and_vector() -> None: - """JSON and VECTOR types should map to expected Arrow types.""" - assert mssql_type_to_arrow("json") == pa.string() - assert mssql_type_to_arrow("vector") == pa.list_(pa.float32()) diff --git a/tests/unit/adapters/test_pymssql/_fakes.py b/tests/unit/adapters/test_pymssql/_fakes.py index 6f1e0969e..8fcd4ec1d 100644 --- a/tests/unit/adapters/test_pymssql/_fakes.py +++ b/tests/unit/adapters/test_pymssql/_fakes.py @@ -61,7 +61,6 @@ def __init__(self, cursor: "FakeCursor | None" = None) -> None: self.rollbacks = 0 self.autocommit_values: list[bool] = [] self.autocommit_state = True - self.bulk_copy_calls: list[dict[str, Any]] = [] def cursor(self, *args: Any, **kwargs: Any) -> FakeCursor: self.cursor_args = args @@ -83,26 +82,6 @@ def autocommit(self, value: bool) -> None: self.autocommit_values.append(value) self.autocommit_state = value - def bulk_copy( - self, - table_name: str, - elements: Any, - column_ids: Any = None, - batch_size: int = 1000, - tablock: bool = False, - check_constraints: bool = False, - fire_triggers: bool = False, - ) -> None: - self.bulk_copy_calls.append({ - "table_name": table_name, - "elements": list(elements), - "column_ids": column_ids, - "batch_size": batch_size, - "tablock": tablock, - "check_constraints": check_constraints, - "fire_triggers": fire_triggers, - }) - class FakePymssqlModule: """Patch target that behaves like the pymssql module surface used by SQLSpec.""" diff --git a/tests/unit/adapters/test_pymssql/test_config.py b/tests/unit/adapters/test_pymssql/test_config.py index 53d5ee1c1..7ddf6e783 100644 --- a/tests/unit/adapters/test_pymssql/test_config.py +++ b/tests/unit/adapters/test_pymssql/test_config.py @@ -4,9 +4,7 @@ import pytest -from sqlspec.adapters.pymssql import PymssqlConnection as PublicPymssqlConnection from sqlspec.adapters.pymssql import build_connection_config -from sqlspec.adapters.pymssql._typing import PymssqlConnection, PymssqlRawCursor, pymssql_module from sqlspec.adapters.pymssql.config import PymssqlConfig, PymssqlConnectionParams from sqlspec.adapters.pymssql.driver import PymssqlDriver from sqlspec.adapters.pymssql.pool import PymssqlConnectionPool @@ -33,7 +31,6 @@ def test_connection_params_cover_common_pymssql_keywords() -> None: "tds_version", "pool_recycle_seconds", "health_check_interval", - "encryption", } assert expected_keys <= set(annotations) @@ -48,7 +45,6 @@ def test_config_defaults_server_port_and_features() -> None: assert config.driver_type is PymssqlDriver assert config.supports_transactional_ddl is True assert config.supports_native_arrow_export is False - assert config.supports_native_arrow_import is True assert config.driver_features["enable_events"] is True @@ -123,10 +119,12 @@ def test_signature_namespace_exposes_public_adapter_types() -> None: def test_pymssql_runtime_aliases_resolve_to_installed_classes() -> None: """pymssql public runtime aliases should expose installed pymssql classes.""" pymssql = pytest.importorskip("pymssql") + from sqlspec.adapters.pymssql import PymssqlConnection as PublicPymssqlConnection + from sqlspec.adapters.pymssql._typing import PYMSSQL_MODULE, PymssqlConnection, PymssqlRawCursor namespace = PymssqlConfig().get_signature_namespace() - assert pymssql_module is pymssql + assert PYMSSQL_MODULE is pymssql assert PymssqlConnection is pymssql.Connection assert PublicPymssqlConnection is pymssql.Connection assert PymssqlRawCursor is pymssql.Cursor diff --git a/tests/unit/adapters/test_pymssql/test_core.py b/tests/unit/adapters/test_pymssql/test_core.py index d7aaedd94..ee1e947a5 100644 --- a/tests/unit/adapters/test_pymssql/test_core.py +++ b/tests/unit/adapters/test_pymssql/test_core.py @@ -6,7 +6,6 @@ from sqlspec.adapters.pymssql.core import ( build_insert_statement, - build_multi_row_insert, collect_rows, create_mapped_exception, default_statement_config, @@ -15,7 +14,6 @@ format_identifier, normalize_execute_many_parameters, normalize_execute_parameters, - quote_tsql_identifier, ) from sqlspec.core import SQL, ParameterStyle from sqlspec.exceptions import ( @@ -149,14 +147,6 @@ def test_normalize_execute_many_parameters_passes_through() -> None: assert normalize_execute_many_parameters(rows) is rows -def test_quote_tsql_identifier() -> None: - """quote_tsql_identifier wraps identifiers in brackets and escapes closing brackets.""" - assert quote_tsql_identifier("users") == "[users]" - assert quote_tsql_identifier("[users]") == "[users]" - assert quote_tsql_identifier("dbo.users") == "[dbo.users]" - assert quote_tsql_identifier("col]name") == "[col]]name]" - - def test_extract_error_number() -> None: """extract_error_number detects error number from attribute, tuple, or regex.""" @@ -174,12 +164,6 @@ class BoolAttributeException(Exception): assert extract_error_number(Exception("Plain error")) is None -def test_build_multi_row_insert() -> None: - """build_multi_row_insert generates a multi-row VALUES INSERT statement.""" - sql = build_multi_row_insert("dbo.users", ["id", "name"], 3) - assert sql == "INSERT INTO [dbo].[users] ([id], [name]) VALUES (%s, %s), (%s, %s), (%s, %s)" - - def test_collect_rows_preserves_list_identity() -> None: """collect_rows avoids copying when the input rows are already a list.""" input_rows = [(1, "Alice"), (2, "Bob")] diff --git a/tests/unit/adapters/test_pymssql/test_data_dictionary.py b/tests/unit/adapters/test_pymssql/test_data_dictionary.py index 8b9e08956..b9c4aafe2 100644 --- a/tests/unit/adapters/test_pymssql/test_data_dictionary.py +++ b/tests/unit/adapters/test_pymssql/test_data_dictionary.py @@ -2,9 +2,7 @@ from typing import Any, cast -from sqlspec.adapters.mssql_python.data_dictionary import MssqlVersionInfo as MssqlPythonVersionInfo from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo, PymssqlSyncDataDictionary -from sqlspec.data_dictionary.dialects.mssql import MssqlVersionInfo as DialectMssqlVersionInfo class FakeSyncDriver: @@ -131,38 +129,3 @@ def test_sync_data_dictionary_explicit_schema_skips_connection_lookup() -> None: data_dictionary.get_tables(cast(Any, driver), schema="custom") assert driver.select_calls[0][1]["schema_name"] == "custom" assert len(driver.executed) == 0 - - -def test_mssql_version_info_supports_vector() -> None: - """Version 17+ or Azure SQL engine editions support vectors.""" - v16 = MssqlVersionInfo(16, 0, 0, engine_edition=3) - v17 = MssqlVersionInfo(17, 0, 0, engine_edition=3) - azure = MssqlVersionInfo(16, 0, 0, engine_edition=5) - - assert v16.supports_vector() is False - assert v17.supports_vector() is True - assert azure.supports_vector() is True - - -def test_data_dictionary_vector_feature_flag_and_optimal_type() -> None: - """Sync data dictionary resolves supports_vector and optimal type for vector.""" - assert MssqlVersionInfo is MssqlPythonVersionInfo - assert MssqlVersionInfo is DialectMssqlVersionInfo - - class VectorDriver: - def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: - return {"product_version": "17.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} - - class NonVectorDriver: - def select_one_or_none(self, _statement: Any, **_kwargs: Any) -> dict[str, Any]: - return {"product_version": "16.0.1000.1", "edition": "Enterprise Edition", "engine_edition": 3} - - data_dictionary = PymssqlSyncDataDictionary() - driver = VectorDriver() - old_driver = NonVectorDriver() - - assert "supports_vector" in data_dictionary.list_available_features() - assert data_dictionary.get_feature_flag(cast(Any, driver), "supports_vector") is True - assert data_dictionary.get_optimal_type(cast(Any, driver), "vector") == "VECTOR" - assert data_dictionary.get_feature_flag(cast(Any, old_driver), "supports_vector") is False - assert data_dictionary.get_optimal_type(cast(Any, old_driver), "vector") == "VARBINARY(MAX)" diff --git a/tests/unit/adapters/test_pymssql/test_driver.py b/tests/unit/adapters/test_pymssql/test_driver.py index 4cf1f848d..7a75e550c 100644 --- a/tests/unit/adapters/test_pymssql/test_driver.py +++ b/tests/unit/adapters/test_pymssql/test_driver.py @@ -1,15 +1,12 @@ """pymssql driver tests.""" -from typing import Any, cast +from typing import cast -import pyarrow as pa import pytest from pymssql import IntegrityError as PymssqlIntegrityError from sqlspec import StatementStack from sqlspec.adapters.pymssql._typing import PymssqlConnection, PymssqlRawCursor -from sqlspec.adapters.pymssql.core import default_statement_config -from sqlspec.adapters.pymssql.driver import PymssqlDriver, PymssqlExceptionHandler from sqlspec.core import SQL from sqlspec.exceptions import SQLSpecError, StackExecutionError, TransactionError, UniqueViolationError from tests.unit.adapters.test_pymssql._fakes import FakeConnection, FakeCursor @@ -32,6 +29,8 @@ def test_execute_maps_pymssql_row_formats( rows: list[tuple[int, str] | dict[str, int | str]], expected: list[dict[str, int | str]] ) -> None: + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor(rows=rows, description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor))) @@ -45,6 +44,8 @@ def test_execute_maps_pymssql_row_formats( @pytest.mark.parametrize("bad_name", UNSAFE_SAVEPOINT_NAMES) def test_pymssql_savepoint_overrides_reject_unsafe_names(bad_name: str) -> None: """The T-SQL savepoint overrides must reject unsafe identifiers before interpolation.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection())) with pytest.raises(TransactionError): @@ -57,6 +58,8 @@ def test_pymssql_savepoint_overrides_reject_unsafe_names(bad_name: str) -> None: def test_pymssql_savepoint_overrides_accept_valid_name() -> None: """A safe savepoint name should pass validation and reach the underlying execute path.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor() connection = FakeConnection(cursor) driver = PymssqlDriver(cast("PymssqlConnection", connection)) @@ -71,6 +74,9 @@ def test_pymssql_savepoint_overrides_accept_valid_name() -> None: def test_dispatch_execute_select_compiles_to_pyformat_and_collects_rows() -> None: """SELECT dispatch should execute pyformat SQL and return fetched rows.""" + from sqlspec.adapters.pymssql.core import default_statement_config + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor(rows=[(1, "Ada")], description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor)), statement_config=default_statement_config) statement = SQL("SELECT id, name FROM dbo.users WHERE id = ?", 1, statement_config=default_statement_config) @@ -84,19 +90,19 @@ def test_dispatch_execute_select_compiles_to_pyformat_and_collects_rows() -> Non def test_dispatch_execute_many_uses_executemany_and_rowcount() -> None: - """execute_many dispatch should forward non-plain-INSERT batch parameters to pymssql executemany.""" + """execute_many dispatch should forward batch parameters to pymssql.""" + from sqlspec.adapters.pymssql.core import default_statement_config + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor(rowcount=2) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor)), statement_config=default_statement_config) statement = SQL( - "UPDATE dbo.users SET name = ? WHERE id = ?", - [("Ada", 1), ("Grace", 2)], - statement_config=default_statement_config, - is_many=True, + "INSERT INTO dbo.users (id) VALUES (?)", [(1,), (2,)], statement_config=default_statement_config, is_many=True ) result = driver.dispatch_execute_many(cast("PymssqlRawCursor", cursor), statement) - assert cursor.many_calls == [("UPDATE dbo.users SET name = %s WHERE id = %s", [("Ada", 1), ("Grace", 2)])] + assert cursor.many_calls == [("INSERT INTO dbo.users (id) VALUES (%s)", [(1,), (2,)])] assert result.rowcount_override == 2 assert result.is_many_result is True @@ -104,6 +110,8 @@ def test_dispatch_execute_many_uses_executemany_and_rowcount() -> None: @pytest.mark.parametrize("finish", ["commit", "rollback"]) def test_autocommit_transaction_is_ended_with_tsql(finish: str) -> None: """pymssql ignores commit() and rollback() under autocommit, so the driver ends its own transaction.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor() connection = FakeConnection(cursor) driver = PymssqlDriver(cast("PymssqlConnection", connection)) @@ -119,6 +127,8 @@ def test_autocommit_transaction_is_ended_with_tsql(finish: str) -> None: @pytest.mark.parametrize("finish", ["commit", "rollback"]) def test_non_autocommit_transaction_uses_connection_boundaries(finish: str) -> None: """Without autocommit, pymssql's connection commit() and rollback() end the open transaction.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor() connection = FakeConnection(cursor) connection.autocommit(False) @@ -133,6 +143,8 @@ def test_non_autocommit_transaction_uses_connection_boundaries(finish: str) -> N def test_begin_reuses_the_open_transaction_without_autocommit() -> None: """A connection with autocommit disabled already holds a transaction, so begin issues no SQL.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor() connection = FakeConnection(cursor) connection.autocommit(False) @@ -148,6 +160,8 @@ def test_begin_reuses_the_open_transaction_without_autocommit() -> None: def test_exception_handler_maps_pymssql_errors() -> None: """pymssql exception handlers should surface mapped SQLSpec exceptions.""" + from sqlspec.adapters.pymssql.driver import PymssqlExceptionHandler + handler = PymssqlExceptionHandler() handled = handler._handle_exception( @@ -160,6 +174,7 @@ def test_exception_handler_maps_pymssql_errors() -> None: def test_commit_wraps_driver_errors() -> None: """Commit failures should be wrapped in SQLSpecError.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver class FailingConnection(FakeConnection): def commit(self) -> None: @@ -173,6 +188,8 @@ def commit(self) -> None: def test_collect_rows_returns_column_names() -> None: """The direct row collection hook should match SyncDriverAdapterBase expectations.""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor(description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor))) @@ -185,6 +202,8 @@ def test_collect_rows_returns_column_names() -> None: def test_select_stream_uses_fetchmany_chunks() -> None: """The pymssql driver should stream rows with cursor.fetchmany().""" + from sqlspec.adapters.pymssql.driver import PymssqlDriver + cursor = FakeCursor(rows=[(1, "Ada"), (2, "Grace"), (3, "Linus")], description=[("id",), ("name",)]) driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection(cursor))) @@ -199,6 +218,8 @@ def test_select_stream_uses_fetchmany_chunks() -> None: @pytest.mark.parametrize("finish", ["commit", "rollback"]) def test_connection_in_transaction_tracks_successful_boundaries(finish: str) -> None: + from sqlspec.adapters.pymssql.driver import PymssqlDriver + driver = PymssqlDriver(cast("PymssqlConnection", FakeConnection())) assert driver._connection_in_transaction() is False driver.begin() @@ -209,6 +230,8 @@ def test_connection_in_transaction_tracks_successful_boundaries(finish: str) -> @pytest.mark.parametrize("operation", ["begin", "commit", "rollback"]) def test_failed_transaction_boundary_preserves_state(operation: str, monkeypatch: pytest.MonkeyPatch) -> None: + from sqlspec.adapters.pymssql.driver import PymssqlDriver + connection = FakeConnection() driver = PymssqlDriver(cast("PymssqlConnection", connection)) if operation != "begin": @@ -227,6 +250,8 @@ def fail(*_args: object) -> None: @pytest.mark.parametrize("fails", [False, True]) def test_execute_stack_preserves_caller_transaction(fails: bool, monkeypatch: pytest.MonkeyPatch) -> None: + from sqlspec.adapters.pymssql.driver import PymssqlDriver + connection = FakeConnection(FakeCursor(rowcount=1)) driver = PymssqlDriver(cast("PymssqlConnection", connection)) driver.begin() @@ -253,120 +278,3 @@ def fail(sql: str, parameters: object = None) -> None: assert sum(sql == "BEGIN TRANSACTION" for sql, _ in connection.cursor_obj.calls) == 1 driver.rollback() assert connection.cursor_obj.calls[-1] == ("IF @@TRANCOUNT > 0 ROLLBACK TRANSACTION", None) - - -def test_driver_bulk_copy_forwards_options() -> None: - """_bulk_copy forwards batch options to underlying connection.""" - connection = FakeConnection() - driver = PymssqlDriver(cast("PymssqlConnection", connection)) - - result = driver._bulk_copy( - "dbo.users", - [(1, "Ada"), (2, "Grace")], - column_ids=[1, 2], - batch_size=500, - tablock=True, - check_constraints=True, - fire_triggers=True, - ) - - assert result == 2 - assert len(connection.bulk_copy_calls) == 1 - call = connection.bulk_copy_calls[0] - assert call["table_name"] == "[dbo].[users]" - assert call["elements"] == [(1, "Ada"), (2, "Grace")] - assert call["column_ids"] == [1, 2] - assert call["batch_size"] == 500 - assert call["tablock"] is True - assert call["check_constraints"] is True - assert call["fire_triggers"] is True - - -def test_load_from_arrow_bulk_copies_batches() -> None: - """load_from_arrow processes Arrow table in batches via _bulk_copy.""" - connection = FakeConnection() - driver = PymssqlDriver( - cast("PymssqlConnection", connection), driver_features={"storage_capabilities": {"arrow_import_enabled": True}} - ) - table = pa.table({"id": [1, 2], "name": ["Ada", "Grace"]}) - - job = driver.load_from_arrow("dbo.users", table, batch_size=500) - - assert job.telemetry["rows_processed"] == 2 - assert len(connection.bulk_copy_calls) == 1 - - -def test_load_from_arrow_overwrite_truncates_first() -> None: - """load_from_arrow with overwrite=True executes TRUNCATE TABLE.""" - cursor = FakeCursor() - connection = FakeConnection(cursor) - driver = PymssqlDriver( - cast("PymssqlConnection", connection), driver_features={"storage_capabilities": {"arrow_import_enabled": True}} - ) - table = pa.table({"id": [1], "name": ["Ada"]}) - - driver.load_from_arrow("dbo.users", table, overwrite=True) - - executed_sqls = [call[0] for call in cursor.calls] - assert "TRUNCATE TABLE [dbo].[users]" in executed_sqls - - -def test_load_from_arrow_overwrite_falls_back_on_fk_error() -> None: - """load_from_arrow falls back to DELETE FROM when error 4712 is encountered.""" - - class FkError(Exception): - number = 4712 - - cursor = FakeCursor() - connection = FakeConnection(cursor) - driver = PymssqlDriver( - cast("PymssqlConnection", connection), driver_features={"storage_capabilities": {"arrow_import_enabled": True}} - ) - - def execute_with_fk(sql: str, *args: Any) -> None: - cursor.calls.append((sql, args)) - if sql.startswith("TRUNCATE"): - raise FkError("Cannot truncate table referenced by foreign key") - - cursor.execute = cast("Any", execute_with_fk) - table = pa.table({"id": [1], "name": ["Ada"]}) - - driver.load_from_arrow("dbo.users", table, overwrite=True) - - executed_sqls = [call[0] for call in cursor.calls] - assert "TRUNCATE TABLE [dbo].[users]" in executed_sqls - assert "DELETE FROM [dbo].[users]" in executed_sqls - - -def test_execute_many_plain_values_chunks_into_multi_row_insert() -> None: - """execute_many with plain VALUES uses multi-row INSERT for both str and SQL objects.""" - cursor = FakeCursor(rowcount=3) - connection = FakeConnection(cursor) - driver = PymssqlDriver(cast("PymssqlConnection", connection)) - - params = [(1, "Ada"), (2, "Grace"), (3, "Linus")] - result = driver.execute_many("INSERT INTO dbo.users (id, name) VALUES (?, ?)", params) - - assert result.rows_affected == 3 - executed_sqls = [call[0] for call in cursor.calls] - assert len(executed_sqls) == 1 - assert "VALUES (%s, %s), (%s, %s), (%s, %s)" in executed_sqls[0] - - cursor.calls.clear() - sql_obj_result = driver.execute_many(SQL("INSERT INTO dbo.users VALUES (?, ?)"), params) - assert sql_obj_result.rows_affected == 3 - assert len(cursor.calls) == 1 - assert "INSERT INTO [dbo].[users] VALUES (%s, %s), (%s, %s), (%s, %s)" in cursor.calls[0][0] - - -def test_execute_many_non_plain_values_uses_standard_executemany() -> None: - """execute_many with non-plain SQL uses cursor.executemany.""" - cursor = FakeCursor(rowcount=2) - connection = FakeConnection(cursor) - driver = PymssqlDriver(cast("PymssqlConnection", connection)) - - params = [("Ada", 1), ("Grace", 2)] - result = driver.execute_many("UPDATE dbo.users SET name = ? WHERE id = ?", params) - - assert result.rows_affected == 2 - assert len(cursor.many_calls) == 1 From bc10c5771e061f48f06e2d9b10da497980b4c809 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:43:59 +0000 Subject: [PATCH 08/11] fix(mssql): retain native Arrow bulk loading and execution options --- docs/changelog.rst | 7 +- docs/reference/adapters/mssql_python.rst | 12 ++ sqlspec/adapters/mssql_python/core.py | 12 +- .../adapters/mssql_python/data_dictionary.py | 42 +----- sqlspec/adapters/mssql_python/driver.py | 127 +++++++++++++++--- sqlspec/adapters/pymssql/config.py | 1 + sqlspec/adapters/pymssql/data_dictionary.py | 42 +----- .../dialects/mssql/__init__.py | 2 + .../data_dictionary/dialects/mssql/config.py | 42 +++++- .../test_mssql_python/test_load_from_arrow.py | 29 +++- 10 files changed, 203 insertions(+), 113 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index b3a3a64ac..ef3ed8a53 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -1,13 +1,10 @@ -========= Changelog -========= All notable SQLSpec changes are summarized here. Entries are grouped by release and focus on user-visible behavior, public API changes, compatibility notes, and important operational fixes. Recent Updates -============== Unreleased ---------- @@ -17,6 +14,9 @@ Unreleased * BigQuery supports native query resource controls, explicit STRUCT parameters, typed empty arrays, and configurable Storage Write stream modes while retaining the atomic PENDING default. +* SQL Server Arrow loading accepts native record batch readers and BulkCopy + options while retaining name-based mappings and DELETE overwrite behavior. +* Pymssql connection typing includes native encryption settings. * Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown row count since the native driver does not return the number of changed rows. @@ -2277,7 +2277,6 @@ v0.24.0 - Builder consolidation * Refactored builder code to reduce duplication. Previous Versions -================= For releases before ``v0.24.0``, see the repository tag history and GitHub release records. diff --git a/docs/reference/adapters/mssql_python.rst b/docs/reference/adapters/mssql_python.rst index 822c85c30..6751a0ef6 100644 --- a/docs/reference/adapters/mssql_python.rst +++ b/docs/reference/adapters/mssql_python.rst @@ -94,3 +94,15 @@ Use these types inside ``extension_config["adk"]``. .. autoclass:: sqlspec.adapters.mssql_python.adk.MssqlPythonADKConfig :members: :show-inheritance: + +Native Arrow loading +-------------------- + +``load_from_arrow()`` accepts tables, record batches, record batch readers, and +Arrow C stream sources. It forwards ``batch_size``, ``timeout``, ``table_lock``, +``check_constraints``, ``fire_triggers``, ``keep_identity``, ``keep_nulls``, +``use_internal_transaction``, and ``column_mappings`` to native BulkCopy. +Field names supply default mappings; sources without schema metadata require +explicit mappings. Native internal transactions apply per batch, not to the +caller connection transaction. ``overwrite=True`` retains DELETE semantics; +stream consumption failures can leave a partial load. diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index 6aa702a78..dda0049d8 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -177,15 +177,17 @@ def create_mapped_exception(error: Exception, *, logger: "Logger | None" = None) def materialize_tuple_rows(fetched: "Sequence[Any] | None") -> "list[tuple[Any, ...]]": - """Materialize mssql-python ``Row`` objects into plain tuples. + """Materialize mssql-python Row objects into plain tuples. - ``mssql-python`` returns ``mssql_python.Row`` objects that are iterable and - indexable but are not ``tuple`` subclasses. The driver reports - ``row_format="tuple"``, so fetched rows are converted to real tuples to keep - that contract accurate when results are materialized. + Uses native tuple storage when available, avoiding a tuple copy per row. """ if not fetched: return [] + first = fetched[0] + if isinstance(first, tuple): + return list(fetched) if not isinstance(fetched, list) else fetched + if hasattr(first, "_values"): + return [tuple(row._values) if not isinstance(row._values, tuple) else row._values for row in fetched] return [tuple(row) for row in fetched] diff --git a/sqlspec/adapters/mssql_python/data_dictionary.py b/sqlspec/adapters/mssql_python/data_dictionary.py index 8e2e13d6d..74f60cb21 100644 --- a/sqlspec/adapters/mssql_python/data_dictionary.py +++ b/sqlspec/adapters/mssql_python/data_dictionary.py @@ -14,23 +14,21 @@ SystemMetadataRequest, SystemMetadataResult, TableMetadata, - VersionInfo, ensure_system_metadata_request, get_data_dictionary_loader, get_dialect_config, system_metadata_gated_result, ) from sqlspec.data_dictionary.dialects.mssql import ( + MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, build_mssql_table_ddl_result, extract_mssql_version_value, get_mssql_data_dictionary_options, - is_mssql_azure_sql, list_mssql_available_features, merge_mssql_table_lists, - mssql_supports_native_json, parse_mssql_engine_edition, parse_mssql_version_components, resolve_mssql_feature_flag, @@ -51,44 +49,6 @@ logger = get_logger("sqlspec.adapters.mssql_python.data_dictionary") -class MssqlVersionInfo(VersionInfo): - """MSSQL database version info with build, revision, and Azure SQL detection.""" - - def __init__( - self, - major: int, - minor: int = 0, - build: int = 0, - revision: int = 0, - edition: str | None = None, - engine_edition: int | None = None, - ) -> None: - super().__init__(major, minor, 0) - self.build = build - self.revision = revision - self.edition = edition - self.engine_edition = engine_edition - self.is_azure_sql = is_mssql_azure_sql(engine_edition) - - def supports_native_json(self) -> bool: - """Return whether this server supports the native JSON type.""" - return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) - - @property - def version_tuple(self) -> "tuple[int, int, int]": - """Get version tuple using the MSSQL build number as the third component.""" - return (self.major, self.minor, self.build) - - def __str__(self) -> str: - """String representation of version info.""" - version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" - if self.edition: - version_str += f" ({self.edition})" - if self.is_azure_sql: - version_str += " [Azure]" - return version_str - - class _MssqlDataDictionaryMixin: """Shared helpers for MSSQL data dictionaries.""" diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index 5eddec33a..b2a270f0b 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -198,7 +198,7 @@ def dispatch_execute_script(self, cursor: "MssqlPythonRawCursor", statement: "SQ statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True) successful_count = 0 for stmt in statements: - _execute_cursor(cursor, stmt, prepared_parameters) + _execute_cursor(cursor, stmt, prepared_parameters, use_prepare=False) successful_count += 1 return self.create_execution_result( cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True @@ -272,10 +272,12 @@ def set_migration_session_schema(self, schema: str) -> None: _execute_cursor(cursor, "SELECT USER_NAME() AS user_name, SCHEMA_NAME() AS schema_name;", None) row: Any = cursor.fetchone() user_name, current_schema = row[0], row[1] - _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None) + _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None, use_prepare=False) self._migration_schema_restore = (str(user_name), str(current_schema)) return - _execute_cursor(cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None) + _execute_cursor( + cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None, use_prepare=False + ) def reset_migration_session_schema(self) -> None: """Restore the user's default schema captured by set_migration_session_schema and commit it.""" @@ -283,7 +285,7 @@ def reset_migration_session_schema(self) -> None: return user_name, previous_schema = self._migration_schema_restore with self.with_cursor(self.connection) as cursor: - _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None) + _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None, use_prepare=False) self.connection.commit() self._migration_schema_restore = None @@ -419,21 +421,98 @@ def load_from_arrow( partitioner: "dict[str, object] | None" = None, overwrite: bool = False, telemetry: "StorageTelemetry | None" = None, + batch_size: int = 0, + timeout: int = 30, + table_lock: bool = False, + check_constraints: bool = False, + fire_triggers: bool = False, + keep_identity: bool = False, + keep_nulls: bool = False, + use_internal_transaction: bool = False, + column_mappings: "list[str] | list[tuple[int, str]] | None" = None, ) -> "StorageBridgeJob": - """Load Arrow data into SQL Server via BulkCopy.""" + """Load Arrow tables or streams using native BulkCopy options. + + Stream sources without schema metadata require explicit column mappings. + Overwrite deletes existing rows after validating the source shape; errors + while consuming a stream can still leave a partially completed load. + """ self._require_capability("arrow_import_enabled") - arrow_table = self._coerce_arrow_table(source) - if overwrite: - exc_handler = self.handle_database_exceptions() - with exc_handler, self.with_cursor(self.connection) as cursor: + ensure_pyarrow() + import pyarrow as pa + + if batch_size < 0 or timeout < 0: + msg = "batch_size and timeout must be non-negative" + raise ValueError(msg) + is_stream = not isinstance(source, pa.Table) and ( + isinstance(source, (pa.RecordBatchReader, pa.RecordBatch)) or hasattr(source, "__arrow_c_stream__") + ) + if is_stream: + arrow_source = source + has_rows = True + schema = getattr(source, "schema", None) + schema_names = getattr(schema, "names", None) + if column_mappings is None and schema_names is None: + msg = "Arrow stream sources without schema metadata require column_mappings" + raise ValueError(msg) + columns = column_mappings if column_mappings is not None else list(cast("Iterable[str]", schema_names)) + telemetry_payload = cast("StorageTelemetry", {"format": "arrow", "extra": {}}) + else: + arrow_source = self._coerce_arrow_table(source) + has_rows = bool(arrow_source.num_rows) + columns = column_mappings if column_mappings is not None else list(arrow_source.column_names) + telemetry_payload = self._ingest_telemetry(arrow_source) + + options = { + "batch_size": batch_size, + "timeout": timeout, + "table_lock": table_lock, + "check_constraints": check_constraints, + "fire_triggers": fire_triggers, + "keep_identity": keep_identity, + "keep_nulls": keep_nulls, + "use_internal_transaction": use_internal_transaction, + "column_mappings": columns, + } + raw_result: Any = None + use_fallback = False + exc_handler = self.handle_database_exceptions() + with exc_handler, self.with_cursor(self.connection) as cursor: + native_bulkcopy = getattr(cursor, "bulkcopy_arrow", None) + if is_stream and not callable(native_bulkcopy): + msg = "This mssql-python version does not support native Arrow stream bulk copy" + raise SQLSpecError(msg) + if overwrite: cursor.execute(f"DELETE FROM {_quote_mssql_table(table)}") - self._check_pending_exception(exc_handler) - if arrow_table.num_rows: - exc_handler = self.handle_database_exceptions() - with exc_handler, self.with_cursor(self.connection) as cursor: - cursor.bulkcopy_arrow(table, arrow_table, column_mappings=list(arrow_table.column_names)) - self._check_pending_exception(exc_handler) - telemetry_payload = self._ingest_telemetry(arrow_table) + if has_rows: + if callable(native_bulkcopy): + raw_result = native_bulkcopy(table, arrow_source, **options) + else: + use_fallback = True + self._check_pending_exception(exc_handler) + if use_fallback: + _, records = self._arrow_table_to_rows(cast("Any", arrow_source)) + raw_result = self.bulk_copy( + table, + records, + batch_size=batch_size, + timeout=timeout, + table_lock=table_lock, + check_constraints=check_constraints, + fire_triggers=fire_triggers, + keep_identity=keep_identity, + keep_nulls=keep_nulls, + use_internal_transaction=use_internal_transaction, + column_mappings=columns, + ) + if isinstance(raw_result, dict): + extra = telemetry_payload.setdefault("extra", {}) + if "rows_copied" in raw_result: + telemetry_payload["rows_processed"] = raw_result["rows_copied"] + extra["rows_ingested"] = raw_result["rows_copied"] + for key in ("elapsed_time", "rows_per_second", "batch_count"): + if key in raw_result: + extra[key] = raw_result[key] telemetry_payload["destination"] = table self._attach_partition_telemetry(telemetry_payload, partitioner) return self._storage_job(telemetry_payload, telemetry) @@ -476,11 +555,19 @@ def _quote_mssql_table(table: str) -> str: return ".".join(_quote_tsql_identifier(part) for part in split_qualified_identifier(table)) -def _execute_cursor(cursor: "MssqlPythonRawCursor", sql: str, parameters: Any) -> None: - if parameters is None: +def _execute_cursor(cursor: "MssqlPythonRawCursor", sql: str, parameters: Any, *, use_prepare: bool = True) -> None: + if use_prepare or parameters: + if parameters is None: + cursor.execute(sql) + else: + cursor.execute(sql, parameters) + return + try: + cursor.execute(sql, use_prepare=False) + except TypeError as exc: + if "use_prepare" not in str(exc): + raise cursor.execute(sql) - else: - cursor.execute(sql, parameters) def _cursor_rowcount(cursor: "MssqlPythonRawCursor") -> int: diff --git a/sqlspec/adapters/pymssql/config.py b/sqlspec/adapters/pymssql/config.py index c00cb661b..4a6ecab4d 100644 --- a/sqlspec/adapters/pymssql/config.py +++ b/sqlspec/adapters/pymssql/config.py @@ -42,6 +42,7 @@ class PymssqlConnectionParams(TypedDict): conn_properties: NotRequired[str] autocommit: NotRequired[bool] tds_version: NotRequired[str] + encryption: NotRequired[Literal["off", "request", "require"]] use_datetime2: NotRequired[bool] arraysize: NotRequired[int] conv: NotRequired[Mapping[int | type[Any], Callable[..., Any]]] diff --git a/sqlspec/adapters/pymssql/data_dictionary.py b/sqlspec/adapters/pymssql/data_dictionary.py index dc39b562e..e719f7784 100644 --- a/sqlspec/adapters/pymssql/data_dictionary.py +++ b/sqlspec/adapters/pymssql/data_dictionary.py @@ -14,23 +14,21 @@ SystemMetadataRequest, SystemMetadataResult, TableMetadata, - VersionInfo, ensure_system_metadata_request, get_data_dictionary_loader, get_dialect_config, system_metadata_gated_result, ) from sqlspec.data_dictionary.dialects.mssql import ( + MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, build_mssql_table_ddl_result, extract_mssql_version_value, get_mssql_data_dictionary_options, - is_mssql_azure_sql, list_mssql_available_features, merge_mssql_table_lists, - mssql_supports_native_json, parse_mssql_engine_edition, parse_mssql_version_components, resolve_mssql_feature_flag, @@ -51,44 +49,6 @@ logger = get_logger("sqlspec.adapters.pymssql.data_dictionary") -class MssqlVersionInfo(VersionInfo): - """MSSQL database version info with build, revision, and Azure SQL detection.""" - - def __init__( - self, - major: int, - minor: int = 0, - build: int = 0, - revision: int = 0, - edition: str | None = None, - engine_edition: int | None = None, - ) -> None: - super().__init__(major, minor, 0) - self.build = build - self.revision = revision - self.edition = edition - self.engine_edition = engine_edition - self.is_azure_sql = is_mssql_azure_sql(engine_edition) - - def supports_native_json(self) -> bool: - """Return whether this server supports the native JSON type.""" - return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) - - @property - def version_tuple(self) -> "tuple[int, int, int]": - """Get version tuple using the MSSQL build number as the third component.""" - return (self.major, self.minor, self.build) - - def __str__(self) -> str: - """String representation of version info.""" - version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" - if self.edition: - version_str += f" ({self.edition})" - if self.is_azure_sql: - version_str += " [Azure]" - return version_str - - class _MssqlDataDictionaryMixin: """Shared helpers for MSSQL data dictionaries.""" diff --git a/sqlspec/data_dictionary/dialects/mssql/__init__.py b/sqlspec/data_dictionary/dialects/mssql/__init__.py index 43b098e76..4fa25962f 100644 --- a/sqlspec/data_dictionary/dialects/mssql/__init__.py +++ b/sqlspec/data_dictionary/dialects/mssql/__init__.py @@ -4,6 +4,7 @@ MSSQL_CONFIG, MSSQL_PRODUCT_VERSION_PATTERN, MSSQL_VERSION_PATTERN, + MssqlVersionInfo, build_mssql_metadata_capability_profile, build_mssql_system_metadata_capability, build_mssql_system_metadata_result, @@ -28,6 +29,7 @@ "MSSQL_CONFIG", "MSSQL_PRODUCT_VERSION_PATTERN", "MSSQL_VERSION_PATTERN", + "MssqlVersionInfo", "build_mssql_metadata_capability_profile", "build_mssql_system_metadata_capability", "build_mssql_system_metadata_result", diff --git a/sqlspec/data_dictionary/dialects/mssql/config.py b/sqlspec/data_dictionary/dialects/mssql/config.py index 622e3dea7..ad910a403 100644 --- a/sqlspec/data_dictionary/dialects/mssql/config.py +++ b/sqlspec/data_dictionary/dialects/mssql/config.py @@ -18,14 +18,16 @@ SystemMetadataRedactionPolicy, SystemMetadataRequest, SystemMetadataResult, + VersionInfo, register_dialect, system_metadata_gated_result, ) if TYPE_CHECKING: - from sqlspec.data_dictionary import TableMetadata, VersionInfo + from sqlspec.data_dictionary import TableMetadata __all__ = ( + "MssqlVersionInfo", "build_mssql_metadata_capability_profile", "build_mssql_system_metadata_capability", "build_mssql_system_metadata_result", @@ -156,6 +158,44 @@ register_dialect(MSSQL_CONFIG) +class MssqlVersionInfo(VersionInfo): + """MSSQL database version info with build, revision, and Azure SQL detection.""" + + def __init__( + self, + major: int, + minor: int = 0, + build: int = 0, + revision: int = 0, + edition: str | None = None, + engine_edition: int | None = None, + ) -> None: + super().__init__(major, minor, 0) + self.build = build + self.revision = revision + self.edition = edition + self.engine_edition = engine_edition + self.is_azure_sql = is_mssql_azure_sql(engine_edition) + + def supports_native_json(self) -> bool: + """Return whether this server supports the native JSON type.""" + return mssql_supports_native_json(self.major, is_azure_sql=self.is_azure_sql) + + @property + def version_tuple(self) -> "tuple[int, int, int]": + """Get version tuple using the MSSQL build number as the third component.""" + return (self.major, self.minor, self.build) + + def __str__(self) -> str: + """String representation of version info.""" + version_str = f"{self.major}.{self.minor}.{self.build}.{self.revision}" + if self.edition: + version_str += f" ({self.edition})" + if self.is_azure_sql: + version_str += " [Azure]" + return version_str + + def extract_mssql_version_value(row: object) -> "str | None": """Extract a SQL Server version string from a row-like object.""" if isinstance(row, dict): diff --git a/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py b/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py index 29cf4ef3c..37a1293f5 100644 --- a/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py +++ b/tests/unit/adapters/test_mssql_python/test_load_from_arrow.py @@ -3,6 +3,7 @@ from typing import Any, cast import pyarrow as pa +import pytest from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver @@ -29,7 +30,9 @@ def bulkcopy(self, target_table: str, rows: Any, **kwargs: Any) -> dict[str, Any def bulkcopy_arrow(self, table_name: str, source: Any, **kwargs: Any) -> dict[str, Any]: self.arrow_calls.append((table_name, source, kwargs)) - return {"rows_copied": source.num_rows} + return { + "rows_copied": source.read_all().num_rows if isinstance(source, pa.RecordBatchReader) else source.num_rows + } def execute(self, sql: str, *_args: Any) -> None: self.execute_calls.append(sql) @@ -96,3 +99,27 @@ def test_sync_load_from_arrow_overwrite_preserves_quoted_dots() -> None: assert conn._cursor.execute_calls == ["DELETE FROM [dbo.schema].[orders.table]"] assert conn._cursor.arrow_calls + + +def test_arrow_stream_preserves_name_mapping_and_bulk_options() -> None: + conn = _FakeConnection() + driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) + reader = pa.table({"second": [2], "first": [1]}).to_reader() + job = driver.load_from_arrow("orders", reader, batch_size=32, timeout=4, keep_identity=True, table_lock=True) + target, source, options = conn._cursor.arrow_calls[0] + assert target == "orders" + assert source is reader + assert options["column_mappings"] == ["second", "first"] + assert options["batch_size"] == 32 + assert options["timeout"] == 4 + assert options["keep_identity"] is True + assert options["table_lock"] is True + assert job.telemetry["rows_processed"] == 1 + + +def test_arrow_overwrite_validates_source_before_delete() -> None: + conn = _FakeConnection() + driver = MssqlPythonDriver(cast("Any", conn), driver_features={"storage_capabilities": _CAPS}) + with pytest.raises((TypeError, ValueError)): + driver.load_from_arrow("orders", object(), overwrite=True) + assert conn._cursor.execute_calls == [] From 7a7f3a0b3aa8c3772de0b7f601f9db6beb67d562 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:05:47 +0000 Subject: [PATCH 09/11] docs: reconcile unreleased adapter changelog after rebase --- docs/changelog.rst | 26 ++++++++++++++++---------- 1 file changed, 16 insertions(+), 10 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index ef3ed8a53..5caedf472 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -1,10 +1,13 @@ +========= Changelog +========= All notable SQLSpec changes are summarized here. Entries are grouped by release and focus on user-visible behavior, public API changes, compatibility notes, and important operational fixes. Recent Updates +============== Unreleased ---------- @@ -14,9 +17,9 @@ Unreleased * BigQuery supports native query resource controls, explicit STRUCT parameters, typed empty arrays, and configurable Storage Write stream modes while retaining the atomic PENDING default. -* SQL Server Arrow loading accepts native record batch readers and BulkCopy - options while retaining name-based mappings and DELETE overwrite behavior. -* Pymssql connection typing includes native encryption settings. +* The mssql-python adapter can load Arrow streams with native BulkCopy options. Columns + map by name by default, and overwrite still uses DELETE. +* Pymssql connection types include native encryption settings. * Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown row count since the native driver does not return the number of changed rows. @@ -45,6 +48,15 @@ Unreleased **Fixed:** +* The mssql-python adapter runs scripts and schema changes without a prepare + step when no values are bound. Arrow read failures use SQLSpec error types. +* ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve + question marks in quoted identifiers, literals, and comments. + +* Arrow ODBC pagination reuses compiled placeholder positions instead of + parsing SQL again. ADBC keeps bound values in its ADK store queries. + DuckDB Arrow loads keep sparse dictionary fields and quote table names. + * 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. @@ -83,13 +95,6 @@ Unreleased ``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. - -* Arrow ODBC pagination reuses compiled placeholder positions instead of - 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. @@ -2277,6 +2282,7 @@ v0.24.0 - Builder consolidation * Refactored builder code to reduce duplication. Previous Versions +================= For releases before ``v0.24.0``, see the repository tag history and GitHub release records. From aae91d2b2c620396dd0a73cd3579e0c86685364a Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:20:25 +0000 Subject: [PATCH 10/11] refactor(mssql): rely on native direct execution defaults --- docs/changelog.rst | 3 +-- sqlspec/adapters/mssql_python/driver.py | 26 ++++++++----------------- 2 files changed, 9 insertions(+), 20 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 5caedf472..51588b38a 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -48,8 +48,7 @@ Unreleased **Fixed:** -* The mssql-python adapter runs scripts and schema changes without a prepare - step when no values are bound. Arrow read failures use SQLSpec error types. +* Arrow read failures in the mssql-python adapter use SQLSpec error types. * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index b2a270f0b..c35982958 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -198,7 +198,7 @@ def dispatch_execute_script(self, cursor: "MssqlPythonRawCursor", statement: "SQ statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True) successful_count = 0 for stmt in statements: - _execute_cursor(cursor, stmt, prepared_parameters, use_prepare=False) + _execute_cursor(cursor, stmt, prepared_parameters) successful_count += 1 return self.create_execution_result( cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True @@ -272,12 +272,10 @@ def set_migration_session_schema(self, schema: str) -> None: _execute_cursor(cursor, "SELECT USER_NAME() AS user_name, SCHEMA_NAME() AS schema_name;", None) row: Any = cursor.fetchone() user_name, current_schema = row[0], row[1] - _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None, use_prepare=False) + _execute_cursor(cursor, _alter_default_schema_sql(str(user_name), schema), None) self._migration_schema_restore = (str(user_name), str(current_schema)) return - _execute_cursor( - cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None, use_prepare=False - ) + _execute_cursor(cursor, _alter_default_schema_sql(self._migration_schema_restore[0], schema), None) def reset_migration_session_schema(self) -> None: """Restore the user's default schema captured by set_migration_session_schema and commit it.""" @@ -285,7 +283,7 @@ def reset_migration_session_schema(self) -> None: return user_name, previous_schema = self._migration_schema_restore with self.with_cursor(self.connection) as cursor: - _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None, use_prepare=False) + _execute_cursor(cursor, _alter_default_schema_sql(user_name, previous_schema), None) self.connection.commit() self._migration_schema_restore = None @@ -555,19 +553,11 @@ def _quote_mssql_table(table: str) -> str: return ".".join(_quote_tsql_identifier(part) for part in split_qualified_identifier(table)) -def _execute_cursor(cursor: "MssqlPythonRawCursor", sql: str, parameters: Any, *, use_prepare: bool = True) -> None: - if use_prepare or parameters: - if parameters is None: - cursor.execute(sql) - else: - cursor.execute(sql, parameters) - return - try: - cursor.execute(sql, use_prepare=False) - except TypeError as exc: - if "use_prepare" not in str(exc): - raise +def _execute_cursor(cursor: "MssqlPythonRawCursor", sql: str, parameters: Any) -> None: + if parameters is None: cursor.execute(sql) + else: + cursor.execute(sql, parameters) def _cursor_rowcount(cursor: "MssqlPythonRawCursor") -> int: From bd5ce5bb860d08674109084cf4394f62d30ffda4 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 28 Sep 2026 00:16:11 +0000 Subject: [PATCH 11/11] fix(events): build SQL Server guards from validated metadata --- docs/changelog.rst | 3 ++ sqlspec/adapters/arrow_odbc/events/store.py | 18 ++---------- sqlspec/adapters/mssql_python/events/store.py | 29 ++++--------------- sqlspec/adapters/pymssql/events/store.py | 29 ++++--------------- .../test_mssql_python/test_events_store.py | 16 +++++++--- .../adapters/test_pymssql/test_extensions.py | 15 +++++----- 6 files changed, 35 insertions(+), 75 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 51588b38a..338b1d7f3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -48,6 +48,9 @@ Unreleased **Fixed:** +* SQL Server event queue DDL guards use the configured table and index names. + Arrow ODBC index checks no longer include column text in the table name. + * Arrow read failures in the mssql-python adapter use SQLSpec error types. * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. diff --git a/sqlspec/adapters/arrow_odbc/events/store.py b/sqlspec/adapters/arrow_odbc/events/store.py index 282ebd942..56d828936 100644 --- a/sqlspec/adapters/arrow_odbc/events/store.py +++ b/sqlspec/adapters/arrow_odbc/events/store.py @@ -1,6 +1,5 @@ """arrow-odbc event queue store with T-SQL and Db2 DDL.""" -import re from typing import Final from sqlspec.adapters.arrow_odbc.config import ArrowOdbcConfig @@ -68,26 +67,15 @@ def _wrap_create_statement(self, statement: str, object_type: str) -> str: if self._db2: return statement if object_type == "table": - match = re.search(r"CREATE TABLE\s+(\S+)", statement, re.IGNORECASE) - if match: - table_name = match.group(1) - return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NULL BEGIN {statement}; END" + return f"IF OBJECT_ID(N'{_object_name(self.table_name)}', N'U') IS NULL BEGIN {statement}; END" if object_type == "index": - match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+([^\s(]+)", statement, re.IGNORECASE) - if match: - index_name = match.group(1).strip("[]") - table_name = match.group(2) - return f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'{index_name}' AND object_id = OBJECT_ID(N'{_object_name(table_name)}')) BEGIN {statement}; END" + return f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'{self._index_name()}' AND object_id = OBJECT_ID(N'{_object_name(self.table_name)}')) BEGIN {statement}; END" return statement def _wrap_drop_statement(self, statement: str) -> str: if self._db2: return statement - match = re.search(r"DROP TABLE\s+(\S+)", statement, re.IGNORECASE) - if match: - table_name = match.group(1) - return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NOT NULL DROP TABLE {table_name};" - return statement + return f"IF OBJECT_ID(N'{_object_name(self.table_name)}', N'U') IS NOT NULL {statement};" def _split_table_name(table_name: str) -> tuple[str, str]: diff --git a/sqlspec/adapters/mssql_python/events/store.py b/sqlspec/adapters/mssql_python/events/store.py index 854a9f22f..b3d273603 100644 --- a/sqlspec/adapters/mssql_python/events/store.py +++ b/sqlspec/adapters/mssql_python/events/store.py @@ -1,7 +1,5 @@ """mssql-python event queue store with T-SQL-specific DDL.""" -import re - from sqlspec.adapters.mssql_python.config import MssqlPythonConfig from sqlspec.extensions.events import BaseEventQueueStore from sqlspec.utils.text import split_qualified_identifier @@ -12,8 +10,8 @@ _QUALIFIED_IDENTIFIER_MIN_PARTS = 2 -class _MssqlPythonEventStoreMixin: - """Shared T-SQL DDL hooks for sync and async event queue stores.""" +class MssqlPythonEventQueueStore(BaseEventQueueStore[MssqlPythonConfig]): + """T-SQL DDL hooks for the event queue store.""" __slots__ = () @@ -33,30 +31,13 @@ def _timestamp_default(self) -> str: def _wrap_create_statement(self, statement: str, object_type: str) -> str: if object_type == "table": - match = re.search(r"CREATE TABLE\s+(\S+)", statement, re.IGNORECASE) - if match: - table_name = match.group(1) - return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NULL BEGIN {statement}; END" + return f"IF OBJECT_ID(N'{_object_name(self.table_name)}', N'U') IS NULL BEGIN {statement}; END" if object_type == "index": - match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+([^\s(]+)", statement, re.IGNORECASE) - if match: - index_name = match.group(1).strip("[]") - table_name = match.group(2) - return f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'{index_name}' AND object_id = OBJECT_ID(N'{_object_name(table_name)}')) BEGIN {statement}; END" + return f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'{self._index_name()}' AND object_id = OBJECT_ID(N'{_object_name(self.table_name)}')) BEGIN {statement}; END" return statement def _wrap_drop_statement(self, statement: str) -> str: - match = re.search(r"DROP TABLE\s+(\S+)", statement, re.IGNORECASE) - if match: - table_name = match.group(1) - return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NOT NULL DROP TABLE {table_name};" - return statement - - -class MssqlPythonEventQueueStore(_MssqlPythonEventStoreMixin, BaseEventQueueStore[MssqlPythonConfig]): - """Event queue DDL for mssql-python sync configs.""" - - __slots__ = () + return f"IF OBJECT_ID(N'{_object_name(self.table_name)}', N'U') IS NOT NULL {statement};" def _split_table_name(table_name: str) -> tuple[str, str]: diff --git a/sqlspec/adapters/pymssql/events/store.py b/sqlspec/adapters/pymssql/events/store.py index 89c17cdc8..ccdd2383e 100644 --- a/sqlspec/adapters/pymssql/events/store.py +++ b/sqlspec/adapters/pymssql/events/store.py @@ -1,7 +1,5 @@ """pymssql event queue store with T-SQL-specific DDL.""" -import re - from sqlspec.adapters.pymssql.config import PymssqlConfig from sqlspec.extensions.events import BaseEventQueueStore from sqlspec.utils.text import split_qualified_identifier @@ -12,8 +10,8 @@ _QUALIFIED_IDENTIFIER_MIN_PARTS = 2 -class _PymssqlEventStoreMixin: - """Shared T-SQL DDL hooks for sync and async event queue stores.""" +class PymssqlEventQueueStore(BaseEventQueueStore[PymssqlConfig]): + """T-SQL DDL hooks for the event queue store.""" __slots__ = () @@ -33,30 +31,13 @@ def _timestamp_default(self) -> str: def _wrap_create_statement(self, statement: str, object_type: str) -> str: if object_type == "table": - match = re.search(r"CREATE TABLE\s+(\S+)", statement, re.IGNORECASE) - if match: - table_name = match.group(1) - return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NULL BEGIN {statement}; END" + return f"IF OBJECT_ID(N'{_object_name(self.table_name)}', N'U') IS NULL BEGIN {statement}; END" if object_type == "index": - match = re.search(r"CREATE INDEX\s+(\S+)\s+ON\s+([^\s(]+)", statement, re.IGNORECASE) - if match: - index_name = match.group(1).strip("[]") - table_name = match.group(2) - return f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'{index_name}' AND object_id = OBJECT_ID(N'{_object_name(table_name)}')) BEGIN {statement}; END" + return f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'{self._index_name()}' AND object_id = OBJECT_ID(N'{_object_name(self.table_name)}')) BEGIN {statement}; END" return statement def _wrap_drop_statement(self, statement: str) -> str: - match = re.search(r"DROP TABLE\s+(\S+)", statement, re.IGNORECASE) - if match: - table_name = match.group(1) - return f"IF OBJECT_ID(N'{_object_name(table_name)}', N'U') IS NOT NULL DROP TABLE {table_name};" - return statement - - -class PymssqlEventQueueStore(_PymssqlEventStoreMixin, BaseEventQueueStore[PymssqlConfig]): - """Event queue DDL for pymssql sync configs.""" - - __slots__ = () + return f"IF OBJECT_ID(N'{_object_name(self.table_name)}', N'U') IS NOT NULL {statement};" def _split_table_name(table_name: str) -> tuple[str, str]: diff --git a/tests/unit/adapters/test_mssql_python/test_events_store.py b/tests/unit/adapters/test_mssql_python/test_events_store.py index ffca80002..5d5e646d1 100644 --- a/tests/unit/adapters/test_mssql_python/test_events_store.py +++ b/tests/unit/adapters/test_mssql_python/test_events_store.py @@ -23,10 +23,6 @@ def test_event_queue_store_uses_tsql_column_types_and_idempotency() -> None: assert "available_at DATETIME2(6) NOT NULL DEFAULT SYSUTCDATETIME()" in ddl assert "IF NOT EXISTS (SELECT 1 FROM sys.indexes" in ddl assert "OBJECT_ID(N'[dbo].[sqlspec_event_queue]')" in ddl - wrapped_no_space = store._wrap_create_statement( - "CREATE INDEX idx_events_channel_status ON app_events(channel, status, available_at)", "index" - ) - assert "OBJECT_ID(N'[dbo].[app_events]')" in wrapped_no_space def test_event_queue_store_drop_uses_object_id_guard() -> None: @@ -40,3 +36,15 @@ def test_event_queue_store_drop_uses_object_id_guard() -> None: def test_object_name_preserves_bracket_quoted_dots() -> None: assert _object_name("[dbo.schema].[sqlspec.event.queue]") == "[dbo.schema].[sqlspec.event.queue]" + + +def test_event_queue_guards_use_configured_schema_and_index() -> None: + store = MssqlPythonEventQueueStore(cast(MssqlPythonConfig, _config({"events": {"queue_table": "app.app_events"}}))) + statements = store.create_statements() + assert "OBJECT_ID(N'[app].[app_events]', N'U')" in statements[0] + assert "name = N'idx_app_app_events_channel_status'" in statements[1] + assert "OBJECT_ID(N'[app].[app_events]')" in statements[1] + assert "ON app.app_events(channel, status, available_at)" in statements[1] + assert store.drop_statements() == [ + "IF OBJECT_ID(N'[app].[app_events]', N'U') IS NOT NULL DROP TABLE app.app_events;" + ] diff --git a/tests/unit/adapters/test_pymssql/test_extensions.py b/tests/unit/adapters/test_pymssql/test_extensions.py index 36837fa72..c1c328480 100644 --- a/tests/unit/adapters/test_pymssql/test_extensions.py +++ b/tests/unit/adapters/test_pymssql/test_extensions.py @@ -13,14 +13,13 @@ def test_event_store_uses_tsql_column_types_and_idempotent_wrappers() -> None: store = PymssqlEventQueueStore(PymssqlConfig(extension_config={"events": {"queue_table": "event_queue"}})) - assert store._column_types() == ("NVARCHAR(MAX)", "NVARCHAR(MAX)", "DATETIME2(6)") - assert store._timestamp_default() == "SYSUTCDATETIME()" - assert "OBJECT_ID" in store._wrap_create_statement("CREATE TABLE event_queue (id INT)", "table") - assert "sys.indexes" in store._wrap_create_statement("CREATE INDEX idx_events ON event_queue (channel)", "index") - wrapped_no_space = store._wrap_create_statement( - "CREATE INDEX idx_events_channel_status ON app_events(channel, status, available_at)", "index" - ) - assert "OBJECT_ID(N'[dbo].[app_events]')" in wrapped_no_space + statements = store.create_statements() + assert "payload_json NVARCHAR(MAX)" in statements[0] + assert "SYSUTCDATETIME()" in statements[0] + assert "OBJECT_ID(N'[dbo].[event_queue]', N'U')" in statements[0] + assert "name = N'idx_event_queue_channel_status'" in statements[1] + assert "OBJECT_ID(N'[dbo].[event_queue]')" in statements[1] + assert store.drop_statements() == ["IF OBJECT_ID(N'[dbo].[event_queue]', N'U') IS NOT NULL DROP TABLE event_queue;"] def test_litestar_store_ddl_is_tsql_idempotent() -> None: