diff --git a/docs/changelog.rst b/docs/changelog.rst index 8769e74ec..4aa6a5cb8 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -17,6 +17,10 @@ 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. +* 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. + * SQLite and aiosqlite can register custom window functions on Python 3.11 and later when the SQLite runtime supports them. Choose a transaction lock mode or set the batch size for Arrow imports. Defaults stay the same. @@ -48,6 +52,16 @@ 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. + +* 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. * SQLite pools replace lost in-memory connections. Arrow imports roll back writes on failure or cancellation when the adapter owns the transaction. @@ -89,13 +103,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. 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/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/adk/store.py b/sqlspec/adapters/mssql_python/adk/store.py index 7540f0a7d..1626a210a 100644 --- a/sqlspec/adapters/mssql_python/adk/store.py +++ b/sqlspec/adapters/mssql_python/adk/store.py @@ -1,12 +1,12 @@ """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 typing_extensions import NotRequired from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor, MssqlPythonError +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 @@ -26,7 +26,6 @@ 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" @@ -848,17 +847,7 @@ def _cursor_rowcount(cursor: Any) -> int: 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: diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index 7d7e26dbf..dda0049d8 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -40,10 +40,11 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "extract_error_number", "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 @@ -98,9 +99,54 @@ } +def extract_error_number(exc: BaseException | None) -> int | None: + """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(",") + 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: + 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] + return int(last_match[0] or last_match[1]) + + 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(): @@ -131,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] @@ -330,16 +378,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..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..c35982958 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -419,21 +419,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) diff --git a/sqlspec/adapters/mssql_python/events/store.py b/sqlspec/adapters/mssql_python/events/store.py index 09bee9b9f..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/mssql_python/litestar/store.py b/sqlspec/adapters/mssql_python/litestar/store.py index 1e4464f96..d48001e5b 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 TYPE_CHECKING, Any +from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -97,12 +98,9 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: + with MssqlPythonCursor(conn) as cursor: cursor.execute(sql, (key,)) row = cursor.fetchone() - finally: - cursor.close() if row is None: return None @@ -111,8 +109,7 @@ 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: + with MssqlPythonCursor(conn) as update_cursor: update_cursor.execute( f""" UPDATE {self._table_name} @@ -121,8 +118,6 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | """, (new_expires_at, key), ) - finally: - update_cursor.close() conn.commit() return _coerce_bytes(_row_value(row, "data", 0)) @@ -143,30 +138,18 @@ 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() + 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_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"DELETE FROM {self._table_name} WHERE session_id = ?", (key,)) - finally: - cursor.close() + 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_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"TRUNCATE TABLE {self._table_name}") - finally: - cursor.close() + 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() @@ -177,22 +160,14 @@ 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() + 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_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_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 @@ -208,13 +183,9 @@ 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() + 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) diff --git a/sqlspec/adapters/mssql_python/pool.py b/sqlspec/adapters/mssql_python/pool.py index f6a1643c4..9d9a15faa 100644 --- a/sqlspec/adapters/mssql_python/pool.py +++ b/sqlspec/adapters/mssql_python/pool.py @@ -51,8 +51,9 @@ 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": if self._closed: diff --git a/sqlspec/adapters/pymssql/adk/store.py b/sqlspec/adapters/pymssql/adk/store.py index 5658e83b4..7d1737f93 100644 --- a/sqlspec/adapters/pymssql/adk/store.py +++ b/sqlspec/adapters/pymssql/adk/store.py @@ -1,12 +1,12 @@ """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 typing_extensions import NotRequired from sqlspec.adapters.pymssql._typing import PymssqlCursor, PymssqlError +from sqlspec.adapters.pymssql.core import extract_error_number, resolve_rowcount from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo from sqlspec.config import ADKConfig from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options @@ -28,7 +28,6 @@ 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" @@ -411,7 +410,7 @@ def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[A 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) + rowcount = resolve_rowcount(cursor) if commit: conn.commit() return rowcount @@ -484,7 +483,7 @@ 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) + inserted += resolve_rowcount(cursor) conn.commit() return inserted @@ -585,7 +584,7 @@ def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[A 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) + rowcount = resolve_rowcount(cursor) if commit: conn.commit() return rowcount @@ -855,24 +854,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: diff --git a/sqlspec/adapters/pymssql/config.py b/sqlspec/adapters/pymssql/config.py index a2e1691da..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]]] @@ -89,6 +90,8 @@ class _PymssqlSessionConnectionHandler(SyncPoolSessionFactory): class PymssqlConfig(SyncDatabaseConfig[PymssqlConnection, PymssqlConnectionPool, PymssqlDriver]): """Configuration for pymssql synchronous connections.""" + __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 diff --git a/sqlspec/adapters/pymssql/core.py b/sqlspec/adapters/pymssql/core.py index a0729717b..6be2f5796 100644 --- a/sqlspec/adapters/pymssql/core.py +++ b/sqlspec/adapters/pymssql/core.py @@ -38,6 +38,7 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "extract_error_number", "format_identifier", "normalize_execute_many_parameters", "normalize_execute_parameters", @@ -46,7 +47,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]]] = { @@ -63,6 +64,14 @@ } +def _quote_bracket_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() @@ -144,7 +153,7 @@ def apply_driver_features( 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(): @@ -202,9 +211,10 @@ def collect_rows( 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: @@ -255,21 +265,26 @@ 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": + 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 not isinstance(val, bool) and val != 0: + return val + 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 not matches: return None - try: - return int(matches[-1]) - except ValueError: - return None + last_match = matches[-1] + return int(last_match[0] or last_match[1]) driver_profile = build_profile() 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/adapters/pymssql/events/store.py b/sqlspec/adapters/pymssql/events/store.py index 717cce409..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/sqlspec/adapters/pymssql/litestar/store.py b/sqlspec/adapters/pymssql/litestar/store.py index 77af491f1..6661d8b2a 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 TYPE_CHECKING, Any +from sqlspec.adapters.pymssql._typing import PymssqlCursor from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -97,12 +98,9 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | AND (expires_at IS NULL OR expires_at > SYSUTCDATETIME()) """ with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: + with PymssqlCursor(conn) as cursor: cursor.execute(sql, (key,)) row = cursor.fetchone() - finally: - cursor.close() if row is None: return None @@ -111,8 +109,7 @@ 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: + with PymssqlCursor(conn) as update_cursor: update_cursor.execute( f""" UPDATE {self._table_name} @@ -121,8 +118,6 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | """, (new_expires_at, key), ) - finally: - update_cursor.close() conn.commit() return _coerce_bytes(_row_value(row, "data", 0)) @@ -143,30 +138,18 @@ 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() + 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_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"DELETE FROM {self._table_name} WHERE session_id = %s", (key,)) - finally: - cursor.close() + 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_connection() as conn: - cursor = conn.cursor() - try: - cursor.execute(f"TRUNCATE TABLE {self._table_name}") - finally: - cursor.close() + 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() @@ -177,22 +160,14 @@ 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() + 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_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_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 @@ -208,13 +183,9 @@ 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() + 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) 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_config.py b/tests/unit/adapters/test_mssql_python/test_config.py index 0007a37fe..4d35a8473 100644 --- a/tests/unit/adapters/test_mssql_python/test_config.py +++ b/tests/unit/adapters/test_mssql_python/test_config.py @@ -346,3 +346,16 @@ def commit(self) -> None: raise RuntimeError assert calls == ["rollback", "release"] + + +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)) + 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 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_events_store.py b/tests/unit/adapters/test_mssql_python/test_events_store.py index bfa11578b..5d5e646d1 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,7 @@ 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 def test_event_queue_store_drop_uses_object_id_guard() -> None: @@ -35,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_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 == [] diff --git a/tests/unit/adapters/test_pymssql/test_core.py b/tests/unit/adapters/test_pymssql/test_core.py index 26cec8aeb..ee1e947a5 100644 --- a/tests/unit/adapters/test_pymssql/test_core.py +++ b/tests/unit/adapters/test_pymssql/test_core.py @@ -4,6 +4,17 @@ import pytest +from sqlspec.adapters.pymssql.core import ( + build_insert_statement, + collect_rows, + create_mapped_exception, + default_statement_config, + driver_profile, + extract_error_number, + format_identifier, + normalize_execute_many_parameters, + normalize_execute_parameters, +) from sqlspec.core import SQL, ParameterStyle from sqlspec.exceptions import ( CheckViolationError, @@ -16,8 +27,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 +44,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 +54,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 +66,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 +82,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 +117,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,16 +136,46 @@ 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,)] assert normalize_execute_many_parameters(rows) is rows + + +def test_extract_error_number() -> None: + """extract_error_number detects error number from attribute, tuple, or regex.""" + + 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 + + +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")] + description = [("id",), ("name",)] + rows, column_names, row_format = collect_rows(input_rows, description) + + assert rows is input_rows + assert column_names == ["id", "name"] + assert row_format == "tuple" + + +def test_normalize_execute_parameters_preserves_tuples() -> None: + """normalize_execute_parameters passes tuples through directly.""" + params = (1, "Alice") + assert normalize_execute_parameters(params) is params diff --git a/tests/unit/adapters/test_pymssql/test_extensions.py b/tests/unit/adapters/test_pymssql/test_extensions.py index f70b0c2a5..c1c328480 100644 --- a/tests/unit/adapters/test_pymssql/test_extensions.py +++ b/tests/unit/adapters/test_pymssql/test_extensions.py @@ -13,10 +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") + 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: