"""pymssql SQL Server driver implementation."""
import contextlib
from collections.abc import Sized
from typing import TYPE_CHECKING, Any, cast
from sqlspec.adapters.pymssql._typing import (
PYMSSQL_MODULE,
PymssqlConnection,
PymssqlCursor,
PymssqlRawCursor,
PymssqlSessionContext,
)
from sqlspec.adapters.pymssql.core import (
collect_rows,
create_mapped_exception,
default_statement_config,
driver_profile,
normalize_execute_many_parameters,
normalize_execute_parameters,
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.driver import (
BaseSyncExceptionHandler,
ExecutionResult,
SyncDriverAdapterBase,
SyncRowStream,
rows_to_dicts,
)
from sqlspec.driver._common import validate_savepoint_name
from sqlspec.exceptions import SQLSpecError
from sqlspec.utils.logging import get_logger
if TYPE_CHECKING:
from collections.abc import Sequence
from pymssql._pymssql import QueryParams
__all__ = ("PymssqlCursor", "PymssqlDriver", "PymssqlExceptionHandler", "PymssqlSessionContext")
logger = get_logger("sqlspec.adapters.pymssql")
pymssql = PYMSSQL_MODULE
class PymssqlExceptionHandler(BaseSyncExceptionHandler):
"""Context manager for handling pymssql exceptions."""
__slots__ = ()
def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool:
if exc_type is None:
return False
error_type = _pymssql_error_type()
if isinstance(exc_val, error_type):
self.pending_exception = create_mapped_exception(cast("Exception", exc_val), logger=logger)
return True
return False
class PymssqlStreamSource:
"""Native pymssql chunk source backed by ``cursor.fetchmany()``."""
__slots__ = ("_chunk_size", "_column_names", "_cursor_manager", "_driver", "_parameters", "_sql")
def __init__(self, driver: "PymssqlDriver", sql: str, parameters: Any, chunk_size: int) -> None:
self._driver = driver
self._sql = sql
self._parameters = parameters
self._chunk_size = chunk_size
self._cursor_manager: PymssqlCursor | None = None
self._column_names: list[str] | None = None
def start(self) -> None:
cursor_manager = self._driver.with_cursor(self._driver.connection)
try:
cursor = cursor_manager.__enter__()
handler = self._driver.handle_database_exceptions()
with handler:
cursor.execute(self._sql, normalize_execute_parameters(self._parameters))
self._driver._check_pending_exception(handler)
except BaseException:
with contextlib.suppress(Exception):
cursor_manager.__exit__(None, None, None)
raise
self._cursor_manager = cursor_manager
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 []
cursor = cursor_manager.cursor
handler = self._driver.handle_database_exceptions()
rows: Any = []
with handler:
rows = cursor.fetchmany(self._chunk_size)
self._driver._check_pending_exception(handler)
if not rows:
return []
column_names = self._column_names
if column_names is None:
column_names = resolve_column_names(cursor.description or None, self._driver._column_name_cache)
self._column_names = column_names
return rows_to_dicts(rows, column_names)
def close(self, error: bool = False) -> None:
cursor_manager = self._cursor_manager
self._cursor_manager = None
if cursor_manager is not None:
with contextlib.suppress(Exception):
cursor_manager.__exit__(None, None, None)
class PymssqlDriver(SyncDriverAdapterBase):
"""SQL Server database driver using pymssql."""
__slots__ = ("_column_name_cache", "_data_dictionary")
dialect = "tsql"
def __init__(
self,
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(
enable_caching=get_cache_config().compiled_cache_enabled
)
super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features)
self._data_dictionary: PymssqlSyncDataDictionary | None = None
self._column_name_cache: dict[int, tuple[Any, list[str]]] = {}