From 203eb0d96bd87180c2e0d32a987d2e2c10d6e035 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Wed, 30 Sep 2026 17:42:50 +0530 Subject: [PATCH] CHORE: Gate PRs on strict mssql_python typing AB#45120 Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .github/copilot-instructions.md | 6 +- .github/prompts/run-tests.prompt.md | 24 +- CONTRIBUTING.md | 39 ++ eng/pipelines/pr-validation-pipeline.yml | 12 + mssql_python/__init__.py | 12 +- mssql_python/_ddbc_types.pyi | 261 ++++++++++++ mssql_python/_pycore_types.pyi | 111 ++++++ mssql_python/async_query/__init__.py | 4 +- mssql_python/async_query/_native.py | 9 +- mssql_python/async_query/async_connection.py | 15 +- mssql_python/async_query/async_cursor.py | 16 +- mssql_python/async_query/async_fetch.py | 3 +- .../async_query/exception_translator.py | 26 +- mssql_python/auth.py | 27 +- mssql_python/connection.py | 60 +-- mssql_python/constants.py | 178 ++++++++- mssql_python/cursor.py | 377 +++++++++++------- mssql_python/db_connection.py | 2 +- mssql_python/ddbc_bindings.py | 23 +- mssql_python/decimal_config.py | 13 +- mssql_python/exceptions.py | 17 +- mssql_python/helpers.py | 8 +- mssql_python/logging.py | 57 +-- mssql_python/mssql_python.pyi | 70 ++-- mssql_python/parameter_helper.py | 16 +- mssql_python/perf_timer.py | 58 ++- mssql_python/pooling.py | 2 +- mssql_python/row.py | 59 ++- mssql_python/type.py | 12 +- pytest.ini | 6 +- requirements.txt | 5 +- setup.py | 1 + tests/test_000_dependencies.py | 52 +++ tests/test_004_cursor.py | 14 + tests/test_004_cursor_arrow.py | 33 ++ tests/test_005_connection_cursor_lifecycle.py | 18 + tests/test_typing.py | 40 ++ 37 files changed, 1313 insertions(+), 373 deletions(-) create mode 100644 mssql_python/_ddbc_types.pyi create mode 100644 mssql_python/_pycore_types.pyi create mode 100644 tests/test_typing.py diff --git a/.github/copilot-instructions.md b/.github/copilot-instructions.md index 6ade59b02..771946391 100644 --- a/.github/copilot-instructions.md +++ b/.github/copilot-instructions.md @@ -39,12 +39,12 @@ Core facts: ```bash black --check --line-length=100 mssql_python/ tests/ # BLOCKING in CI -python -m pytest -v # 'stress' marker excluded by default +python -m pytest -v # 'stress' and 'typing' excluded by default ``` - **`pr-format-check` (BLOCKING):** PR title must start with one of `FEAT: FIX: DOC: CHORE: STYLE: REFACTOR: PERF: RELEASE: AI:`; the body must link a work item/issue and have a `### Summary` of at least 10 characters. - Use `AI:` for AI tooling, agents, skills, prompts, and AI-assisted development workflows, not merely because AI helped write an ordinary fix or feature. -- `flake8`, `pylint`, `mypy`, `clang-format`, and `cpplint` run but are **informational**, not blocking. +- `flake8`, `pylint`, the GitHub lint workflow's `mypy` step, `clang-format`, and `cpplint` are **informational**. The separate source/stub typing harness (`python -m pytest tests/test_typing.py -m typing -v`) checks only `mssql_python` in strict mode and is **blocking** in the ADO PR-validation pipeline. `tests` and `mssql_python_odbc` are not typing-gate targets; the runtime test jobs are unchanged. - The authoritative cross-platform validation runs on **Azure DevOps** (broader OS / Python / arch coverage than the GitHub checks); consult the specific pipeline in `eng/pipelines/` for the exact matrix rather than assuming full coverage. A coverage bot posts a report comment on the PR. ## Code standards @@ -64,7 +64,7 @@ python -m pytest -v # 'stress' marker excl ## Testing conventions -- Test files are mostly numbered `test_NNN_*.py`; `tests/test_000_dependencies.py` runs without a DB, most others need a live SQL Server. `-m "not stress"` is the default. +- Test files are mostly numbered `test_NNN_*.py`; `tests/test_000_dependencies.py` runs without a DB, most others need a live SQL Server. `-m "not stress and not typing"` is the default; run source/stub typing checks separately with `python -m pytest tests/test_typing.py -m typing -v` (no DB required). - Run segfault-prone or ODBC/pool global-state tests in an **isolated subprocess** so a crash or shared state cannot poison the rest of the suite. - **Assert the contract, not just the output.** If a change's value is "we now call X once," assert the call/round-trip count — a correctness-only test won't catch a perf regression. - For global type-mapping changes, add typed-NULL integration cases (VARBINARY, UNIQUEIDENTIFIER, XML, DECIMAL, stored-proc params) before applying the optimization broadly. diff --git a/.github/prompts/run-tests.prompt.md b/.github/prompts/run-tests.prompt.md index da8bcfa88..16a7a1554 100644 --- a/.github/prompts/run-tests.prompt.md +++ b/.github/prompts/run-tests.prompt.md @@ -61,6 +61,11 @@ python -c "import pytest; print('✅ pytest ready:', pytest.__version__)" pip install pytest pytest-cov ``` +The typing harness checks only `mssql_python` source/stub files, not `tests` or +`mssql_python_odbc`, and does not require a database. After building the native +extension, run `python -m pytest tests/test_typing.py -m typing -v` +without completing the database checks below. + ### Step 3: Verify Database Connection String ```bash @@ -126,16 +131,17 @@ Help the developer run tests to validate their changes. Follow this process base | Category | Description | When to Use | |----------|-------------|-------------| -| **All tests** | Full test suite (excluding stress) | Before creating PR | +| **Default tests** | Test suite excluding stress and separately gated typing tests | Before creating PR | | **Specific file** | Single test file | Testing one area | | **Specific test** | Single test function | Debugging a failure | | **Stress tests** | Long-running, resource-intensive | Performance validation | +| **Typing tests** | Strict checking of `mssql_python` sources/stubs only; no database required | Python typing changes | | **With coverage** | Tests + coverage report | Checking coverage | ### Ask the Developer > "What would you like to test?" -> 1. **All tests** - Run full suite (recommended before PR) +> 1. **Default tests** - Run the suite excluding stress and separately gated typing tests > 2. **Specific tests** - Tell me which file(s) or test name(s) > 3. **With coverage** - Generate coverage report @@ -143,16 +149,16 @@ Help the developer run tests to validate their changes. Follow this process base ## STEP 2: Run Tests -### Option A: Run All Tests (Default - Excludes Stress Tests) +### Option A: Run Default Tests (Excludes Stress and Typing Tests) ```bash # From repository root python -m pytest -v -# This automatically applies: -m "not stress" (from pytest.ini) +# This automatically applies: -m "not stress and not typing" (from pytest.ini) ``` -### Option B: Run All Tests Including Stress Tests +### Option B: Run All Tests Including Stress and Typing Tests ```bash python -m pytest -v -m "" @@ -384,7 +390,7 @@ python -m pytest tests/ -v ### Common Commands ```bash -# Run all tests (default, excludes stress) +# Run default tests (excludes stress and separately gated typing tests) python -m pytest -v # Run specific file @@ -417,9 +423,11 @@ The project uses these default settings in `pytest.ini`: [pytest] markers = stress: marks tests as stress tests (long-running, resource-intensive) + slow: marks tests as extra-slow (sustained load, multi-minute duration) + typing: static mssql_python source and stub type checks (run separately in PR validation) -# Default: Skips stress tests -addopts = -m "not stress" +# Default: Skips stress and separately gated typing tests +addopts = -m "not stress and not typing" ``` --- diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index fdb5ba8ec..6186f92c8 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -44,3 +44,42 @@ All pull requests must include: Use `AI:` for changes to AI tooling, agents, skills, prompts, or AI-assisted development workflows. It describes the subject of the change, not whether AI helped write it; ordinary driver fixes and features keep their usual prefixes. + +## Type Checking + +The typing harness checks Python source (`.py`) and stub (`.pyi`) files recursively +under `mssql_python` only, excluding generated `build` directories and their +downloaded third-party sources. `tests` and `mssql_python_odbc` are not targets of +this typing gate. It does not execute the checked code or require a live database; +the existing runtime test jobs are unchanged. + +After installing `requirements.txt` and building the native extension using +`.github/prompts/build-ddbc.prompt.md`, run: + +```console +python -m pytest tests/test_typing.py -m typing -v +``` + +The harness runs mypy in strict mode, including checks inside unannotated function +bodies. It reports source, stub, and import errors without suppressing errors in +the driver. Explicit package bases resolve module names from the repository root. +Keep the mypy version pinned in `requirements.txt` aligned +with the existing locked development dependencies. + +To run the same package check directly: + +```console +python -m mypy --config-file= --strict --explicit-package-bases --exclude "(^|/)build/" --no-incremental mssql_python +``` + +Private `_ddbc_types.pyi` and `_pycore_types.pyi` declarations describe the native +boundaries; they do not replace or hide the Python implementations from mypy. +Keep these declarations aligned with the C++/Rust APIs, and keep the static +constant declarations aligned with the dynamically exported integer aliases. +The dependency tests check the native export names and constant declaration parity. + +The existing required `MSSQL-Python-PR-Validation` pipeline runs the harness once, +in the Ubuntu CodeQL job immediately after its native build. Typing failures fail +the pipeline and block PR merging; no separate GitHub workflow or required check +is needed. The `typing` marker is excluded from default pytest runs so the same +static checks are not repeated across the database/OS matrix. diff --git a/eng/pipelines/pr-validation-pipeline.yml b/eng/pipelines/pr-validation-pipeline.yml index 99507c8e6..cc07df01d 100644 --- a/eng/pipelines/pr-validation-pipeline.yml +++ b/eng/pipelines/pr-validation-pipeline.yml @@ -45,6 +45,18 @@ jobs: ./build.sh displayName: 'Build C++ extension for CodeQL analysis' + - script: | + python -m pytest tests/test_typing.py -m typing -v --junitxml=typing-test-results.xml + displayName: 'Gate mssql_python source and stub typing' + + - task: PublishTestResults@2 + condition: succeededOrFailed() + inputs: + testResultsFiles: 'typing-test-results.xml' + testRunTitle: 'mssql_python source and stub typing (Linux x64)' + failTaskOnFailedTests: true + displayName: 'Publish typing regression results' + - task: CodeQL3000Finalize@0 condition: always() displayName: 'Finalize CodeQL' diff --git a/mssql_python/__init__.py b/mssql_python/__init__.py index 06b382c9e..32deef4d8 100644 --- a/mssql_python/__init__.py +++ b/mssql_python/__init__.py @@ -80,7 +80,7 @@ from .odbc_provider import ProviderManager -def get_native_provider_info() -> dict: +def get_native_provider_info() -> dict[str, object]: """Return the selected native provider for diagnostics. Reports the provider ``id``, package ``version``, resolved ``driver_path``, @@ -91,17 +91,17 @@ def get_native_provider_info() -> dict: # Global registry for tracking active connections (using weak references) -_active_connections = weakref.WeakSet() +_active_connections: weakref.WeakSet[Connection] = weakref.WeakSet() _connections_lock = threading.Lock() -def _register_connection(conn): +def _register_connection(conn: Connection) -> None: """Register a connection for cleanup before shutdown.""" with _connections_lock: _active_connections.add(conn) -def _cleanup_connections(): +def _cleanup_connections() -> None: """ Cleanup function called by atexit to close all active connections. @@ -579,7 +579,7 @@ def pooling(max_size: int = 100, idle_timeout: int = 600, enabled: bool = True) _original_module_setattr = sys.modules[__name__].__setattr__ -def _custom_setattr(name, value): +def _custom_setattr(name: str, value: object) -> None: if name == "lowercase": with _settings_lock: _settings.lowercase = bool(value) @@ -590,7 +590,7 @@ def _custom_setattr(name, value): # Replace the module's __setattr__ with our custom version -sys.modules[__name__].__setattr__ = _custom_setattr +setattr(sys.modules[__name__], "__setattr__", _custom_setattr) # Create a custom module class that uses properties instead of __setattr__ diff --git a/mssql_python/_ddbc_types.pyi b/mssql_python/_ddbc_types.pyi new file mode 100644 index 000000000..9daa973d0 --- /dev/null +++ b/mssql_python/_ddbc_types.pyi @@ -0,0 +1,261 @@ +"""Native declarations; kept separate so mypy also checks the Python loader.""" + +from collections.abc import Callable, Mapping, Sequence +from typing import Any, NoReturn, TypedDict, overload + +from .row import Row + +__all__ = [ + "ARCHITECTURE", + "Connection", + "DDBCSQLCheckError", + "DDBCSQLColumns", + "DDBCSQLDescribeCol", + "DDBCSQLExecDirect", + "DDBCSQLExecute", + "DDBCSQLFetch", + "DDBCSQLFetchAll", + "DDBCSQLFetchArrowBatch", + "DDBCSQLFetchMany", + "DDBCSQLFetchOne", + "DDBCSQLFetchScroll", + "DDBCSQLForeignKeys", + "DDBCSQLFreeHandle", + "DDBCSQLGetAllDiagRecords", + "DDBCSQLGetData", + "DDBCSQLGetTypeInfo", + "DDBCSQLMoreResults", + "DDBCSQLNumResultCols", + "DDBCSQLPrimaryKeys", + "DDBCSQLProcedures", + "DDBCSQLResetStmt", + "DDBCSQLRowCount", + "DDBCSQLSetStmtAttr", + "DDBCSQLSpecialColumns", + "DDBCSQLStatistics", + "DDBCSQLTables", + "DDBCSetDecimalSeparator", + "ErrorInfo", + "GetDriverPathCpp", + "NumericData", + "ParamInfo", + "SQLExecuteMany", + "SQL_NO_TOTAL", + "SqlHandle", + "ThrowStdException", + "close_pooling", + "construct_rows", + "disable_pooling", + "enable_pooling", + "update_log_level", +] + +class EncodingSettings(TypedDict): + encoding: str + ctype: int + +class ColumnMetadata(TypedDict): + ColumnName: str + DataType: int + ColumnSize: int + DecimalDigits: int + Nullable: int + +class InfoResult(TypedDict): + data: bytes + length: int + info_type: int + +class SqlHandle: + def free(self) -> None: ... + def _close_cursor(self) -> None: ... + def _cancel(self) -> None: ... + +class Connection: + def __init__( + self, + conn_str: str, + use_pool: bool, + attrs_before: dict[int, int | str | bytes] = ..., + pool_key: str = "", + token_factory: Callable[[], tuple[dict[int, int | str | bytes], int | None]] | None = None, + ) -> None: ... + def alloc_statement_handle(self) -> SqlHandle: ... + def close(self, transaction_already_rolled_back: bool = False) -> None: ... + def commit(self) -> None: ... + def rollback(self) -> None: ... + def get_autocommit(self) -> bool: ... + def set_autocommit(self, value: bool) -> None: ... + def set_attr(self, attribute: int, value: int | str | bytes | bytearray) -> None: ... + def get_info(self, info_type: int) -> InfoResult | None: ... + +class ParamInfo: + inputOutputType: int + paramCType: int + paramSQLType: int + columnSize: int + decimalDigits: int + strLenOrInd: int + dataPtr: object + isDAE: bool + def __init__(self) -> None: ... + +class NumericData: + precision: int + scale: int + sign: int + val: str | bytes + @overload + def __init__(self) -> None: ... + @overload + def __init__(self, precision: int, scale: int, sign: int, val: str | bytes) -> None: ... + +class ErrorInfo: + sqlState: str + ddbcErrorMsg: str + +ARCHITECTURE: str +SQL_NO_TOTAL: int + +def _get_odbc_driver_path(base_dir: str, provider: str) -> str: ... +def _set_odbc_provider(provider: str) -> None: ... +def GetDriverPathCpp(base_dir: str) -> str: ... +def ThrowStdException(message: str) -> NoReturn: ... +def enable_pooling(max_size: int, idle_timeout: int) -> None: ... +def disable_pooling() -> None: ... +def close_pooling() -> None: ... +def update_log_level(level: int) -> None: ... +def DDBCSetDecimalSeparator(separator: str) -> None: ... +def DDBCSQLCheckError(handle_type: int, handle: SqlHandle | None, ret: int) -> ErrorInfo: ... +def DDBCSQLGetAllDiagRecords(handle: SqlHandle | None) -> list[tuple[str, str]]: ... +def DDBCSQLExecDirect(handle: SqlHandle | None, query: str) -> int: ... +def DDBCSQLExecute( + statementHandle: SqlHandle | None, + query: str, + params: list[Any], + inputSizes: list[tuple[int, int, int, int]] | None, + isStmtPrepared: list[bool], + usePrepare: bool, + encodingSettings: EncodingSettings, +) -> int: ... +def SQLExecuteMany( + statementHandle: SqlHandle | None, + query: str, + columnwise_params: list[list[Any]], + paramInfos: Sequence[ParamInfo], + paramSetSize: int, + encodingSettings: EncodingSettings, +) -> int: ... +def DDBCSQLRowCount(handle: SqlHandle | None) -> int: ... +def DDBCSQLFetch(handle: SqlHandle | None) -> int: ... +def DDBCSQLNumResultCols( + statementHandle: SqlHandle | None, messages: list[tuple[str, str]] | None = None +) -> int: ... +def DDBCSQLDescribeCol( + StatementHandle: SqlHandle | None, + ColumnMetadata: list[ColumnMetadata], + messages: list[tuple[str, str]] | None = None, +) -> int: ... +def DDBCSQLGetData( + StatementHandle: SqlHandle | None, + colCount: int, + row: list[Any], + charEncoding: str, + wcharEncoding: str, + charCtype: int, + messages: list[tuple[str, str]] | None = None, +) -> int: ... +def DDBCSQLMoreResults(handle: SqlHandle | None) -> int: ... +def DDBCSQLFetchOne( + StatementHandle: SqlHandle | None, + row: list[Any], + charEncoding: str = "utf-16le", + wcharEncoding: str = "utf-16le", + charCtype: int = -8, + messages: list[tuple[str, str]] | None = None, +) -> int: ... +def DDBCSQLFetchMany( + StatementHandle: SqlHandle | None, + rows: list[list[Any]], + fetchSize: int, + charEncoding: str = "utf-16le", + wcharEncoding: str = "utf-16le", + charCtype: int = -8, + messages: list[tuple[str, str]] | None = None, +) -> int: ... +def DDBCSQLFetchAll( + StatementHandle: SqlHandle | None, + rows: list[list[Any]], + charEncoding: str = "utf-16le", + wcharEncoding: str = "utf-16le", + charCtype: int = -8, + messages: list[tuple[str, str]] | None = None, +) -> int: ... +def DDBCSQLFetchArrowBatch( + StatementHandle: SqlHandle | None, + capsules: list[object], + arrowBatchSize: int, + charCtype: int, + messages: list[tuple[str, str]] | None = None, +) -> int: ... +def DDBCSQLFreeHandle(handle_type: int, handle: SqlHandle | None) -> int: ... +def DDBCSQLResetStmt(handle: SqlHandle | None) -> int: ... +def DDBCSQLSetStmtAttr(handle: SqlHandle | None, attribute: int, value: object) -> int: ... +def DDBCSQLTables( + StatementHandle: SqlHandle | None, + catalog: str = "", + schema: str = "", + table: str = "", + tableType: str = "", +) -> int: ... +def DDBCSQLFetchScroll( + handle: SqlHandle | None, orientation: int, offset: int, row: list[Any] +) -> int: ... +def DDBCSQLGetTypeInfo(StatementHandle: SqlHandle | None, DataType: int) -> int: ... +def DDBCSQLProcedures( + handle: SqlHandle | None, catalog: str | None, schema: str | None, procedure: str | None +) -> int: ... +def DDBCSQLForeignKeys( + handle: SqlHandle | None, + pk_catalog: str | None, + pk_schema: str | None, + pk_table: str | None, + fk_catalog: str | None, + fk_schema: str | None, + fk_table: str | None, +) -> int: ... +def DDBCSQLPrimaryKeys( + handle: SqlHandle | None, catalog: str | None, schema: str | None, table: str +) -> int: ... +def DDBCSQLSpecialColumns( + handle: SqlHandle | None, + identifier: int, + catalog: str | None, + schema: str | None, + table: str, + scope: int, + nullable: int, +) -> int: ... +def DDBCSQLStatistics( + handle: SqlHandle | None, + catalog: str | None, + schema: str | None, + table: str, + unique: int, + accuracy: int, +) -> int: ... +def DDBCSQLColumns( + handle: SqlHandle | None, + catalog: str | None, + schema: str | None, + table: str | None, + column: str | None, +) -> int: ... +def construct_rows( + rows_data: list[list[Any]], + row_class: type[Row], + column_map: Mapping[str, int] | None, + cursor: object, + column_map_lower: Mapping[str, int] | None = None, + column_names: tuple[str, ...] | None = None, +) -> list[Row]: ... diff --git a/mssql_python/_pycore_types.pyi b/mssql_python/_pycore_types.pyi new file mode 100644 index 000000000..49cca45b0 --- /dev/null +++ b/mssql_python/_pycore_types.pyi @@ -0,0 +1,111 @@ +"""Structural contracts for the dynamically loaded Rust extension.""" + +from collections.abc import Callable, Iterable, Mapping, Sequence +from types import TracebackType +from typing import Any, Protocol, TypedDict + +CoreContext = dict[str, str | int | Callable[[str, str, str], bytes]] + +class BulkCopyResult(TypedDict): + rows_copied: int + batch_count: int + elapsed_time: float + rows_per_second: float + +class CoreCursor(Protocol): + def close(self) -> None: ... + def bulkcopy( + self, + table_name: str, + data_source: Iterable[tuple[Any, ...]], + batch_size: int = 0, + timeout: int = 30, + column_mappings: list[str] | list[tuple[int, str]] | None = None, + keep_identity: bool = False, + check_constraints: bool = False, + table_lock: bool = False, + keep_nulls: bool = False, + fire_triggers: bool = False, + use_internal_transaction: bool = False, + python_logger: object = None, + ) -> BulkCopyResult: ... + def bulkcopy_arrow( + self, + table_name: str, + source: object, + batch_size: int = 0, + timeout: int = 30, + column_mappings: list[str] | list[tuple[int, str]] | None = None, + keep_identity: bool = False, + check_constraints: bool = False, + table_lock: bool = False, + keep_nulls: bool = False, + fire_triggers: bool = False, + use_internal_transaction: bool = False, + python_logger: object = None, + ) -> BulkCopyResult: ... + +class CoreConnection(Protocol): + def __init__( + self, client_context_dict: Mapping[str, object], python_logger: object = None + ) -> None: ... + def cursor(self) -> CoreCursor: ... + def close(self) -> None: ... + +class AsyncCoreCursor(Protocol): + arraysize: int + @property + def timeout(self) -> int: ... + @property + def rowcount(self) -> int: ... + @property + def description(self) -> list[tuple[Any, ...]] | None: ... + async def execute( + self, operation: str, *parameters: Any, use_prepare: bool = True, reset_cursor: bool = True + ) -> object: ... + async def executemany( + self, + operation: str, + seq_of_parameters: Sequence[Sequence[Any]] | Sequence[Mapping[str, Any]], + *, + use_prepare: bool = True, + ) -> None: ... + def setinputsizes(self, sizes: Sequence[int | tuple[int, ...]]) -> None: ... + async def fetchone(self) -> tuple[Any, ...] | None: ... + async def fetchmany(self, size: int | None = None) -> list[tuple[Any, ...]]: ... + async def fetchall(self) -> list[tuple[Any, ...]]: ... + async def nextset(self) -> bool: ... + async def close(self) -> None: ... + +class AsyncCoreConnection(Protocol): + timeout: int + @property + def autocommit(self) -> bool: ... + @property + def closed(self) -> bool: ... + @classmethod + async def connect( + cls, + client_context_dict: Mapping[str, object], + python_logger: object = None, + autocommit: bool = False, + ) -> AsyncCoreConnection: ... + def cursor(self) -> AsyncCoreCursor: ... + async def commit(self) -> None: ... + async def rollback(self) -> None: ... + async def close(self) -> None: ... + async def __aenter__(self) -> AsyncCoreConnection: ... + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: ... + def is_connected(self) -> bool: ... + +class PyCoreModule(Protocol): + __name__: str + PyCoreConnection: type[CoreConnection] + PyCoreCursor: type[CoreCursor] + PyAsyncConnection: type[AsyncCoreConnection] + PyAsyncCursor: type[AsyncCoreCursor] diff --git a/mssql_python/async_query/__init__.py b/mssql_python/async_query/__init__.py index 95f8f0ecf..5287d8b02 100644 --- a/mssql_python/async_query/__init__.py +++ b/mssql_python/async_query/__init__.py @@ -10,8 +10,8 @@ """ from ._native import load_py_core -from .async_connection import _AsyncConnection # pyright: ignore[reportPrivateUsage] -from .async_cursor import _AsyncCursor # pyright: ignore[reportPrivateUsage] +from .async_connection import _AsyncConnection as _AsyncConnection +from .async_cursor import _AsyncCursor as _AsyncCursor from .exception_translator import ( DataError, DatabaseError, diff --git a/mssql_python/async_query/_native.py b/mssql_python/async_query/_native.py index 8cd282335..490001fd3 100644 --- a/mssql_python/async_query/_native.py +++ b/mssql_python/async_query/_native.py @@ -1,10 +1,13 @@ """Native mssql-py-core dependency boundary for asynchronous queries.""" from importlib import import_module -from types import ModuleType +from typing import TYPE_CHECKING, cast +if TYPE_CHECKING: + from .._pycore_types import PyCoreModule -def load_py_core() -> ModuleType: + +def load_py_core() -> "PyCoreModule": """Load the PyO3 extension that owns asynchronous TDS operations.""" try: py_core = import_module("mssql_py_core") @@ -20,4 +23,4 @@ def load_py_core() -> ModuleType: missing = ", ".join(missing_types) raise ImportError(f"mssql_py_core does not provide the required async types: {missing}") - return py_core + return cast("PyCoreModule", py_core) diff --git a/mssql_python/async_query/async_connection.py b/mssql_python/async_query/async_connection.py index 865d172b8..3abc567a6 100644 --- a/mssql_python/async_query/async_connection.py +++ b/mssql_python/async_query/async_connection.py @@ -6,7 +6,11 @@ may change without notice. """ -from typing import Any, Optional +from types import TracebackType +from typing import Any, Optional, TYPE_CHECKING + +if TYPE_CHECKING: + from .._pycore_types import AsyncCoreConnection from ..logging import logger from ._native import load_py_core @@ -46,7 +50,7 @@ class _AsyncConnection: ProgrammingError = ProgrammingError NotSupportedError = NotSupportedError - def __init__(self, py_core_async_connection: Any) -> None: + def __init__(self, py_core_async_connection: "AsyncCoreConnection") -> None: self._py_core_async_connection = py_core_async_connection @classmethod @@ -116,7 +120,12 @@ async def __aenter__(self) -> "_AsyncConnection": logger.debug("AsyncConnection.__aenter__: context entered") return self - async def __aexit__(self, exc_type, exc_value, traceback) -> Any: + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: logger.debug( "AsyncConnection.__aexit__: exiting context; block_error=%s", exc_type is not None, diff --git a/mssql_python/async_query/async_cursor.py b/mssql_python/async_query/async_cursor.py index 2547e2fae..d5d69095b 100644 --- a/mssql_python/async_query/async_cursor.py +++ b/mssql_python/async_query/async_cursor.py @@ -9,7 +9,7 @@ import asyncio from collections.abc import Mapping, Sequence from contextlib import asynccontextmanager -from typing import Any, Optional +from typing import Any, Optional, TYPE_CHECKING, AsyncIterator import uuid from ..exceptions import OperationalError @@ -19,6 +19,10 @@ from . import async_execute, async_fetch from .exception_translator import translate_py_core_exceptions +if TYPE_CHECKING: + from .._pycore_types import AsyncCoreCursor + from .async_connection import _AsyncConnection + class _AsyncCursor: """Internal Python wrapper over ``mssql_py_core.PyAsyncCursor``. @@ -28,7 +32,9 @@ class _AsyncCursor: Its signatures, behavior, error handling, and compatibility may change without notice. """ - def __init__(self, py_core_async_cursor: Any, connection: Any = None) -> None: + def __init__( + self, py_core_async_cursor: "AsyncCoreCursor", connection: "_AsyncConnection | None" = None + ) -> None: self._py_core_async_cursor = py_core_async_cursor self._connection = connection self._closed = False @@ -82,7 +88,7 @@ def _reset_fetch_tracking(self) -> None: self._fetch_rowcount = None @asynccontextmanager - async def _result_transition(self): + async def _result_transition(self) -> AsyncIterator[None]: async with self._result_transition_lock: self._result_ready.clear() try: @@ -199,7 +205,7 @@ async def close(self) -> None: self._clear_result_metadata() logger.debug("AsyncCursor.close: completed") - def setinputsizes(self, sizes: Any) -> None: + def setinputsizes(self, sizes: Sequence[int | tuple[int, ...]]) -> None: with translate_py_core_exceptions(): self._py_core_async_cursor.setinputsizes(sizes) @@ -209,7 +215,7 @@ def timeout(self) -> int: return self._py_core_async_cursor.timeout @property - def description(self) -> Any: + def description(self) -> list[tuple[Any, ...]] | None: return self._description @property diff --git a/mssql_python/async_query/async_fetch.py b/mssql_python/async_query/async_fetch.py index 404baa5e6..07845036e 100644 --- a/mssql_python/async_query/async_fetch.py +++ b/mssql_python/async_query/async_fetch.py @@ -10,6 +10,7 @@ from .exception_translator import translate_py_core_exceptions if TYPE_CHECKING: + from .._pycore_types import AsyncCoreCursor from .async_cursor import _AsyncCursor # pyright: ignore[reportPrivateUsage] _ResultSnapshot = tuple[ @@ -21,7 +22,7 @@ ] -def _get_py_core_async_cursor(cursor: "_AsyncCursor") -> Any: +def _get_py_core_async_cursor(cursor: "_AsyncCursor") -> "AsyncCoreCursor": return cursor._py_core_async_cursor # pyright: ignore[reportPrivateUsage] diff --git a/mssql_python/async_query/exception_translator.py b/mssql_python/async_query/exception_translator.py index 824650315..5d54260f6 100644 --- a/mssql_python/async_query/exception_translator.py +++ b/mssql_python/async_query/exception_translator.py @@ -4,16 +4,16 @@ from typing import Iterator from ..exceptions import ( - DataError, - DatabaseError, - Error, - IntegrityError, - InterfaceError, - InternalError, - NotSupportedError, - OperationalError, - ProgrammingError, - Warning, + DataError as DataError, + DatabaseError as DatabaseError, + Error as Error, + IntegrityError as IntegrityError, + InterfaceError as InterfaceError, + InternalError as InternalError, + NotSupportedError as NotSupportedError, + OperationalError as OperationalError, + ProgrammingError as ProgrammingError, + Warning as Warning, ) from ..logging import logger @@ -65,7 +65,9 @@ def _translate_known_builtin_error(error: Exception) -> Exception: return error -def _classify_database_error(error: Exception, default_type: type[DatabaseError]): +def _classify_database_error( + error: Exception, default_type: type[DatabaseError] +) -> type[DatabaseError]: diagnostics = getattr(error, "sql_errors", ()) numbers = {item.get("number") for item in diagnostics if isinstance(item, dict)} if numbers & _DATA_ERROR_NUMBERS: @@ -86,7 +88,7 @@ def translate_py_core_exception(error: Exception) -> Exception: if public_type is None: continue if public_type is DatabaseError: - public_type = _classify_database_error(error, public_type) + public_type = _classify_database_error(error, DatabaseError) logger.debug( "Async exception translation: %s -> %s", diff --git a/mssql_python/auth.py b/mssql_python/auth.py index 13a17f94e..8765b0b7a 100644 --- a/mssql_python/auth.py +++ b/mssql_python/auth.py @@ -11,10 +11,12 @@ import sys import threading import time -from typing import Tuple, Dict, NamedTuple, Optional, TYPE_CHECKING +from types import FrameType +from typing import Callable, Tuple, Dict, NamedTuple, Optional, TYPE_CHECKING if TYPE_CHECKING: from azure.core.credentials import TokenCredential + from mssql_python.connection import TokenProvider from mssql_python.logging import logger from mssql_python.constants import ( @@ -43,14 +45,15 @@ # within a single process. Multi-user apps must bring their own per-user token # provider instead of relying on interactive auth. (Documented for users in the # Connection Pooling section of README.md.) -_credential_cache: Dict[object, object] = {} +_CredentialCacheKey = str | tuple[str, tuple[tuple[str, str], ...]] +_credential_cache: Dict[_CredentialCacheKey, "TokenCredential"] = {} _credential_cache_lock = threading.Lock() # Stable home_account_id captured the first time an interactive/device-code # credential runs authenticate(). Lets later *silent* get_token() acquisitions # still key the pool on the account without re-authenticating. Keyed like # _credential_cache and guarded by the same lock. -_account_id_cache: Dict[object, Optional[str]] = {} +_account_id_cache: Dict[_CredentialCacheKey, Optional[str]] = {} # Canonical keys to strip when handing an Entra-token connection to ODBC. _SENSITIVE_KEYS = frozenset({_KEY_UID, _KEY_PWD, _KEY_TRUSTED_CONNECTION, _KEY_AUTHENTICATION}) @@ -124,7 +127,9 @@ def _warn_default_credential_pooling() -> None: } -def _credential_cache_key(auth_type: str, credential_kwargs: Optional[Dict[str, str]]): +def _credential_cache_key( + auth_type: str, credential_kwargs: Optional[Dict[str, str]] +) -> _CredentialCacheKey: """Build a hashable cache key from auth_type and optional credential kwargs. Returns the plain auth_type string when no kwargs are provided so that @@ -193,7 +198,7 @@ def get_raw_token(auth_type: str, credential_kwargs: Optional[Dict[str, str]] = return raw_token @staticmethod - def _authenticate_interactive(credential) -> Optional[str]: + def _authenticate_interactive(credential: "TokenCredential") -> Optional[str]: """Run the interactive ``authenticate()`` step and return the resulting ``home_account_id``. @@ -250,7 +255,7 @@ def _acquire_token( ) from e # Mapping of auth types to credential classes - credential_map = { + credential_map: dict[str, Callable[..., "TokenCredential"]] = { _AuthInternal.DEFAULT: DefaultAzureCredential, _AuthInternal.DEVICE_CODE: DeviceCodeCredential, _AuthInternal.INTERACTIVE: InteractiveBrowserCredential, @@ -401,7 +406,7 @@ class ServicePrincipalAuth: """ @staticmethod - def make_token_factory(client_id: str, client_secret: str): + def make_token_factory(client_id: str, client_secret: str) -> Callable[[str, str, str], bytes]: """Return a callable suitable for ``entra_id_token_factory``. Signature: ``(spn: str, sts_url: str, auth_method: str) -> bytes``. @@ -715,7 +720,7 @@ def _user_facing_stacklevel() -> int: """ # sys._getframe(1) is this helper's caller — the warnings.warn call site, # which corresponds to stacklevel=1. - frame = sys._getframe(1) + frame: FrameType | None = sys._getframe(1) level = 1 while frame is not None: if not frame.f_globals.get("__name__", "").startswith("mssql_python"): @@ -728,7 +733,7 @@ def _user_facing_stacklevel() -> int: def _get_token_from_credential( - credential: "TokenCredential", + credential: "TokenProvider", ) -> Tuple[str, Optional[int]]: """Internal: call credential.get_token() and return ``(raw_jwt, expires_on)``. @@ -858,7 +863,7 @@ def _get_token_from_credential( return raw_token, expires_on -def acquire_token_from_credential(credential: "TokenCredential") -> Tuple[bytes, Optional[int]]: +def acquire_token_from_credential(credential: "TokenProvider") -> Tuple[bytes, Optional[int]]: """Acquire an ODBC token struct from a user-supplied credential object. The credential must follow the Azure ``TokenCredential`` protocol — i.e. @@ -888,7 +893,7 @@ def acquire_token_from_credential(credential: "TokenCredential") -> Tuple[bytes, return AADAuth.get_token_struct(raw_token), expires_on -def acquire_raw_token_from_credential(credential: "TokenCredential") -> Tuple[str, Optional[int]]: +def acquire_raw_token_from_credential(credential: "TokenProvider") -> Tuple[str, Optional[int]]: """Acquire a raw JWT string from a user-supplied credential object. Used by bulk copy, which needs the raw JWT rather than the ODBC struct. diff --git a/mssql_python/connection.py b/mssql_python/connection.py index 1984f4979..98250ef7e 100644 --- a/mssql_python/connection.py +++ b/mssql_python/connection.py @@ -11,6 +11,8 @@ - Cursors are also cleaned up automatically when no longer referenced, to prevent memory leaks. """ +from __future__ import annotations + import weakref import re import codecs @@ -18,6 +20,7 @@ import struct from types import MappingProxyType from typing import Any, Dict, Optional, Union, List, Tuple, Callable, Protocol, TYPE_CHECKING +from typing import NoReturn, cast import threading import mssql_python @@ -65,7 +68,9 @@ ) if TYPE_CHECKING: - from mssql_python.row import Row + from mssql_python.row import Row, OutputConverter + from mssql_python.auth import TokenInfo + from mssql_python._ddbc_types import EncodingSettings class TokenProvider(Protocol): @@ -212,7 +217,7 @@ def get_token(self, scope: str) -> Any: _SQLSTATE_RE = re.compile(r"^SQLSTATE:([A-Z0-9]{0,5}):(.*)", re.DOTALL) -def _raise_connection_error(e: RuntimeError) -> None: +def _raise_connection_error(e: RuntimeError) -> NoReturn: """Map a RuntimeError from the C++ pybind layer to the correct DB-API 2.0 exception. Connection::checkError() throws "SQLSTATE:XXXXX:" so the SQLSTATE @@ -516,7 +521,7 @@ def __init__( # Initialize encoding settings with defaults for Python 3 # Python 3 only has str (which is Unicode), so we use utf-16le by default - self._encoding_settings = { + self._encoding_settings: EncodingSettings = { "encoding": "utf-16le", "ctype": ConstantsDDBC.SQL_WCHAR.value, } @@ -526,7 +531,7 @@ def __init__( # UTF-16 data for VARCHAR columns. This avoids encoding mismatches on # Windows where the driver returns raw bytes in the server's native # code page (e.g. CP-1252) that may fail to decode as UTF-8. - self._decoding_settings = { + self._decoding_settings: dict[int, EncodingSettings] = { ConstantsDDBC.SQL_CHAR.value: { "encoding": "utf-16le", "ctype": ConstantsDDBC.SQL_WCHAR.value, @@ -627,7 +632,7 @@ def __init__( token_attr = ConstantsDDBC.SQL_COPT_SS_ACCESS_TOKEN.value base_attrs = self._attrs_before - def _acquire_token_info(): + def _acquire_token_info() -> TokenInfo | None: # DB-API boundary: get_auth_token_info fails closed by # letting the underlying Azure error propagate as a # ValueError (unsupported auth type) or RuntimeError @@ -646,7 +651,9 @@ def _acquire_token_info(): ddbc_error=str(e), ) from e - def _make_token_factory(expected_account: Optional[str] = None): + def _make_token_factory( + expected_account: Optional[str] = None, + ) -> Callable[[], tuple[dict[int, int | str | bytes], int | None]]: # Build the deferred connect-attrs provider handed to native. # Native invokes the returned callable only when it actually # opens a physical connection (a pool miss, non-pooled @@ -661,7 +668,7 @@ def _make_token_factory(expected_account: Optional[str] = None): # pool would hand a caller a connection authenticated as the # wrong account. MSI pools pass ``None`` (their identity is # fixed by params, not by a mutable signed-in account). - def _token_factory(): + def _token_factory() -> tuple[dict[int, int | str | bytes], int | None]: attrs = dict(base_attrs) info = _acquire_token_info() if info and info.token_struct: @@ -779,10 +786,10 @@ def _token_factory(): # when no longer in use without requiring explicit deletion. # TODO: Think and implement scenarios for multi-threaded access # to cursors - self._cursors = weakref.WeakSet() + self._cursors: weakref.WeakSet[Cursor] = weakref.WeakSet() # Initialize output converters dictionary and its lock for thread safety - self._output_converters = {} + self._output_converters: dict[int | type, OutputConverter] = {} self._converters_generation = 0 self._converters_lock = threading.Lock() @@ -796,7 +803,7 @@ def _token_factory(): self._encoding_lock = threading.Lock() # Initialize search escape character - self._searchescape = None + self._searchescape: str | None = None # Safety net for the raw access-token pattern: a caller may pass # SQL_COPT_SS_ACCESS_TOKEN directly in attrs_before with no @@ -859,7 +866,7 @@ def _token_factory(): ddbc_bindings._set_odbc_provider(_provider) try: - self._conn = ddbc_bindings.Connection( + self._conn: ddbc_bindings.Connection | None = ddbc_bindings.Connection( self.connection_str, self._pooling, self._attrs_before, @@ -1133,6 +1140,8 @@ def autocommit(self) -> bool: Returns: bool: True if autocommit is enabled, False otherwise. """ + if self._conn is None: + raise InterfaceError("Connection is closed", "Connection is closed") try: return self._conn.get_autocommit() except RuntimeError as e: @@ -1176,6 +1185,8 @@ def setautocommit(self, value: bool = False) -> None: Raises: DatabaseError: If there is an error while setting the autocommit mode. """ + if self._conn is None: + raise InterfaceError("Connection is closed", "Connection is closed") try: self._conn.set_autocommit(value) except RuntimeError as e: @@ -1287,7 +1298,7 @@ def setencoding(self, encoding: Optional[str] = None, ctype: Optional[int] = Non sanitize_user_input(str(ctype)), ) - def getencoding(self) -> Dict[str, Union[str, int]]: + def getencoding(self) -> EncodingSettings: """ Gets the current text encoding settings (thread-safe). @@ -1464,7 +1475,7 @@ def setdecoding( sanitize_user_input(str(ctype)), ) - def getdecoding(self, sqltype: int) -> Dict[str, Union[str, int]]: + def getdecoding(self, sqltype: int) -> EncodingSettings: """ Gets the current text decoding settings for the specified SQL type (thread-safe). @@ -1543,7 +1554,7 @@ def set_attr(self, attribute: int, value: Union[int, str, bytes, bytearray]) -> must be provided in the attrs_before parameter when creating the connection. Attempting to set these attributes after connection will raise a ProgrammingError. """ - if self._closed: + if self._closed or self._conn is None: raise InterfaceError( "Cannot set attribute on closed connection", "Connection is closed" ) @@ -1561,7 +1572,7 @@ def set_attr(self, attribute: int, value: Union[int, str, bytes, bytearray]) -> ) raise ProgrammingError( driver_error=f"Invalid attribute or value: {error_message}", - ddbc_error=error_message, + ddbc_error=str(error_message), ) # Log with sanitized values @@ -1697,7 +1708,7 @@ def add_output_converter(self, sqltype: Union[int, type], func: Callable[[Any], self._output_converters[sqltype] = func self._converters_generation += 1 # Pass to the underlying connection if native implementation supports it - if hasattr(self._conn, "add_output_converter"): + if self._conn is not None and hasattr(self._conn, "add_output_converter"): self._conn.add_output_converter(sqltype, func) logger.info(f"Added output converter for SQL type {sqltype}") @@ -1739,7 +1750,7 @@ def remove_output_converter(self, sqltype: Union[int, type]) -> None: del self._output_converters[sqltype] self._converters_generation += 1 # Pass to the underlying connection if native implementation supports it - if hasattr(self._conn, "remove_output_converter"): + if self._conn is not None and hasattr(self._conn, "remove_output_converter"): self._conn.remove_output_converter(sqltype) logger.info(f"Removed output converter for SQL type {sqltype}") @@ -1758,7 +1769,7 @@ def clear_output_converters(self) -> None: self._output_converters.clear() self._converters_generation += 1 # Pass to the underlying connection if native implementation supports it - if hasattr(self._conn, "clear_output_converters"): + if self._conn is not None and hasattr(self._conn, "clear_output_converters"): self._conn.clear_output_converters() logger.info("Cleared all output converters") @@ -1888,10 +1899,10 @@ def batch_execute( # Determine which cursor to use is_new_cursor = reuse_cursor is None - cursor = self.cursor() if is_new_cursor else reuse_cursor + cursor = self.cursor() if reuse_cursor is None else reuse_cursor # Execute statements and collect results - results = [] + results: list[list[Row] | int] = [] try: for i, (stmt, param) in enumerate(zip(statements, params)): try: @@ -1945,7 +1956,7 @@ def batch_execute( return results, cursor - def getinfo(self, info_type: int) -> Union[str, int, bool, None]: + def getinfo(self, info_type: int) -> str | int | bool | bytes | None: """ Return general information about the driver and data source. @@ -1957,7 +1968,8 @@ def getinfo(self, info_type: int) -> Union[str, int, bool, None]: The requested information. The type of the returned value depends on the information requested. For registered ODBC types, character values (including "Y"/"N") return strings; numeric values and bitmasks return unsigned - integers. Native retrieval failures, including unsupported types, + integers. Unregistered, driver-specific types can return raw bytes. + Native retrieval failures, including unsupported types, timeouts, and connection loss, are logged and return None. Note: @@ -1973,7 +1985,7 @@ def getinfo(self, info_type: int) -> Union[str, int, bool, None]: DatabaseError: If a numeric byte result does not match its ODBC type's width. InterfaceError: If the connection is closed. """ - if self._closed: + if self._closed or self._conn is None: raise InterfaceError( driver_error="Cannot get info on closed connection", ddbc_error="Cannot get info on closed connection", @@ -2060,7 +2072,7 @@ def getinfo(self, info_type: int) -> Union[str, int, bool, None]: f"got length={length} with {len(data)} bytes of data" ), ) - return return_type.unpack_from(data)[0] + return cast(int, return_type.unpack_from(data)[0]) # Legacy non-byte payloads must not lose precision or change bool to int. if isinstance(data, str) and data.isdecimal(): try: diff --git a/mssql_python/constants.py b/mssql_python/constants.py index 54a51b9ce..f9dfa69da 100644 --- a/mssql_python/constants.py +++ b/mssql_python/constants.py @@ -5,7 +5,7 @@ """ from enum import Enum -from typing import Dict, Optional, Tuple +from typing import Dict, Optional, Tuple, TYPE_CHECKING class ConstantsDDBC(Enum): @@ -384,7 +384,7 @@ class SQLTypes: """Constants for valid SQL data types to use with setinputsizes""" @classmethod - def get_valid_types(cls) -> set: + def get_valid_types(cls) -> set[int]: """Returns a set of all valid SQL type constants""" return { @@ -423,7 +423,7 @@ def get_valid_types(cls) -> set: # Could also add category methods for convenience @classmethod - def get_string_types(cls) -> set: + def get_string_types(cls) -> set[int]: """Returns a set of string SQL type constants""" return { @@ -436,7 +436,7 @@ def get_string_types(cls) -> set: } @classmethod - def get_numeric_types(cls) -> set: + def get_numeric_types(cls) -> set[int]: """Returns a set of numeric SQL type constants""" return { @@ -492,7 +492,7 @@ class AttributeSetTime(Enum): } -def get_attribute_set_timing(attribute): +def get_attribute_set_timing(attribute: int) -> AttributeSetTime: """ Get when an attribute can be set (before connection, after, or either). @@ -658,6 +658,168 @@ def get_info_constants() -> Dict[str, int]: "SQL_MODE_READ_ONLY", } +# Declarations for the integer aliases populated dynamically below. +if TYPE_CHECKING: + SQL_SMALLINT: int + SQL_CHAR: int + SQL_WCHAR: int + SQL_WVARCHAR: int + SQL_BIT: int + SQL_TINYINT: int + SQL_BIGINT: int + SQL_BINARY: int + SQL_VARBINARY: int + SQL_LONGVARBINARY: int + SQL_LONGVARCHAR: int + SQL_NUMERIC: int + SQL_DECIMAL: int + SQL_INTEGER: int + SQL_FLOAT: int + SQL_REAL: int + SQL_DOUBLE: int + SQL_TIMESTAMP: int + SQL_DATE: int + SQL_TIME: int + SQL_VARCHAR: int + SQL_TYPE_DATE: int + SQL_TYPE_TIME: int + SQL_TYPE_TIMESTAMP: int + SQL_GUID: int + SQL_XML: int + SQL_WLONGVARCHAR: int + SQL_SS_TIME2: int + SQL_SS_XML: int + SQL_SS_VARIANT: int + SQL_ATTR_ACCESS_MODE: int + SQL_ATTR_CONNECTION_TIMEOUT: int + SQL_ATTR_CURRENT_CATALOG: int + SQL_ATTR_LOGIN_TIMEOUT: int + SQL_ATTR_PACKET_SIZE: int + SQL_ATTR_TXN_ISOLATION: int + SQL_TXN_ISOLATION_LEVEL: int + SQL_TXN_READ_UNCOMMITTED: int + SQL_TXN_READ_COMMITTED: int + SQL_TXN_REPEATABLE_READ: int + SQL_TXN_SERIALIZABLE: int + SQL_MODE_READ_WRITE: int + SQL_MODE_READ_ONLY: int + SQL_CONCURRENCY: int + SQL_ROWSET_SIZE: int + SQL_ROW_NUMBER: int + SQL_IC_UPPER: int + SQL_IC_LOWER: int + SQL_IC_SENSITIVE: int + SQL_IC_MIXED: int + SQL_SC_SQL92_ENTRY: int + SQL_SC_FIPS127_2_TRANSITIONAL: int + SQL_SC_SQL92_INTERMEDIATE: int + SQL_SC_SQL92_FULL: int + SQL_SQL92_ENTRY_SQL: int + SQL_SQL92_INTERMEDIATE_SQL: int + SQL_SQL92_FULL_SQL: int + SQL_DRIVER_NAME: int + SQL_DRIVER_VER: int + SQL_DRIVER_ODBC_VER: int + SQL_DRIVER_HLIB: int + SQL_DRIVER_HENV: int + SQL_DRIVER_HDBC: int + SQL_DATA_SOURCE_NAME: int + SQL_DATABASE_NAME: int + SQL_SERVER_NAME: int + SQL_USER_NAME: int + SQL_SQL_CONFORMANCE: int + SQL_KEYWORDS: int + SQL_IDENTIFIER_CASE: int + SQL_IDENTIFIER_QUOTE_CHAR: int + SQL_SPECIAL_CHARACTERS: int + SQL_SUBQUERIES: int + SQL_EXPRESSIONS_IN_ORDERBY: int + SQL_CORRELATION_NAME: int + SQL_SEARCH_PATTERN_ESCAPE: int + SQL_CATALOG_TERM: int + SQL_CATALOG_NAME_SEPARATOR: int + SQL_SCHEMA_TERM: int + SQL_TABLE_TERM: int + SQL_PROCEDURES: int + SQL_ACCESSIBLE_TABLES: int + SQL_ACCESSIBLE_PROCEDURES: int + SQL_CATALOG_NAME: int + SQL_CATALOG_USAGE: int + SQL_SCHEMA_USAGE: int + SQL_COLUMN_ALIAS: int + SQL_DESCRIBE_PARAMETER: int + SQL_TXN_CAPABLE: int + SQL_TXN_ISOLATION_OPTION: int + SQL_DEFAULT_TXN_ISOLATION: int + SQL_MULTIPLE_ACTIVE_TXN: int + SQL_NUMERIC_FUNCTIONS: int + SQL_STRING_FUNCTIONS: int + SQL_DATETIME_FUNCTIONS: int + SQL_TIMEDATE_FUNCTIONS: int + SQL_SYSTEM_FUNCTIONS: int + SQL_CONVERT_FUNCTIONS: int + SQL_LIKE_ESCAPE_CLAUSE: int + SQL_MAX_COLUMN_NAME_LEN: int + SQL_MAX_TABLE_NAME_LEN: int + SQL_MAX_SCHEMA_NAME_LEN: int + SQL_MAX_CATALOG_NAME_LEN: int + SQL_MAX_IDENTIFIER_LEN: int + SQL_MAX_STATEMENT_LEN: int + SQL_MAX_CHAR_LITERAL_LEN: int + SQL_MAX_BINARY_LITERAL_LEN: int + SQL_MAX_COLUMNS_IN_TABLE: int + SQL_MAX_COLUMNS_IN_SELECT: int + SQL_MAX_COLUMNS_IN_GROUP_BY: int + SQL_MAX_COLUMNS_IN_ORDER_BY: int + SQL_MAX_COLUMNS_IN_INDEX: int + SQL_MAX_TABLES_IN_SELECT: int + SQL_MAX_CONCURRENT_ACTIVITIES: int + SQL_MAX_DRIVER_CONNECTIONS: int + SQL_MAX_ROW_SIZE: int + SQL_MAX_USER_NAME_LEN: int + SQL_ACTIVE_CONNECTIONS: int + SQL_ACTIVE_STATEMENTS: int + SQL_DATA_SOURCE_READ_ONLY: int + SQL_NEED_LONG_DATA_LEN: int + SQL_GETDATA_EXTENSIONS: int + SQL_CURSOR_COMMIT_BEHAVIOR: int + SQL_CURSOR_ROLLBACK_BEHAVIOR: int + SQL_CURSOR_SENSITIVITY: int + SQL_BOOKMARK_PERSISTENCE: int + SQL_DYNAMIC_CURSOR_ATTRIBUTES1: int + SQL_DYNAMIC_CURSOR_ATTRIBUTES2: int + SQL_FORWARD_ONLY_CURSOR_ATTRIBUTES1: int + SQL_FORWARD_ONLY_CURSOR_ATTRIBUTES2: int + SQL_STATIC_CURSOR_ATTRIBUTES1: int + SQL_STATIC_CURSOR_ATTRIBUTES2: int + SQL_KEYSET_CURSOR_ATTRIBUTES1: int + SQL_KEYSET_CURSOR_ATTRIBUTES2: int + SQL_SCROLL_OPTIONS: int + SQL_SCROLL_CONCURRENCY: int + SQL_FETCH_DIRECTION: int + SQL_STATIC_SENSITIVITY: int + SQL_BATCH_SUPPORT: int + SQL_BATCH_ROW_COUNT: int + SQL_PARAM_ARRAY_ROW_COUNTS: int + SQL_PARAM_ARRAY_SELECTS: int + SQL_PROCEDURE_TERM: int + SQL_POSITIONED_STATEMENTS: int + SQL_GROUP_BY: int + SQL_OJ_CAPABILITIES: int + SQL_ORDER_BY_COLUMNS_IN_SELECT: int + SQL_OUTER_JOINS: int + SQL_QUOTED_IDENTIFIER_CASE: int + SQL_CONCAT_NULL_BEHAVIOR: int + SQL_NULL_COLLATION: int + SQL_ALTER_TABLE: int + SQL_UNION: int + SQL_DDL_INDEX: int + SQL_MULT_RESULT_SETS: int + SQL_OWNER_USAGE: int + SQL_QUALIFIER_USAGE: int + SQL_TIMEDATE_ADD_INTERVALS: int + SQL_TIMEDATE_DIFF_INTERVALS: int + # Get current module's globals for dynamic export _module_globals = globals() _exported_names = [] @@ -669,9 +831,9 @@ def get_info_constants() -> Dict[str, int]: _exported_names.append(_name) # Export GetInfoConstants members not already exported from ConstantsDDBC. -for _name, _member in GetInfoConstants.__members__.items(): +for _name, _info_member in GetInfoConstants.__members__.items(): if _name not in _DDBC_PUBLIC_API: - _module_globals[_name] = _member.value + _module_globals[_name] = _info_member.value _exported_names.append(_name) # AuthType enum is exported as a class only (not individual members) @@ -697,4 +859,4 @@ def get_info_constants() -> Dict[str, int]: ] # Clean up temporary variables -del _module_globals, _exported_names, _name, _member +del _module_globals, _exported_names, _name, _member, _info_member diff --git a/mssql_python/cursor.py b/mssql_python/cursor.py index 0825ea1b5..28a3c2085 100644 --- a/mssql_python/cursor.py +++ b/mssql_python/cursor.py @@ -11,12 +11,16 @@ # pylint: disable=too-many-lines # Large file due to comprehensive DB-API 2.0 implementation +from __future__ import annotations + import decimal import logging import uuid import datetime import warnings from typing import List, Mapping, Union, Any, Optional, Tuple, Sequence, TYPE_CHECKING, Iterable +from typing import ClassVar, Generator, Iterator, Literal, Protocol, cast +from importlib import import_module from mssql_python.constants import ConstantsDDBC as ddbc_sql_const, SQLTypes from mssql_python.helpers import check_error, connstr_to_pycore_params from mssql_python.logging import logger @@ -28,7 +32,7 @@ OperationalError, DatabaseError, ) -from mssql_python.row import Row +from mssql_python.row import Row, OutputConverter from mssql_python.perf_timer import perf_phase from mssql_python import get_settings from mssql_python.parameter_helper import ( @@ -38,8 +42,23 @@ ) if TYPE_CHECKING: - import pyarrow # type: ignore + import pyarrow + from mssql_python._ddbc_types import ColumnMetadata, EncodingSettings + from mssql_python._pycore_types import ( + BulkCopyResult, + CoreContext, + CoreConnection, + CoreCursor, + PyCoreModule, + ) from mssql_python.connection import Connection + + class _ArrowModule(Protocol): + RecordBatch: type[pyarrow.RecordBatch] + RecordBatchReader: type[pyarrow.RecordBatchReader] + Table: type[pyarrow.Table] + ArrowInvalid: type[Exception] + else: pyarrow = None @@ -62,10 +81,13 @@ } -def _string_only_output_converter(converter): +_Description = list[tuple[str, type, int | None, int | None, int | None, int | None, bool | None]] + + +def _string_only_output_converter(converter: OutputConverter) -> OutputConverter: """Gate fallback conversion on the fetched value, not its column metadata.""" - def convert(value): + def convert(value: Any) -> Any: if isinstance(value, (str, bytes)): return converter(value) return value @@ -73,7 +95,7 @@ def convert(value): return convert -def _normalize_time_param(value, c_type): +def _normalize_time_param(value: object, c_type: int) -> str | None: """Convert a datetime.time to its isoformat string when bound via text C-types. Returns the isoformat string if conversion applies, otherwise *None*. @@ -151,13 +173,13 @@ def __init__( self, cursor: "Cursor", inner: "pyarrow.RecordBatchReader", - generator, - arrow_invalid_exc: type, + generator: Generator[pyarrow.RecordBatch, None, None], + arrow_invalid_exc: type[Exception], close_requested: list[bool], ) -> None: - self._cursor = cursor - self._inner = inner - self._generator = generator + self._cursor: Cursor | None = cursor + self._inner: pyarrow.RecordBatchReader | None = inner + self._generator: Generator[pyarrow.RecordBatch, None, None] | None = generator self._closed = False self._close_requested = close_requested # Cache the exception class so post-close reads in a hot loop don't @@ -171,7 +193,7 @@ def closed(self) -> bool: """True once ``close()`` has been called.""" return self._closed - def __getattr__(self, name): + def __getattr__(self, name: str) -> Any: """Delegate any attribute we don't explicitly define to the inner ``pyarrow.RecordBatchReader``. @@ -199,7 +221,7 @@ def __getattr__(self, name): raise self._arrow_invalid("Reader is closed") return getattr(self._inner, name) - def __arrow_c_stream__(self, requested_schema=None): + def __arrow_c_stream__(self, requested_schema: object = None) -> object: """Arrow PyCapsule Protocol — export as an Arrow C stream. Implements the Arrow PyCapsule Protocol for streams (pyarrow >= 14), @@ -230,24 +252,24 @@ def __arrow_c_stream__(self, requested_schema=None): ) return inner_export(requested_schema) - def __iter__(self): + def __iter__(self) -> _ArrowReader: return self - def __next__(self): - if self._closed: + def __next__(self) -> pyarrow.RecordBatch: + if self._closed or self._inner is None: raise self._arrow_invalid("Reader is closed") return self._inner.read_next_batch() - def __enter__(self): + def __enter__(self) -> _ArrowReader: if self._closed: raise self._arrow_invalid("Reader is closed") return self - def __exit__(self, exc_type, exc_val, exc_tb): + def __exit__(self, exc_type: object, exc_val: object, exc_tb: object) -> Literal[False]: self.close() return False - def __del__(self): + def __del__(self) -> None: # Best-effort cleanup if the user never called close() (or a previous # close() attempt failed to release the generator and left cleanup # incomplete) and the reader is being garbage-collected. Skip during @@ -385,26 +407,14 @@ def __init__(self, connection: "Connection", timeout: int = 0) -> None: # ``hstmt=None`` keeps ``close()`` safe when allocation fails before # an HSTMT exists. self.closed: bool = False - self.hstmt: Optional[Any] = None + self.hstmt: ddbc_bindings.SqlHandle | None = None self._connection: "Connection" = connection # Store as private attribute self._timeout: int = timeout self._inputsizes: Optional[List[Tuple[int, int, int, int]]] = None # self.connection.autocommit = False self._initialize_cursor() - self.description: Optional[ - List[ - Tuple[ - str, - Any, - Optional[int], - Optional[int], - Optional[int], - Optional[int], - Optional[bool], - ] - ] - ] = None + self.description: _Description | None = None self.rowcount: int = -1 self.arraysize: int = ( 1 # Default number of rows to fetch at a time is 1, user can change it @@ -424,22 +434,22 @@ def __init__(self, connection: "Connection", timeout: int = 0) -> None: # Column-name -> index map for the current result set. For catalog/metadata # result sets this also carries lowercase and friendly aliases (see # _prepare_metadata_result_set). - self._cached_column_map = None - self._cached_column_map_lower = None - self._cached_converter_map = None + self._cached_column_map: dict[str, int] | None = None + self._cached_column_map_lower: dict[str, int] | None = None + self._cached_converter_map: Sequence[OutputConverter | None] | None = None self._cached_converters_generation = self._connection._converters_generation # Canonical, order-preserving column names snapshotted once per result set # and handed to each Row so mapping views never read the live cursor.description # (which changes when the cursor is reused for another query). _result_columns_src # tracks the self.description identity the snapshot was built from, so a new # result set rebuilds it exactly once. - self._cached_result_columns = None - self._result_columns_src = None + self._cached_result_columns: tuple[str, ...] | None = None + self._result_columns_src: _Description | None = None # Raw ODBC SQL type codes (from SQLDescribeCol) per column, parallel to # self.description. Kept so output-converter dispatch can key on the integer # ODBC SQL type code (pyodbc-compatible), not just the mapped Python type. See #684. - self._column_sql_types = None - self._uuid_str_indices = None # Pre-computed UUID column indices for str conversion + self._column_sql_types: list[int] | None = None + self._uuid_str_indices: tuple[int, ...] | None = None # Cache the effective native_uuid setting for this cursor's connection. # Resolution order: connection._native_uuid (if not None) → module-level setting. self._conn_native_uuid = getattr(self.connection, "_native_uuid", None) @@ -526,7 +536,7 @@ def _parse_time(self, param: str) -> Optional[datetime.time]: continue return None - def _get_numeric_data(self, param: decimal.Decimal) -> Any: + def _get_numeric_data(self, param: decimal.Decimal) -> ddbc_bindings.NumericData: """ Get the data for a numeric parameter. @@ -601,7 +611,7 @@ def _get_numeric_data(self, param: decimal.Decimal) -> Any: numeric_data.val = bytes(byte_array) return numeric_data - def _get_encoding_settings(self): + def _get_encoding_settings(self) -> EncodingSettings: """ Get the encoding settings from the connection. @@ -636,7 +646,7 @@ def _get_encoding_settings(self): # This is the only case where defaults are appropriate (method doesn't exist) return {"encoding": "utf-16le", "ctype": ddbc_sql_const.SQL_WCHAR.value} - def _refresh_decoding_cache(self): + def _refresh_decoding_cache(self) -> None: """Read decoding settings only when the connection configuration changes.""" generation = self._connection._decoding_generation char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) @@ -648,7 +658,7 @@ def _refresh_decoding_cache(self): self._cached_wchar_encoding = wchar_encoding self._cached_decoding_generation = generation - def _get_decoding_settings(self, sql_type): + def _get_decoding_settings(self, sql_type: int) -> EncodingSettings: """ Get decoding settings for a specific SQL type. @@ -1032,7 +1042,12 @@ def _allocate_statement_handle(self) -> None: """ Allocate the DDBC statement handle. """ - self.hstmt = self._connection._conn.alloc_statement_handle() + native_connection = self._connection._conn + if native_connection is None: + raise InterfaceError( + "Cannot create cursor on closed connection", "Connection is closed" + ) + self.hstmt = native_connection.alloc_statement_handle() def _set_timeout(self) -> None: """ @@ -1157,20 +1172,20 @@ def _capture_diagnostics(self, ret: int) -> None: ): self.messages.extend(ddbc_bindings.DDBCSQLGetAllDiagRecords(self.hstmt)) - def _ensure_pyarrow(self) -> Any: + def _ensure_pyarrow(self) -> _ArrowModule: """ Import and return pyarrow or raise ImportError accordingly. """ try: import pyarrow - return pyarrow + return cast("_ArrowModule", pyarrow) except ImportError as e: raise ImportError( "pyarrow is required for Arrow fetch methods. Please install pyarrow." ) from e - def setinputsizes(self, sizes: List[Union[int, tuple]]) -> None: + def setinputsizes(self, sizes: Sequence[int | tuple[int, ...]]) -> None: """ Sets the type information to be used for parameters in execute and executemany. @@ -1279,10 +1294,10 @@ def _reset_inputsizes(self) -> None: # Pre-built constant lookup table — avoids rebuilding ~30 entries on every call. # Used by setinputsizes fallback path (PR #549 fast path doesn't need this). - _SQL_TO_C_TYPE = None + _SQL_TO_C_TYPE: ClassVar[dict[int, int] | None] = None @classmethod - def _get_sql_to_c_type_map(cls): + def _get_sql_to_c_type_map(cls) -> dict[int, int]: if cls._SQL_TO_C_TYPE is None: cls._SQL_TO_C_TYPE = { ddbc_sql_const.SQL_CHAR.value: ddbc_sql_const.SQL_C_CHAR.value, @@ -1326,13 +1341,13 @@ def _get_c_type_for_sql_type(self, sql_type: int) -> int: def _create_parameter_types_list( # pylint: disable=too-many-arguments,too-many-positional-arguments self, parameter: Any, - param_info: Optional[Tuple[Any, ...]], + param_info: type[ddbc_bindings.ParamInfo], parameters_list: List[Any], i: int, min_val: Optional[Any] = None, max_val: Optional[Any] = None, decimal_as_numeric: bool = False, - ) -> Tuple[int, int, int, int, bool]: + ) -> ddbc_bindings.ParamInfo: """ Maps parameter types for the given parameter. @@ -1375,14 +1390,14 @@ def _create_parameter_types_list( # pylint: disable=too-many-arguments,too-many return paraminfo - def _initialize_description(self, column_metadata: Optional[Any] = None) -> None: + def _initialize_description(self, column_metadata: list[ColumnMetadata] | None = None) -> None: """Initialize the description attribute from column metadata.""" if not column_metadata: self.description = None self._column_sql_types = None return - description = [] + description: _Description = [] # Raw ODBC SQL type codes, parallel to description, for output-converter # dispatch by integer SQL type (see _build_converter_map / #684). sql_type_codes = [] @@ -1410,7 +1425,7 @@ def _initialize_description(self, column_metadata: Optional[Any] = None) -> None self.description = description self._column_sql_types = sql_type_codes - def _build_converter_map(self): + def _build_converter_map(self) -> Sequence[OutputConverter | None]: """ Build a pre-computed converter map for output converters. Returns a list where each element is either a converter function or None. @@ -1430,7 +1445,7 @@ def _build_converter_map(self): return () sql_type_codes = self._column_sql_types - converter_map = [] + converter_map: list[OutputConverter | None] = [] for i, desc in enumerate(self.description): if desc is None: @@ -1457,7 +1472,7 @@ def _build_converter_map(self): self._cached_converters_generation = generation return converter_map if any(converter is not None for converter in converter_map) else () - def _compute_uuid_str_indices(self): + def _compute_uuid_str_indices(self) -> tuple[int, ...] | None: """ Compute the tuple of column indices whose uuid.UUID values should be stringified (as uppercase), based on the effective native_uuid setting. @@ -1483,7 +1498,11 @@ def _compute_uuid_str_indices(self): return indices if indices else None return None - def _get_column_and_converter_maps(self): + def _get_column_and_converter_maps( + self, + ) -> tuple[ + dict[str, int] | None, Sequence[OutputConverter | None] | None, dict[str, int] | None + ]: """ Get column map and converter map for Row construction (thread-safe). This centralizes the column map building logic to eliminate duplication @@ -1518,7 +1537,7 @@ def _get_column_and_converter_maps(self): return column_map, converter_map, self._cached_column_map_lower - def _map_data_type(self, sql_type): + def _map_data_type(self, sql_type: int) -> type: """ Map SQL data type to Python data type. @@ -1644,7 +1663,7 @@ def _reset_rownumber(self) -> None: self._has_result_set = True self._skip_increment_for_next_fetch = False - def _increment_rownumber(self): + def _increment_rownumber(self) -> None: """ Called after a successful fetch from the driver. Keep both counters consistent. """ @@ -1660,7 +1679,7 @@ def _increment_rownumber(self): ) # Will be used when we add support for scrollable cursors - def _decrement_rownumber(self): + def _decrement_rownumber(self) -> None: """ Decrement the rownumber by 1. @@ -1677,7 +1696,7 @@ def _decrement_rownumber(self): "No active result set.", ) - def _clear_rownumber(self): + def _clear_rownumber(self) -> None: """ Clear the rownumber tracking. @@ -1687,7 +1706,7 @@ def _clear_rownumber(self): self._has_result_set = False self._skip_increment_for_next_fetch = False - def __iter__(self): + def __iter__(self) -> Cursor: """ Return the cursor itself as an iterator. @@ -1699,7 +1718,7 @@ def __iter__(self): self._check_closed() return self - def __next__(self): + def __next__(self) -> Row: """ Fetch the next row when iterating over the cursor. @@ -1715,7 +1734,7 @@ def __next__(self): raise StopIteration return row - def next(self): + def next(self) -> Row: """ Fetch the next row from the cursor. @@ -1732,7 +1751,7 @@ def next(self): def execute( # pylint: disable=too-many-locals,too-many-branches,too-many-statements self, operation: str, - *parameters, + *parameters: Any, use_prepare: bool = True, reset_cursor: bool = True, ) -> "Cursor": @@ -1824,25 +1843,25 @@ def execute( # pylint: disable=too-many-locals,too-many-branches,too-many-state if operation == self.last_executed_stmt and isinstance( actual_params, (tuple, list) ): - parameters = list(actual_params) + bound_parameters = list(actual_params) else: operation, converted_params = detect_and_convert_parameters( operation, actual_params ) - parameters = list(converted_params) + bound_parameters = list(converted_params) else: - parameters = [] + bound_parameters = [] # Getting encoding setting encoding_settings = self._get_encoding_settings() # Validate that inputsizes matches parameter count if both are present - if parameters and self._inputsizes: - if len(self._inputsizes) != len(parameters): + if bound_parameters and self._inputsizes: + if len(self._inputsizes) != len(bound_parameters): warnings.warn( f"Number of input sizes ({len(self._inputsizes)}) does not match " - f"number of parameters ({len(parameters)}). " + f"number of parameters ({len(bound_parameters)}). " f"This may lead to unexpected behavior.", Warning, ) @@ -1850,17 +1869,19 @@ def execute( # pylint: disable=too-many-locals,too-many-branches,too-many-state # Prepare caching: skip SQLPrepare when re-executing the same SQL # with parameters. The HSTMT is reused via _soft_reset_cursor, so the # server-side plan from the previous SQLPrepare is still valid. - same_sql = parameters and operation == self.last_executed_stmt and self.is_stmt_prepared[0] + same_sql = ( + bound_parameters and operation == self.last_executed_stmt and self.is_stmt_prepared[0] + ) if not same_sql: self.is_stmt_prepared = [False] effective_use_prepare = use_prepare and not same_sql with perf_phase("py::execute::cpp_call"): - if parameters: + if bound_parameters: ret = ddbc_bindings.DDBCSQLExecute( self.hstmt, operation, - parameters, + bound_parameters, self._inputsizes, self.is_stmt_prepared, effective_use_prepare, @@ -1891,7 +1912,7 @@ def execute( # pylint: disable=too-many-locals,too-many-branches,too-many-state # Initialize description after execution # After successful execution, initialize description if there are results - column_metadata = [] + column_metadata: list[ColumnMetadata] = [] try: ddbc_bindings.DDBCSQLDescribeCol(self.hstmt, column_metadata) self._initialize_description(column_metadata) @@ -1928,8 +1949,11 @@ def execute( # pylint: disable=too-many-locals,too-many-branches,too-many-state return self def _prepare_metadata_result_set( # pylint: disable=too-many-statements - self, column_metadata=None, fallback_description=None, specialized_mapping=None - ): + self, + column_metadata: list[ColumnMetadata] | None = None, + fallback_description: _Description | None = None, + specialized_mapping: Mapping[str, int] | None = None, + ) -> Cursor: """ Prepares a metadata result set by: 1. Retrieving column metadata if not provided @@ -2016,7 +2040,7 @@ def _prepare_metadata_result_set( # pylint: disable=too-many-statements # Return the cursor itself for method chaining return self - def getTypeInfo(self, sqlType=None): + def getTypeInfo(self, sqlType: int | None = None) -> Cursor: """ Executes SQLGetTypeInfo and creates a result set with information about the specified data type or all data types supported by the ODBC driver if not specified. @@ -2039,7 +2063,9 @@ def getTypeInfo(self, sqlType=None): self._reset_cursor() raise e - def procedures(self, procedure=None, catalog=None, schema=None): + def procedures( + self, procedure: str | None = None, catalog: str | None = None, schema: str | None = None + ) -> Cursor: """ Executes SQLProcedures and creates a result set of information about procedures in the data source. @@ -2057,7 +2083,7 @@ def procedures(self, procedure=None, catalog=None, schema=None): check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Define fallback description for procedures - fallback_description = [ + fallback_description: _Description = [ ("procedure_cat", str, None, 128, 128, 0, True), ("procedure_schem", str, None, 128, 128, 0, True), ("procedure_name", str, None, 128, 128, 0, False), @@ -2071,7 +2097,9 @@ def procedures(self, procedure=None, catalog=None, schema=None): # Use the helper method to prepare the result set return self._prepare_metadata_result_set(fallback_description=fallback_description) - def primaryKeys(self, table, catalog=None, schema=None): + def primaryKeys( + self, table: str, catalog: str | None = None, schema: str | None = None + ) -> Cursor: """ Creates a result set of column names that make up the primary key for a table by executing the SQLPrimaryKeys function. @@ -2092,7 +2120,7 @@ def primaryKeys(self, table, catalog=None, schema=None): check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Define fallback description for primary keys - fallback_description = [ + fallback_description: _Description = [ ("table_cat", str, None, 128, 128, 0, True), ("table_schem", str, None, 128, 128, 0, True), ("table_name", str, None, 128, 128, 0, False), @@ -2106,13 +2134,13 @@ def primaryKeys(self, table, catalog=None, schema=None): def foreignKeys( # pylint: disable=too-many-arguments,too-many-positional-arguments self, - table=None, - catalog=None, - schema=None, - foreignTable=None, - foreignCatalog=None, - foreignSchema=None, - ): + table: str | None = None, + catalog: str | None = None, + schema: str | None = None, + foreignTable: str | None = None, + foreignCatalog: str | None = None, + foreignSchema: str | None = None, + ) -> Cursor: """ Executes the SQLForeignKeys function and creates a result set of column names that are foreign keys. @@ -2141,7 +2169,7 @@ def foreignKeys( # pylint: disable=too-many-arguments,too-many-positional-argum check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Define fallback description for foreign keys - fallback_description = [ + fallback_description: _Description = [ ("pktable_cat", str, None, 128, 128, 0, True), ("pktable_schem", str, None, 128, 128, 0, True), ("pktable_name", str, None, 128, 128, 0, False), @@ -2161,7 +2189,13 @@ def foreignKeys( # pylint: disable=too-many-arguments,too-many-positional-argum # Use the helper method to prepare the result set return self._prepare_metadata_result_set(fallback_description=fallback_description) - def rowIdColumns(self, table, catalog=None, schema=None, nullable=True): + def rowIdColumns( + self, + table: str, + catalog: str | None = None, + schema: str | None = None, + nullable: bool = True, + ) -> Cursor: """ Executes SQLSpecialColumns with SQL_BEST_ROWID which creates a result set of columns that uniquely identify a row. @@ -2186,7 +2220,7 @@ def rowIdColumns(self, table, catalog=None, schema=None, nullable=True): check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Define fallback description for special columns - fallback_description = [ + fallback_description: _Description = [ ("scope", int, None, 10, 10, 0, False), ("column_name", str, None, 128, 128, 0, False), ("data_type", int, None, 10, 10, 0, False), @@ -2200,7 +2234,13 @@ def rowIdColumns(self, table, catalog=None, schema=None, nullable=True): # Use the helper method to prepare the result set return self._prepare_metadata_result_set(fallback_description=fallback_description) - def rowVerColumns(self, table, catalog=None, schema=None, nullable=True): + def rowVerColumns( + self, + table: str, + catalog: str | None = None, + schema: str | None = None, + nullable: bool = True, + ) -> Cursor: """ Executes SQLSpecialColumns with SQL_ROWVER which creates a result set of columns that are automatically updated when any value in the row is updated. @@ -2225,7 +2265,7 @@ def rowVerColumns(self, table, catalog=None, schema=None, nullable=True): check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Same fallback description as rowIdColumns - fallback_description = [ + fallback_description: _Description = [ ("scope", int, None, 10, 10, 0, False), ("column_name", str, None, 128, 128, 0, False), ("data_type", int, None, 10, 10, 0, False), @@ -2242,8 +2282,8 @@ def rowVerColumns(self, table, catalog=None, schema=None, nullable=True): def statistics( # pylint: disable=too-many-arguments,too-many-positional-arguments self, table: str, - catalog: str = None, - schema: str = None, + catalog: str | None = None, + schema: str | None = None, unique: bool = False, quick: bool = True, ) -> "Cursor": @@ -2272,7 +2312,7 @@ def statistics( # pylint: disable=too-many-arguments,too-many-positional-argume check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Define fallback description for statistics - fallback_description = [ + fallback_description: _Description = [ ("table_cat", str, None, 128, 128, 0, True), ("table_schem", str, None, 128, 128, 0, True), ("table_name", str, None, 128, 128, 0, False), @@ -2291,7 +2331,13 @@ def statistics( # pylint: disable=too-many-arguments,too-many-positional-argume # Use the helper method to prepare the result set return self._prepare_metadata_result_set(fallback_description=fallback_description) - def columns(self, table=None, catalog=None, schema=None, column=None): + def columns( + self, + table: str | None = None, + catalog: str | None = None, + schema: str | None = None, + column: str | None = None, + ) -> Cursor: """ Creates a result set of column information in the specified tables using the SQLColumns function. @@ -2304,7 +2350,7 @@ def columns(self, table=None, catalog=None, schema=None, column=None): check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, retcode) # Define fallback description for columns - fallback_description = [ + fallback_description: _Description = [ ("table_cat", str, None, 128, 128, 0, True), ("table_schem", str, None, 128, 128, 0, True), ("table_name", str, None, 128, 128, 0, False), @@ -2331,7 +2377,7 @@ def columns(self, table=None, catalog=None, schema=None, column=None): def _transpose_rowwise_to_columnwise( self, seq_of_parameters: Sequence[Sequence[Any]], - ) -> tuple[list, int]: + ) -> tuple[list[list[Any]], int]: """ Convert sequence of rows (row-wise) into list of columns (column-wise), for array binding via ODBC. Works with both iterables and generators. @@ -2342,7 +2388,7 @@ def _transpose_rowwise_to_columnwise( Returns: tuple: (columnwise_data, row_count) """ - columnwise = [] + columnwise: list[list[Any]] = [] first_row = True row_count = 0 @@ -2364,7 +2410,9 @@ def _transpose_rowwise_to_columnwise( return columnwise, row_count - def _compute_column_type(self, column): + def _compute_column_type( + self, column: Sequence[Any] + ) -> tuple[Any, int | None, int | None, int]: """ Determine representative value and integer min/max for a column. @@ -2385,7 +2433,7 @@ def _compute_column_type(self, column): int_values = [v for v in non_nulls if isinstance(v, int)] if int_values: min_val, max_val = min(int_values), max(int_values) - sample_value = max(int_values, key=abs) + sample_value: Any = max(int_values, key=abs) return sample_value, min_val, max_val, 0 sample_value = None @@ -2410,6 +2458,12 @@ def _compute_column_type(self, column): # For Decimal objects, prefer the one that requires higher precision or scale v_tuple = v.as_tuple() sample_tuple = sample_value.as_tuple() + v_exponent = v_tuple.exponent + sample_exponent = sample_tuple.exponent + if isinstance(v_exponent, str) or isinstance(sample_exponent, str): + raise ValueError( + "Cannot infer precision/scale from non-finite Decimal (NaN/Infinity)" + ) # Calculate precision (total significant digits) and scale (decimal places) # For a number like 0.000123456789, we need precision = 9, scale = 12 @@ -2417,14 +2471,14 @@ def _compute_column_type(self, column): # The scale is the number of decimal places needed to represent the number v_precision = len(v_tuple.digits) - if v_tuple.exponent < 0: - v_scale = -v_tuple.exponent + if v_exponent < 0: + v_scale = -v_exponent else: v_scale = 0 sample_precision = len(sample_tuple.digits) - if sample_tuple.exponent < 0: - sample_scale = -sample_tuple.exponent + if sample_exponent < 0: + sample_scale = -sample_exponent else: sample_scale = 0 @@ -2511,6 +2565,9 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s "executemany: Converted %d rows from pyformat to qmark", len(seq_of_parameters) ) + # Named rows have been normalized; array binding indexes positional rows. + seq_of_parameters = cast("Sequence[Sequence[Any]]", seq_of_parameters) + # Apply timeout if set (non-zero) if self._timeout > 0: try: @@ -2806,7 +2863,7 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s self.last_executed_stmt = operation # Fetch column metadata (e.g. for INSERT … OUTPUT) - column_metadata = [] + column_metadata: list[ColumnMetadata] = [] try: ddbc_bindings.DDBCSQLDescribeCol(self.hstmt, column_metadata) self._initialize_description(column_metadata) @@ -2853,7 +2910,7 @@ def fetchone(self) -> Union[None, Row]: wchar_enc = self._cached_wchar_encoding # Fetch raw data - row_data = [] + row_data: list[Any] = [] try: with perf_phase("py::fetchone::cpp_call"): ret = ddbc_bindings.DDBCSQLFetchOne( @@ -2928,7 +2985,7 @@ def fetchmany(self, size: Optional[int] = None) -> List[Row]: wchar_enc = self._cached_wchar_encoding # Fetch raw data - rows_data = [] + rows_data: list[list[Any]] = [] try: with perf_phase("py::fetchmany::cpp_call"): ret = ddbc_bindings.DDBCSQLFetchMany( @@ -3003,7 +3060,7 @@ def fetchall(self) -> List[Row]: wchar_enc = self._cached_wchar_encoding # Fetch raw data - rows_data = [] + rows_data: list[list[Any]] = [] try: with perf_phase("py::fetchall::cpp_call"): ret = ddbc_bindings.DDBCSQLFetchAll( @@ -3072,12 +3129,12 @@ def arrow_batch(self, batch_size: int = 8192) -> "pyarrow.RecordBatch": A pyarrow RecordBatch object containing up to batch_size rows. """ self._check_closed() # Check if the cursor is closed - pyarrow = self._ensure_pyarrow() + pa = self._ensure_pyarrow() if not self._has_result_set and self.description: self._reset_rownumber() - capsules = [] + capsules: list[object] = [] char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) char_c_type = char_decoding.get("ctype", ddbc_sql_const.SQL_WCHAR.value) ret = ddbc_bindings.DDBCSQLFetchArrowBatch( @@ -3085,7 +3142,7 @@ def arrow_batch(self, batch_size: int = 8192) -> "pyarrow.RecordBatch": ) check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret) - batch = pyarrow.RecordBatch._import_from_c_capsule(*capsules) + batch = pa.RecordBatch._import_from_c_capsule(*capsules) # Update rownumber for the number of rows actually fetched num_fetched = batch.num_rows @@ -3112,7 +3169,7 @@ def arrow(self, batch_size: int = 8192) -> "pyarrow.Table": A pyarrow Table containing all remaining rows from the result set. """ self._check_closed() # Check if the cursor is closed - pyarrow = self._ensure_pyarrow() + pa = self._ensure_pyarrow() batches: list["pyarrow.RecordBatch"] = [] while True: @@ -3122,7 +3179,7 @@ def arrow(self, batch_size: int = 8192) -> "pyarrow.Table": batches.append(batch) break batches.append(batch) - return pyarrow.Table.from_batches(batches, schema=batches[0].schema) + return pa.Table.from_batches(batches, schema=batches[0].schema) def arrow_reader(self, batch_size: int = 8192) -> "_ArrowReader": """ @@ -3151,7 +3208,7 @@ def arrow_reader(self, batch_size: int = 8192) -> "_ArrowReader": A pyarrow-compatible RecordBatchReader for the result set. """ self._check_closed() # Check if the cursor is closed - pyarrow = self._ensure_pyarrow() + pa = self._ensure_pyarrow() # Fetch schema without advancing cursor schema_batch = self.arrow_batch(0) @@ -3160,20 +3217,21 @@ def arrow_reader(self, batch_size: int = 8192) -> "_ArrowReader": # Capture the parent cursor in a closure cell that the generator # can null out after cleanup, so a GC'd reader does not keep the # cursor pinned. - cursor_ref = [self] + cursor_ref: list[Cursor | None] = [self] close_requested = [False] - def batch_generator(): + def batch_generator() -> Generator[pyarrow.RecordBatch, None, None]: + cur = cursor_ref[0] + assert cur is not None exhausted = False try: - while (batch := cursor_ref[0].arrow_batch(batch_size)).num_rows > 0: + while (batch := cur.arrow_batch(batch_size)).num_rows > 0: yield batch exhausted = True finally: # Symmetric server-side teardown — runs on exhaustion, # GeneratorExit (from close()), or an exception inside the # body. This is the single canonical cleanup site. - cur = cursor_ref[0] cursor_ref[0] = None if not cur.closed and cur.hstmt is not None: # Natural EOF was captured natively. A close/cancel request @@ -3224,8 +3282,8 @@ def batch_generator(): logger.debug("arrow_reader cleanup: bookkeeping reset failed: %s", e) gen = batch_generator() - inner = pyarrow.RecordBatchReader.from_batches(schema, gen) - return _ArrowReader(self, inner, gen, pyarrow.ArrowInvalid, close_requested) + inner = pa.RecordBatchReader.from_batches(schema, gen) + return _ArrowReader(self, inner, gen, pa.ArrowInvalid, close_requested) def nextset(self) -> Optional[bool]: """ @@ -3273,7 +3331,7 @@ def nextset(self) -> Optional[bool]: self._reset_rownumber() # Initialize description for the new result set - column_metadata = [] + column_metadata: list[ColumnMetadata] = [] try: ddbc_bindings.DDBCSQLDescribeCol(self.hstmt, column_metadata) self._initialize_description(column_metadata) @@ -3301,7 +3359,7 @@ def nextset(self) -> Optional[bool]: return True # ── Mapping from ODBC connection-string keywords (lowercase, as _parse returns) - def _build_pycore_context(self) -> dict: + def _build_pycore_context(self) -> CoreContext: """Build the connection context dict expected by mssql_py_core. Parses the underlying ODBC connection string, validates the SERVER @@ -3331,7 +3389,7 @@ def _build_pycore_context(self) -> dict: raise ValueError("SERVER parameter is required in connection string") # Translate parsed connection string into the dict py-core expects. - pycore_context = connstr_to_pycore_params(params) + pycore_context = cast("CoreContext", connstr_to_pycore_params(params)) # Forward the cursor's query timeout to py-core so the bulkcopy # connection uses the same limit instead of py-core's compiled-in 15s @@ -3434,7 +3492,7 @@ def _build_pycore_context(self) -> dict: return pycore_context - def _looks_like_arrow_source(self, data) -> bool: + def _looks_like_arrow_source(self, data: object) -> bool: """Return True if ``data`` should be routed to :meth:`bulkcopy_arrow`. Soft check: never imports pyarrow eagerly and never raises. Anything @@ -3452,11 +3510,11 @@ def _looks_like_arrow_source(self, data) -> bool: return isinstance(data, (pa.Table, pa.RecordBatch, pa.RecordBatchReader)) @staticmethod - def _bulkcopy_core_and_validate(table_name, batch_size, timeout): + def _bulkcopy_core_and_validate(table_name: str, batch_size: int, timeout: int) -> PyCoreModule: """Import the native core and validate the args shared by ``bulkcopy`` and ``bulkcopy_arrow``. Returns the imported ``mssql_py_core`` module.""" try: - import mssql_py_core + mssql_py_core = cast("PyCoreModule", import_module("mssql_py_core")) except ImportError as exc: logger.error("bulkcopy: Failed to import mssql_py_core module") raise ImportError( @@ -3483,7 +3541,11 @@ def _bulkcopy_core_and_validate(table_name, batch_size, timeout): return mssql_py_core @staticmethod - def _bulkcopy_teardown(pycore_context, pycore_cursor, pycore_connection): + def _bulkcopy_teardown( + pycore_context: CoreContext, + pycore_cursor: CoreCursor | None, + pycore_connection: CoreConnection | None, + ) -> None: """Scrub credential material from the context and close native bulk-copy resources. Safe to call with partially-initialized state.""" if pycore_context: @@ -3503,7 +3565,7 @@ def _bulkcopy_teardown(pycore_context, pycore_cursor, pycore_connection): def bulkcopy( self, table_name: str, - data: Iterable[Union[Tuple, "Row"]], + data: Iterable[tuple[Any, ...]] | Iterable[Row], batch_size: int = 0, timeout: int = 30, column_mappings: Optional[Union[List[str], List[Tuple[int, str]]]] = None, @@ -3513,7 +3575,7 @@ def bulkcopy( keep_nulls: bool = False, fire_triggers: bool = False, use_internal_transaction: bool = False, - ): # pragma: no cover + ) -> BulkCopyResult: # pragma: no cover """ Perform bulk copy operation for high-performance data loading. @@ -3524,6 +3586,7 @@ def bulkcopy( data: Iterable of tuples or Row objects containing row data to be inserted. Row objects from fetchone/fetchmany/fetchall are automatically converted to tuples. Lists and other types are not accepted. + Rows must use the same representation throughout the iterable. Data Format Requirements: - Each element in the iterable represents one row @@ -3627,7 +3690,9 @@ def bulkcopy( # using direct _values access (4x faster than __iter__ protocol). # Uses itertools.chain for C-level iteration (avoids Python # generator frame overhead on the tuple passthrough path). - def _prepare_row_iterator(iterable): + def _prepare_row_iterator( + iterable: Iterable[tuple[Any, ...]] | Iterable[Row], + ) -> Iterator[tuple[Any, ...]]: from itertools import chain it = iter(iterable) @@ -3635,9 +3700,11 @@ def _prepare_row_iterator(iterable): if first is None: return iter(()) if isinstance(first, tuple): - return chain((first,), it) + return chain((first,), cast("Iterator[tuple[Any, ...]]", it)) if isinstance(first, Row): - return (tuple(item._values) for item in chain((first,), it)) + return ( + tuple(item._values) for item in chain((first,), cast("Iterator[Row]", it)) + ) raise TypeError( f"bulkcopy data rows must be tuples or Row objects, " f"got {type(first).__name__}" @@ -3688,7 +3755,7 @@ def _prepare_row_iterator(iterable): def bulkcopy_arrow( self, table_name: str, - source, + source: object, batch_size: int = 0, timeout: int = 30, column_mappings: Optional[Union[List[str], List[Tuple[int, str]]]] = None, @@ -3698,7 +3765,7 @@ def bulkcopy_arrow( keep_nulls: bool = False, fire_triggers: bool = False, use_internal_transaction: bool = False, - ): + ) -> BulkCopyResult: """Bulk-copy from an Apache Arrow source straight into TDS. ``source`` may be any of: @@ -3808,7 +3875,7 @@ def bulkcopy_arrow( finally: self._bulkcopy_teardown(pycore_context, pycore_cursor, pycore_connection) - def __enter__(self): + def __enter__(self) -> Cursor: """ Enter the runtime context for the cursor. @@ -3818,12 +3885,12 @@ def __enter__(self): self._check_closed() return self - def __exit__(self, *args): + def __exit__(self, *args: object) -> None: """Closes the cursor when exiting the context, ensuring proper resource cleanup.""" if not self.closed: self.close() - def fetchval(self): + def fetchval(self) -> Any: """ Fetch the first column of the first row if there are results. @@ -3866,7 +3933,7 @@ def fetchval(self): logger.debug("fetchval: Value retrieved successfully") return row[0] - def commit(self): + def commit(self) -> None: """ Commit all SQL statements executed on the connection that created this cursor. @@ -3895,7 +3962,7 @@ def commit(self): # Delegate to the connection's commit method self._connection.commit() - def rollback(self): + def rollback(self) -> None: """ Roll back all SQL statements executed on the connection that created this cursor. @@ -3924,7 +3991,7 @@ def rollback(self): # Delegate to the connection's rollback method self._connection.rollback() - def __del__(self): + def __del__(self) -> None: """ Destructor to ensure the cursor is closed when it is no longer needed. This is a safety net to ensure resources are cleaned up @@ -3999,7 +4066,7 @@ def scroll( f"Cannot move backward by {value} rows on a forward-only cursor", ) - row_data: list = [] + row_data: list[Any] = [] # Absolute positioning not supported with forward-only cursors if mode == "absolute": @@ -4061,12 +4128,12 @@ def skip(self, count: int) -> None: def _execute_tables( # pylint: disable=too-many-arguments,too-many-positional-arguments self, - stmt_handle, - catalog_name=None, - schema_name=None, - table_name=None, - table_type=None, - ): + stmt_handle: ddbc_bindings.SqlHandle | None, + catalog_name: str | None = None, + schema_name: str | None = None, + table_name: str | None = None, + table_type: str | None = None, + ) -> None: """ Execute SQLTables ODBC function to retrieve table metadata. @@ -4098,8 +4165,12 @@ def _execute_tables( # pylint: disable=too-many-arguments,too-many-positional-a self.messages.extend(ddbc_bindings.DDBCSQLGetAllDiagRecords(stmt_handle)) def tables( - self, table=None, catalog=None, schema=None, tableType=None - ): # pylint: disable=too-many-arguments,too-many-positional-arguments + self, + table: str | None = None, + catalog: str | None = None, + schema: str | None = None, + tableType: str | list[str] | tuple[str, ...] | None = None, + ) -> Cursor: # pylint: disable=too-many-arguments,too-many-positional-arguments """ Returns information about tables in the database that match the given criteria using the SQLTables ODBC function. @@ -4136,7 +4207,7 @@ def tables( ) # Define fallback description for tables - fallback_description = [ + fallback_description: _Description = [ ("table_cat", str, None, 128, 128, 0, True), ("table_schem", str, None, 128, 128, 0, True), ("table_name", str, None, 128, 128, 0, False), diff --git a/mssql_python/db_connection.py b/mssql_python/db_connection.py index ec7067093..c9eb562e7 100644 --- a/mssql_python/db_connection.py +++ b/mssql_python/db_connection.py @@ -6,7 +6,7 @@ from typing import Any, Dict, Optional, Union -from mssql_python.connection import Connection, TokenProvider +from mssql_python.connection import Connection as Connection, TokenProvider def connect( diff --git a/mssql_python/ddbc_bindings.py b/mssql_python/ddbc_bindings.py index 90eaf81a8..d4bd5e9f5 100644 --- a/mssql_python/ddbc_bindings.py +++ b/mssql_python/ddbc_bindings.py @@ -11,9 +11,17 @@ import platform import sysconfig import warnings +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from ._ddbc_types import * + from ._ddbc_types import ( + _get_odbc_driver_path as _get_odbc_driver_path, + _set_odbc_provider as _set_odbc_provider, + ) -def normalize_architecture(platform_name_param, architecture_param): +def normalize_architecture(platform_name_param: str, architecture_param: str) -> str: """ Normalize architecture names for the given platform. @@ -74,7 +82,7 @@ def normalize_architecture(platform_name_param, architecture_param): ) -def get_interpreter_architecture(platform_name_param): +def get_interpreter_architecture(platform_name_param: str) -> str: """ Get the raw architecture string of the running interpreter. @@ -98,7 +106,7 @@ def get_interpreter_architecture(platform_name_param): return platform.machine().lower() -def get_module_architecture(platform_name_param): +def get_module_architecture(platform_name_param: str) -> str: """ Get the architecture token used in the compiled ddbc_bindings filename. @@ -128,7 +136,12 @@ def get_module_architecture(platform_name_param): return architecture_name -def find_module_path(module_dir_param, python_version_param, architecture_param, extension_param): +def find_module_path( + module_dir_param: str, + python_version_param: str, + architecture_param: str, + extension_param: str, +) -> str: """ Find the compiled ddbc_bindings module file for the running interpreter. @@ -205,6 +218,8 @@ def find_module_path(module_dir_param, python_version_param, architecture_param, # Use the original module name 'ddbc_bindings' that the C extension was compiled with module_name = "ddbc_bindings" spec = importlib.util.spec_from_file_location(module_name, module_path) +if spec is None or spec.loader is None: + raise ImportError(f"Cannot create a loader for ddbc_bindings at {module_path}") module = importlib.util.module_from_spec(spec) sys.modules[module_name] = module spec.loader.exec_module(module) diff --git a/mssql_python/decimal_config.py b/mssql_python/decimal_config.py index 8b8caf447..985d05851 100644 --- a/mssql_python/decimal_config.py +++ b/mssql_python/decimal_config.py @@ -5,13 +5,17 @@ This module provides functions for managing decimal separator configuration. """ -from typing import TYPE_CHECKING +from typing import Callable, TYPE_CHECKING if TYPE_CHECKING: # pragma: no cover from mssql_python.helpers import Settings -def _setDecimalSeparator(separator: str, settings: "Settings", set_in_cpp_func=None) -> None: +def _setDecimalSeparator( + separator: str, + settings: "Settings", + set_in_cpp_func: Callable[[str], None] | None = None, +) -> None: """ Internal implementation for setting the decimal separator. @@ -65,7 +69,9 @@ def _getDecimalSeparator(settings: "Settings") -> str: return settings.decimal_separator -def create_decimal_separator_functions(settings: "Settings"): +def create_decimal_separator_functions( + settings: "Settings", +) -> tuple[Callable[[str], None], Callable[[], str]]: """ Factory function to create decimal separator getter/setter bound to specific settings. @@ -78,6 +84,7 @@ def create_decimal_separator_functions(settings: "Settings"): Tuple of (setDecimalSeparator, getDecimalSeparator) functions """ # Try to import and initialize the C++ binding + cpp_binding: Callable[[str], None] | None try: from mssql_python.ddbc_bindings import DDBCSetDecimalSeparator diff --git a/mssql_python/exceptions.py b/mssql_python/exceptions.py index ddf0fde08..602d7c1bd 100644 --- a/mssql_python/exceptions.py +++ b/mssql_python/exceptions.py @@ -5,7 +5,7 @@ These classes are used to raise exceptions when an error occurs while executing a query. """ -from typing import Optional +from typing import Callable, Optional from mssql_python.logging import logger import builtins @@ -19,7 +19,7 @@ class ConnectionStringParseError(builtins.Exception): failures. It collects all errors and reports them together. """ - def __init__(self, errors: list) -> None: + def __init__(self, errors: list[str]) -> None: """ Initialize the error with a list of validation errors. @@ -30,7 +30,7 @@ def __init__(self, errors: list) -> None: message = "Connection string parsing failed:\n " + "\n ".join(errors) super().__init__(message) - def __reduce__(self): + def __reduce__(self) -> tuple[type["ConnectionStringParseError"], tuple[list[str]]]: return (self.__class__, (self.errors,)) @@ -50,7 +50,12 @@ def __init__(self, driver_error: str, ddbc_error: str) -> None: self.message = f"Driver Error: {self.driver_error}" super().__init__(self.message) - def __reduce__(self): + def __reduce__( + self, + ) -> tuple[ + Callable[[type["Exception"], str, str, str], "Exception"], + tuple[type["Exception"], str, str, str], + ]: # Reconstruct without re-running __init__/truncate_error_message() to avoid # emitting warnings for already-truncated "[Microsoft]..." messages. return ( @@ -59,7 +64,9 @@ def __reduce__(self): ) @staticmethod - def _unpickle(cls, driver_error: str, ddbc_error: str, message: str): + def _unpickle( + cls: type["Exception"], driver_error: str, ddbc_error: str, message: str + ) -> "Exception": obj = cls.__new__(cls) obj.driver_error = driver_error obj.ddbc_error = ddbc_error diff --git a/mssql_python/helpers.py b/mssql_python/helpers.py index 1e2f41f42..236d8b89f 100644 --- a/mssql_python/helpers.py +++ b/mssql_python/helpers.py @@ -7,7 +7,7 @@ import re import threading import locale -from typing import Any, Union, Tuple, Optional +from typing import Any, Mapping, Union, Tuple, Optional from mssql_python import ddbc_bindings from mssql_python.exceptions import raise_exception from mssql_python.logging import logger @@ -300,7 +300,9 @@ def _sanitize_for_logging(input_val: Any, max_length: int = max_log_length) -> s _PYCORE_UINT32_MAX = 2**32 - 1 -def connstr_to_pycore_params(params: dict, *, strict: bool = False) -> dict: +def connstr_to_pycore_params( + params: Mapping[str, str | int | None], *, strict: bool = False +) -> dict[str, str | int]: """Translate parsed ODBC connection-string parameters for mssql-py-core. Used by async connection setup and by bulk copy when it opens a separate @@ -317,7 +319,7 @@ def connstr_to_pycore_params(params: dict, *, strict: bool = False) -> dict: # path the parser validates keywords first (validate_keywords=True), # but bulkcopy parses with validation off, so this mapping is the # authoritative filter in that path. - pycore_params: dict = {} + pycore_params: dict[str, str | int] = {} seen_pycore_keys = set() for connstr_key, raw_value in params.items(): diff --git a/mssql_python/logging.py b/mssql_python/logging.py index 2130f3a88..e4ff933d9 100644 --- a/mssql_python/logging.py +++ b/mssql_python/logging.py @@ -15,7 +15,7 @@ import re import platform import atexit -from typing import Optional +from typing import Any, Optional, TextIO # Single DEBUG level - all or nothing philosophy # If you need logging, you need to see everything @@ -33,9 +33,10 @@ class ThreadIDFilter(logging.Filter): """Filter that adds thread_id to all log records.""" - def filter(self, record): + def filter(self, record: logging.LogRecord) -> bool: """Add thread_id (OS native) attribute to log record.""" # Use OS native thread ID for debugging compatibility + thread_id: int | None try: thread_id = threading.get_native_id() except AttributeError: @@ -75,7 +76,7 @@ def __new__(cls) -> "MSSQLLogger": cls._instance = super(MSSQLLogger, cls).__new__(cls) return cls._instance - def __init__(self): + def __init__(self) -> None: """Initialize the logger (only once) - thread-safe""" # Use separate lock for initialization check to prevent race condition # This ensures hasattr check and assignment are atomic @@ -96,10 +97,10 @@ def __init__(self): # Output mode and handlers self._output_mode = FILE # Default to file only - self._file_handler = None - self._stdout_handler = None - self._log_file = None - self._custom_log_path = None # Custom log file path (if specified) + self._file_handler: RotatingFileHandler | None = None + self._stdout_handler: logging.StreamHandler[TextIO] | None = None + self._log_file: str | None = None + self._custom_log_path: str | None = None # Custom log file path (if specified) self._handlers_initialized = False self._handler_lock = threading.RLock() # Reentrant lock for handler operations self._cleanup_registered = False # Track if atexit cleanup is registered @@ -122,7 +123,7 @@ def __init__(self): # Don't setup full handlers yet - do it lazily when setLevel is called # This prevents creating log files when user changes output mode before enabling logging - def _setup_handlers(self): + def _setup_handlers(self) -> None: """ Setup handlers based on output mode. Creates file handler and/or stdout handler as needed. @@ -159,7 +160,7 @@ def _setup_handlers(self): # Create CSV formatter # Custom formatter to extract source from message and format as CSV class CSVFormatter(logging.Formatter): - def format(self, record): + def format(self, record: logging.LogRecord) -> str: # Check if this is from py-core (via py_core_log method) if hasattr(record, "funcName") and record.funcName == "py-core": source = "py-core" @@ -234,14 +235,14 @@ def format(self, record): self._stdout_handler.setFormatter(formatter) self._logger.addHandler(self._stdout_handler) - def _reconfigure_handlers(self): + def _reconfigure_handlers(self) -> None: """ Reconfigure handlers when output mode changes. Closes existing handlers and creates new ones based on current output mode. """ self._setup_handlers() - def _cleanup_handlers(self): + def _cleanup_handlers(self) -> None: """ Cleanup all handlers on process exit. Registered with atexit to ensure proper file handle cleanup. @@ -321,7 +322,7 @@ def _validate_log_file_path(self, file_path: str) -> str: return resolved - def _write_log_header(self): + def _write_log_header(self) -> None: """ Write CSV header and metadata to the log file. Called once when log file is created. @@ -375,7 +376,9 @@ def _write_log_header(self): pass # Even stderr notification failed # Don't crash - logging continues without header - def py_core_log(self, level: int, msg: str, filename: str = "cursor.rs", lineno: int = 0): + def py_core_log( + self, level: int, msg: str, filename: str = "cursor.rs", lineno: int = 0 + ) -> None: """ Logging method for py-core (Rust/TDS) code with custom source location. @@ -413,7 +416,9 @@ def py_core_log(self, level: int, msg: str, filename: str = "cursor.rs", lineno: except: pass - def _log(self, level: int, msg: str, add_prefix: bool = True, *args, **kwargs): + def _log( + self, level: int, msg: str, add_prefix: bool = True, *args: object, **kwargs: Any + ) -> None: """ Internal logging method with exception safety. @@ -472,19 +477,19 @@ def _log(self, level: int, msg: str, add_prefix: bool = True, *args, **kwargs): # Convenience methods for logging - def debug(self, msg: str, *args, **kwargs): + def debug(self, msg: str, *args: object, **kwargs: Any) -> None: """Log at DEBUG level (all diagnostic messages)""" self._log(logging.DEBUG, msg, True, *args, **kwargs) - def info(self, msg: str, *args, **kwargs): + def info(self, msg: str, *args: object, **kwargs: Any) -> None: """Log at INFO level""" self._log(logging.INFO, msg, True, *args, **kwargs) - def warning(self, msg: str, *args, **kwargs): + def warning(self, msg: str, *args: object, **kwargs: Any) -> None: """Log at WARNING level""" self._log(logging.WARNING, msg, True, *args, **kwargs) - def error(self, msg: str, *args, **kwargs): + def error(self, msg: str, *args: object, **kwargs: Any) -> None: """Log at ERROR level""" self._log(logging.ERROR, msg, True, *args, **kwargs) @@ -492,7 +497,7 @@ def error(self, msg: str, *args, **kwargs): def _setLevel( self, level: int, output: Optional[str] = None, log_file_path: Optional[str] = None - ): + ) -> None: """ Internal method to set logging level (use setup_logging() instead). @@ -567,30 +572,30 @@ def isEnabledFor(self, level: int) -> bool: # Handler management - def addHandler(self, handler: logging.Handler): + def addHandler(self, handler: logging.Handler) -> None: """Add a handler to the logger (thread-safe)""" with self._handler_lock: self._logger.addHandler(handler) - def removeHandler(self, handler: logging.Handler): + def removeHandler(self, handler: logging.Handler) -> None: """Remove a handler from the logger (thread-safe)""" with self._handler_lock: self._logger.removeHandler(handler) @property - def handlers(self) -> list: + def handlers(self) -> list[logging.Handler]: """Get list of handlers attached to the logger (thread-safe)""" with self._handler_lock: return self._logger.handlers[:] # Return copy to prevent external modification - def reset_handlers(self): + def reset_handlers(self) -> None: """ Reset/recreate handlers. Useful when log file has been deleted or needs to be recreated. """ self._setup_handlers() - def _notify_cpp_level_change(self, level: int): + def _notify_cpp_level_change(self, level: int) -> None: """ Notify C++ bridge that log level has changed. This updates the cached level in C++ for fast checks. @@ -616,7 +621,7 @@ def output(self) -> str: return self._output_mode @output.setter - def output(self, mode: str): + def output(self, mode: str) -> None: """ Set the output mode. @@ -669,7 +674,7 @@ def is_debug_enabled(self) -> bool: # ============================================================================ -def setup_logging(output: str = "file", log_file_path: Optional[str] = None): +def setup_logging(output: str = "file", log_file_path: Optional[str] = None) -> MSSQLLogger: """ Enable DEBUG logging for troubleshooting. diff --git a/mssql_python/mssql_python.pyi b/mssql_python/mssql_python.pyi index 18f70c28c..27b30daff 100644 --- a/mssql_python/mssql_python.pyi +++ b/mssql_python/mssql_python.pyi @@ -16,7 +16,11 @@ from typing import ( Callable, Iterator, Iterable, + Literal, ) +from ._ddbc_types import EncodingSettings +from ._pycore_types import BulkCopyResult +from .logging import MSSQLLogger import datetime import logging import pyarrow @@ -48,7 +52,7 @@ def get_native_provider_info() -> Dict[str, object]: ... def get_info_constants() -> Dict[str, int]: ... # Logging Functions -def setup_logging(mode: str = "file", log_level: int = logging.DEBUG) -> None: ... +def setup_logging(output: str = "file", log_file_path: Optional[str] = None) -> MSSQLLogger: ... def get_logger() -> Optional[logging.Logger]: ... # DB-API 2.0 Type Objects @@ -144,19 +148,19 @@ class Row: def __init__( self, - values: List[Any], - column_map: Dict[str, int], + values: Sequence[Any], + column_map: Mapping[str, int] | None, cursor: Optional["Cursor"] = None, - converter_map: Optional[List[Any]] = None, + converter_map: Sequence[Callable[[Any], Any] | None] | None = None, uuid_str_indices: Optional[Tuple[int, ...]] = None, - column_map_lower: Optional[Dict[str, int]] = None, + column_map_lower: Optional[Mapping[str, int]] = None, column_names: Optional[Tuple[str, ...]] = None, ) -> None: ... @property def _mapping(self) -> "RowMapping": ... - def __getitem__(self, index: int) -> Any: ... + def __getitem__(self, index: int | str | slice) -> Any: ... def __getattr__(self, name: str) -> Any: ... - def __eq__(self, other: Any) -> bool: ... + def __eq__(self, other: object) -> bool: ... def __len__(self) -> int: ... def __iter__(self) -> Iterator[Any]: ... def __str__(self) -> str: ... @@ -228,7 +232,7 @@ class Cursor: def fetchmany(self, size: Optional[int] = None) -> List[Row]: ... def fetchall(self) -> List[Row]: ... def nextset(self) -> Optional[bool]: ... - def setinputsizes(self, sizes: List[Union[int, Tuple[Any, ...]]]) -> None: ... + def setinputsizes(self, sizes: Sequence[int | tuple[int, ...]]) -> None: ... def setoutputsize(self, size: int, column: Optional[int] = None) -> None: ... # Arrow Extension Methods (requires pyarrow) @@ -236,31 +240,11 @@ class Cursor: def arrow(self, batch_size: int = 8192) -> pyarrow.Table: ... def arrow_reader(self, batch_size: int = 8192) -> "_ArrowReader": ... -# pyarrow.RecordBatchReader-compatible wrapper returned by Cursor.arrow_reader. -# Not part of the DB-API 2.0 surface and not intended to be instantiated -# directly by users; declared here so the return type of arrow_reader is -# accurate for static type checkers. Attributes not listed here (e.g. -# ``read_all``, ``read_pandas``, ``cast``) are delegated at runtime to the -# wrapped ``pyarrow.RecordBatchReader`` via ``__getattr__``. -class _ArrowReader: - @property - def closed(self) -> bool: ... - @property - def schema(self) -> pyarrow.Schema: ... - def read_next_batch(self) -> pyarrow.RecordBatch: ... - def close(self) -> None: ... - def __arrow_c_stream__(self, requested_schema: Any = ...) -> Any: ... - def __iter__(self) -> "_ArrowReader": ... - def __next__(self) -> pyarrow.RecordBatch: ... - def __enter__(self) -> "_ArrowReader": ... - def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> None: ... - def __getattr__(self, name: str) -> Any: ... - # Bulk Copy def bulkcopy( self, table_name: str, - data: Iterable[Union[Tuple[Any, ...], Row]], + data: Iterable[Tuple[Any, ...]] | Iterable[Row], batch_size: int = 0, timeout: int = 30, column_mappings: Optional[Union[List[str], List[Tuple[int, str]]]] = None, @@ -270,11 +254,11 @@ class _ArrowReader: keep_nulls: bool = False, fire_triggers: bool = False, use_internal_transaction: bool = False, - ) -> Dict[str, Any]: ... + ) -> BulkCopyResult: ... def bulkcopy_arrow( self, table_name: str, - source: Any, + source: object, batch_size: int = 0, timeout: int = 30, column_mappings: Optional[Union[List[str], List[Tuple[int, str]]]] = None, @@ -284,7 +268,23 @@ class _ArrowReader: keep_nulls: bool = False, fire_triggers: bool = False, use_internal_transaction: bool = False, - ) -> Dict[str, Any]: ... + ) -> BulkCopyResult: ... + +# pyarrow.RecordBatchReader-compatible wrapper returned by Cursor.arrow_reader. +# Methods not listed here are delegated to the wrapped pyarrow reader. +class _ArrowReader: + @property + def closed(self) -> bool: ... + @property + def schema(self) -> pyarrow.Schema: ... + def read_next_batch(self) -> pyarrow.RecordBatch: ... + def close(self) -> None: ... + def __arrow_c_stream__(self, requested_schema: object = None) -> object: ... + def __iter__(self) -> "_ArrowReader": ... + def __next__(self) -> pyarrow.RecordBatch: ... + def __enter__(self) -> "_ArrowReader": ... + def __exit__(self, exc_type: object, exc: object, tb: object) -> Literal[False]: ... + def __getattr__(self, name: str) -> Any: ... # DB-API 2.0 Connection Object # https://www.python.org/dev/peps/pep-0249/#connection-objects @@ -337,11 +337,11 @@ class Connection: # Extension Methods def setautocommit(self, value: bool = False) -> None: ... def setencoding(self, encoding: Optional[str] = None, ctype: Optional[int] = None) -> None: ... - def getencoding(self) -> Dict[str, Union[str, int]]: ... + def getencoding(self) -> EncodingSettings: ... def setdecoding( self, sqltype: int, encoding: Optional[str] = None, ctype: Optional[int] = None ) -> None: ... - def getdecoding(self, sqltype: int) -> Dict[str, Union[str, int]]: ... + def getdecoding(self, sqltype: int) -> EncodingSettings: ... def set_attr(self, attribute: int, value: Union[int, str, bytes, bytearray]) -> None: ... def add_output_converter( self, sqltype: Union[int, type], func: Callable[[Any], Any] @@ -357,7 +357,7 @@ class Connection: reuse_cursor: Optional[Cursor] = None, auto_close: bool = False, ) -> Tuple[List[Union[List[Row], int]], Cursor]: ... - def getinfo(self, info_type: int) -> Union[str, int, bool, None]: ... + def getinfo(self, info_type: int) -> Union[str, int, bool, bytes, None]: ... # Context Manager Support def __enter__(self) -> "Connection": ... diff --git a/mssql_python/parameter_helper.py b/mssql_python/parameter_helper.py index 11a6ddbb6..2337c369b 100644 --- a/mssql_python/parameter_helper.py +++ b/mssql_python/parameter_helper.py @@ -12,7 +12,7 @@ Reference: https://www.python.org/dev/peps/pep-0249/#paramstyle """ -from typing import Dict, List, Tuple, Any, Union +from typing import Dict, List, Tuple, Any, Union, overload from mssql_python.logging import logger # Distinctive marker for escaped percent signs during pyformat conversion @@ -372,9 +372,19 @@ def convert_pyformat_to_qmark(sql: str, param_dict: Dict[str, Any]) -> Tuple[str return rewritten_sql, positional_params +@overload +def detect_and_convert_parameters(sql: str, parameters: None) -> Tuple[str, None]: ... + + +@overload +def detect_and_convert_parameters( + sql: str, parameters: Union[Tuple[Any, ...], List[Any], Dict[str, Any]] +) -> Tuple[str, Union[Tuple[Any, ...], List[Any]]]: ... + + def detect_and_convert_parameters( - sql: str, parameters: Union[None, Tuple, List, Dict] -) -> Tuple[str, Union[None, Tuple, List]]: + sql: str, parameters: Union[None, Tuple[Any, ...], List[Any], Dict[str, Any]] +) -> Tuple[str, Union[None, Tuple[Any, ...], List[Any]]]: """ Auto-detect parameter style and convert to qmark if needed. diff --git a/mssql_python/perf_timer.py b/mssql_python/perf_timer.py index 81724fbf7..351ee0e28 100644 --- a/mssql_python/perf_timer.py +++ b/mssql_python/perf_timer.py @@ -20,7 +20,20 @@ import threading import time from contextlib import contextmanager -from typing import NamedTuple +from typing import Iterator, Literal, NamedTuple, TypedDict + + +class PhaseStats(TypedDict): + calls: int + total_us: float + min_us: float + max_us: float + + +class TimelineEvent(TypedDict): + name: str + start_us: int + duration_us: int class _Counter(NamedTuple): @@ -47,7 +60,7 @@ class _Event(NamedTuple): @contextmanager -def _bookkeeping(): +def _bookkeeping() -> Iterator[None]: # GC can run SQL-cleanup finalizers during our own allocations. Suppress only # recursive samples on this thread, not the cleanup or ordinary nested phases. depth = getattr(_local, "depth", 0) @@ -58,7 +71,7 @@ def _bookkeeping(): _local.depth = depth -def enable(): +def enable() -> None: global _enabled, _window_start_ns with _bookkeeping(): release = _lock.release @@ -70,7 +83,7 @@ def enable(): release() -def disable(): +def disable() -> None: global _enabled with _bookkeeping(): release = _lock.release @@ -85,10 +98,11 @@ def is_enabled() -> bool: return _enabled -def reset(): +def reset() -> None: global _window_start_ns, _stats, _timeline with _bookkeeping(): - stats, timeline = {}, [] + stats: dict[str, _Counter] = {} + timeline: list[_Event] = [] release = _lock.release _lock.acquire() try: @@ -99,10 +113,10 @@ def reset(): release() -def reset_stats_only(): +def reset_stats_only() -> None: global _window_start_ns, _stats with _bookkeeping(): - stats = {} + stats: dict[str, _Counter] = {} release = _lock.release _lock.acquire() try: @@ -112,14 +126,14 @@ def reset_stats_only(): release() -def enable_timeline(): +def enable_timeline() -> None: global _timeline_enabled, _epoch_ns, _timeline # Clear any previously recorded events when (re)setting the epoch, so every # event in _timeline shares the current epoch. Otherwise a second # enable_timeline() without an intervening reset() would leave stale events # whose offsets were computed from an older epoch, corrupting the sort. with _bookkeeping(): - timeline = [] + timeline: list[_Event] = [] release = _lock.release _lock.acquire() try: @@ -130,7 +144,7 @@ def enable_timeline(): release() -def disable_timeline(): +def disable_timeline() -> None: global _timeline_enabled with _bookkeeping(): release = _lock.release @@ -141,7 +155,7 @@ def disable_timeline(): release() -def get_timeline() -> list[dict]: +def get_timeline() -> list[TimelineEvent]: with _bookkeeping(): # Built-in container copies hold the GIL on supported CPython builds. # Immutable entries stay stable; no profiler lock surrounds GC allocations. @@ -156,10 +170,10 @@ def get_timeline() -> list[dict]: ] -def get_stats() -> dict: +def get_stats() -> dict[str, PhaseStats]: with _bookkeeping(): snapshot = _stats.copy() - out = {} + out: dict[str, PhaseStats] = {} for name, s in snapshot.items(): # Keep fractional microseconds when converting accumulated samples. out[name] = { @@ -182,10 +196,10 @@ class _NullPhase: __slots__ = () - def __enter__(self): + def __enter__(self) -> None: return None - def __exit__(self, *exc): + def __exit__(self, *exc: object) -> Literal[False]: return False @@ -199,14 +213,14 @@ class _Phase: __slots__ = ("_name", "_t0") - def __init__(self, name: str): + def __init__(self, name: str) -> None: self._name = name - def __enter__(self): + def __enter__(self) -> None: self._t0 = perf_start() return None - def __exit__(self, *exc): + def __exit__(self, *exc: object) -> Literal[False]: perf_stop(self._name, self._t0) return False @@ -214,7 +228,7 @@ def __exit__(self, *exc): _NULL_PHASE = _NullPhase() -def perf_phase(name: str): +def perf_phase(name: str) -> _NullPhase | _Phase: if not _enabled: return _NULL_PHASE return _Phase(name) @@ -226,7 +240,7 @@ def perf_start() -> int: return time.perf_counter_ns() -def perf_stop(name: str, t0: int): +def perf_stop(name: str, t0: int) -> None: # t0 == 0 means perf_start() ran while disabled (or was never called); a # falsy start has no valid interval, so record nothing rather than a bogus # "now - 0" duration. @@ -235,7 +249,7 @@ def perf_stop(name: str, t0: int): _record(name, time.perf_counter_ns() - t0, t0) -def _record(name: str, elapsed: int, start_ns: int = 0): +def _record(name: str, elapsed: int, start_ns: int = 0) -> None: if getattr(_local, "depth", 0): return with _bookkeeping(): diff --git a/mssql_python/pooling.py b/mssql_python/pooling.py index e56ec01df..388ea56c4 100644 --- a/mssql_python/pooling.py +++ b/mssql_python/pooling.py @@ -133,7 +133,7 @@ def _reset_for_testing(cls) -> None: @atexit.register -def shutdown_pooling(): +def shutdown_pooling() -> None: """ Shutdown pooling during application exit. diff --git a/mssql_python/row.py b/mssql_python/row.py index ccec53787..d127ad52a 100644 --- a/mssql_python/row.py +++ b/mssql_python/row.py @@ -7,10 +7,15 @@ import decimal import uuid as _uuid -from collections.abc import Mapping -from typing import Any +from collections.abc import Callable, Iterator, Mapping, Sequence +from typing import Any, TYPE_CHECKING from mssql_python.logging import logger +if TYPE_CHECKING: + from mssql_python.cursor import Cursor + +OutputConverter = Callable[[Any], Any] + class Row: """ @@ -41,6 +46,8 @@ class Row: print(value) """ + _values: Sequence[Any] + # Slot internal fields while preserving dynamic attributes and weak references. __slots__ = ( "_values", @@ -53,7 +60,13 @@ class Row: ) @staticmethod - def _fast_create(values, column_map, cursor, column_map_lower=None, column_names=None): + def _fast_create( + values: Sequence[Any], + column_map: Mapping[str, int] | None, + cursor: "Cursor | None", + column_map_lower: Mapping[str, int] | None = None, + column_names: tuple[str, ...] | None = None, + ) -> "Row": """Construct a Row bypassing __init__ — for the common fast path. Used by fetchall/fetchmany when no output converters and no UUID @@ -70,14 +83,14 @@ def _fast_create(values, column_map, cursor, column_map_lower=None, column_names def __init__( self, - values, - column_map, - cursor=None, - converter_map=None, - uuid_str_indices=None, - column_map_lower=None, - column_names=None, - ): + values: Sequence[Any], + column_map: Mapping[str, int] | None, + cursor: "Cursor | None" = None, + converter_map: Sequence[OutputConverter | None] | None = None, + uuid_str_indices: tuple[int, ...] | None = None, + column_map_lower: Mapping[str, int] | None = None, + column_names: tuple[str, ...] | None = None, + ) -> None: """ Initialize a Row object with values and pre-built column map. Args: @@ -125,7 +138,7 @@ def __init__( # test constructions); _mapping_keys() then reconstructs names from _column_map. self._column_names = column_names - def _stringify_uuids(self, indices): + def _stringify_uuids(self, indices: tuple[int, ...]) -> None: """ Convert uuid.UUID values at the given column indices to uppercase str in-place. @@ -143,7 +156,7 @@ def _stringify_uuids(self, indices): if v is not None and isinstance(v, _uuid.UUID): vals[i] = str(v).upper() - def _apply_output_converters(self, values, cursor): + def _apply_output_converters(self, values: Sequence[Any], cursor: "Cursor") -> Sequence[Any]: """ Apply output converters to raw values. @@ -194,7 +207,9 @@ def _apply_output_converters(self, values, cursor): return converted_values - def _apply_output_converters_optimized(self, values, converter_map): + def _apply_output_converters_optimized( + self, values: Sequence[Any], converter_map: Sequence[OutputConverter | None] + ) -> list[Any]: """ Apply output converters using pre-computed converter map for optimal performance. @@ -220,7 +235,7 @@ def _apply_output_converters_optimized(self, values, converter_map): return converted_values - def __getitem__(self, index) -> Any: + def __getitem__(self, index: int | str | slice) -> Any: """Allow accessing by numeric index (row[0]) or column name (row["col"]).""" if type(index) is int: return self._values[index] @@ -253,7 +268,7 @@ def __getattr__(self, name: str) -> Any: """ # Handle lowercase attribute access - if lowercase is enabled, # try to match attribute names case-insensitively - if name in self._column_map: + if self._column_map is not None and name in self._column_map: return self._values[self._column_map[name]] # O(1) case-insensitive lookup when lowercase is enabled @@ -292,7 +307,7 @@ def _mapping(self) -> "RowMapping": """ return RowMapping(self) - def _mapping_keys(self) -> tuple: + def _mapping_keys(self) -> tuple[str, ...]: """Canonical, order-preserving column names backing ``_mapping``. Prefers the names snapshotted once by the cursor for the result set, which @@ -307,13 +322,13 @@ def _mapping_keys(self) -> tuple: if self._column_names is not None: return self._column_names if self._column_map: - idx_to_name: dict = {} + idx_to_name: dict[int, str] = {} for name, idx in self._column_map.items(): idx_to_name.setdefault(idx, name) return tuple(idx_to_name[i] for i in sorted(idx_to_name)) return () - def __eq__(self, other: Any) -> bool: + def __eq__(self, other: object) -> bool: """ Support comparison with lists for test compatibility. This is the key change needed to fix the tests. @@ -328,7 +343,7 @@ def __len__(self) -> int: """Return the number of values in the row""" return len(self._values) - def __iter__(self) -> Any: + def __iter__(self) -> Iterator[Any]: """Allow iteration through values""" return iter(self._values) @@ -359,7 +374,7 @@ def __repr__(self) -> str: return repr(tuple(self._values)) -class RowMapping(Mapping): +class RowMapping(Mapping[str, Any]): """Read-only ``Mapping`` view over a :class:`Row` (column name -> value). Created via :attr:`Row._mapping`. Keys are the row's canonical column names, @@ -384,7 +399,7 @@ def __getitem__(self, key: str) -> Any: return self._row[key] raise KeyError(key) - def __iter__(self): + def __iter__(self) -> Iterator[str]: seen = set() for name in self._row._mapping_keys(): if name not in seen: diff --git a/mssql_python/type.py b/mssql_python/type.py index 2b1b392d1..3b813c485 100644 --- a/mssql_python/type.py +++ b/mssql_python/type.py @@ -15,7 +15,7 @@ class STRING(str): This type object is used to describe columns in a database that are string-based (e.g. CHAR). """ - def __new__(cls): + def __new__(cls) -> "STRING": return str.__new__(cls, "") @@ -25,7 +25,7 @@ class BINARY(bytearray): binary columns in a database (e.g. LONG, RAW, BLOBs). """ - def __new__(cls): + def __new__(cls) -> "BINARY": return bytearray.__new__(cls) @@ -34,7 +34,7 @@ class NUMBER(float): This type object is used to describe numeric columns in a database. """ - def __new__(cls): + def __new__(cls) -> "NUMBER": return float.__new__(cls, 0.0) @@ -52,10 +52,10 @@ def __new__( minute: int = 0, second: int = 0, microsecond: int = 0, - tzinfo=None, + tzinfo: datetime.tzinfo | None = None, *, fold: int = 0, - ): + ) -> "DATETIME": return datetime.datetime.__new__( cls, year, month, day, hour, minute, second, microsecond, tzinfo, fold=fold ) @@ -66,7 +66,7 @@ class ROWID(int): This type object is used to describe the "Row ID" column in a database. """ - def __new__(cls): + def __new__(cls) -> "ROWID": return int.__new__(cls, 0) diff --git a/pytest.ini b/pytest.ini index 55827f1f0..11eaf9f64 100644 --- a/pytest.ini +++ b/pytest.ini @@ -3,9 +3,11 @@ markers = stress: marks tests as stress tests (long-running, resource-intensive) slow: marks tests as extra-slow (sustained load, multi-minute duration) + typing: static mssql_python source and stub type checks (run separately in PR validation) # Default options applied to all pytest runs -# Default: pytest -v → Skips stress tests (fast) +# Default: pytest -v → Skips stress and separately gated typing tests # To run ONLY stress tests: pytest -m stress +# To run ONLY typing tests: pytest tests/test_typing.py -m typing # To run ALL tests: pytest -v -m "" -addopts = -m "not stress" +addopts = -m "not stress and not typing" diff --git a/requirements.txt b/requirements.txt index daffd1a1c..bbbc91a5b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,10 @@ autopep8 flake8 pylint cpplint -mypy +mypy==2.3.1 # Type checking stubs types-setuptools +types-psutil +pandas-stubs +pyarrow-stubs diff --git a/setup.py b/setup.py index 45a526db7..640a63423 100644 --- a/setup.py +++ b/setup.py @@ -150,6 +150,7 @@ def finalize_options(self): package_data = { "mssql_python": [ "py.typed", + "*.pyi", "ddbc_bindings.cp*.pyd", "ddbc_bindings.cp*.so", # msvcp140.dll (VC++ runtime) is copied next to the compiled extension by diff --git a/tests/test_000_dependencies.py b/tests/test_000_dependencies.py index 633dafa75..c2521c428 100644 --- a/tests/test_000_dependencies.py +++ b/tests/test_000_dependencies.py @@ -4,6 +4,7 @@ """ import pytest +import ast import platform import os import re @@ -20,6 +21,43 @@ ) +def test_native_typing_exports_match_extension() -> None: + from mssql_python import ddbc_bindings + + declarations = Path(ddbc_bindings.__file__).with_name("_ddbc_types.pyi") + tree = ast.parse(declarations.read_text(encoding="utf-8")) + exports = next( + node.value + for node in tree.body + if isinstance(node, ast.Assign) + and any(isinstance(target, ast.Name) and target.id == "__all__" for target in node.targets) + ) + names = ast.literal_eval(exports) + ["_set_odbc_provider", "_get_odbc_driver_path"] + missing = [name for name in names if not hasattr(ddbc_bindings, name)] + assert not missing, f"Native declarations missing at runtime: {missing}" + + +def test_constant_typing_exports_match_runtime() -> None: + from mssql_python import constants + + tree = ast.parse(Path(constants.__file__).read_text(encoding="utf-8")) + declarations = { + node.target.id + for block in tree.body + if isinstance(block, ast.If) + and isinstance(block.test, ast.Name) + and block.test.id == "TYPE_CHECKING" + for node in block.body + if isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name) + } + expected = { + name + for name in constants.__all__ + if isinstance(getattr(constants, name), int) and name != "SQL_WMETADATA" + } + assert declarations == expected + + class DependencyTester: """Helper class to test platform-specific dependencies.""" @@ -807,6 +845,20 @@ def test_ddbc_bindings_import_error_scenarios(): normalize_architecture(platform_name, arch) +@pytest.mark.parametrize("missing_spec", [True, False]) +def test_ddbc_bindings_missing_loader(monkeypatch: pytest.MonkeyPatch, missing_spec: bool) -> None: + import importlib.machinery + import importlib.util + import runpy + import mssql_python + + spec = None if missing_spec else importlib.machinery.ModuleSpec("ddbc_bindings", None) + monkeypatch.setattr(importlib.util, "spec_from_file_location", lambda *args, **kwargs: spec) + loader_path = Path(mssql_python.__file__).with_name("ddbc_bindings.py") + with pytest.raises(ImportError, match="Cannot create a loader for ddbc_bindings"): + runpy.run_path(str(loader_path)) + + def test_ddbc_bindings_exact_module_match_is_silent(tmp_path, capsys): """find_module_path returns the exact match without any warning or stdout output.""" diff --git a/tests/test_004_cursor.py b/tests/test_004_cursor.py index f2050376f..eaec36a3e 100644 --- a/tests/test_004_cursor.py +++ b/tests/test_004_cursor.py @@ -111,6 +111,18 @@ ] +@pytest.mark.parametrize("nonfinite", ["NaN", "sNaN", "Infinity", "-Infinity"]) +@pytest.mark.parametrize("reverse", [False, True]) +def test_compute_column_type_rejects_nonfinite_precision(nonfinite: str, reverse: bool) -> None: + cur = mssql_python.Cursor.__new__(mssql_python.Cursor) + cur.closed = True + values = [decimal.Decimal("1.25"), decimal.Decimal(nonfinite)] + if reverse: + values.reverse() + with pytest.raises(ValueError, match="non-finite Decimal"): + cur._compute_column_type(values) + + def test_package_sources_compile_with_warnings_as_errors(): """Every package source must compile when warnings are promoted to errors.""" package_dir = Path(__file__).parents[1] / "mssql_python" @@ -3656,6 +3668,8 @@ def test_row_mapping_none_column_map(): assert mapping.get("missing", "fallback") == "fallback" assert mapping.get("missing") is None assert ("missing" in mapping) is False + with pytest.raises(AttributeError, match="missing"): + _ = row.missing def test_row_mapping_dedup_fallback(): diff --git a/tests/test_004_cursor_arrow.py b/tests/test_004_cursor_arrow.py index 37a336789..2b316326b 100644 --- a/tests/test_004_cursor_arrow.py +++ b/tests/test_004_cursor_arrow.py @@ -7,6 +7,8 @@ import pytest import decimal import io +from typing import Generator +from unittest.mock import Mock from datetime import datetime, date, time, timezone import mssql_python @@ -22,6 +24,37 @@ pytestmark = pytest.mark.skipif(pa is None, reason="pyarrow is not installed") +def test_arrow_reader_optional_state_after_close() -> None: + from mssql_python.cursor import _ArrowReader + + cursor = Mock(spec=mssql_python.Cursor) + cursor.closed = False + cursor.hstmt = Mock() + batch = pa.record_batch({"value": [1]}) + released = [] + + def batches() -> Generator[pa.RecordBatch, None, None]: + try: + yield batch + finally: + released.append(True) + + generator = batches() + inner = pa.RecordBatchReader.from_batches(batch.schema, generator) + reader = _ArrowReader(cursor, inner, generator, pa.ArrowInvalid, [False]) + assert next(reader).equals(batch) + reader.close() + reader.close() + assert reader.closed + assert released == [True] + assert reader._cursor is None + assert reader._inner is None + assert reader._generator is None + cursor.hstmt._cancel.assert_called_once() + with pytest.raises(pa.ArrowInvalid, match="Reader is closed"): + next(reader) + + def get_arrow_test_data(include_lobs: bool, batch_length: int): arrow_test_data = [ (pa.uint8(), "tinyint", [1, 2, None, 4, 5, 0, 2**8 - 1]), diff --git a/tests/test_005_connection_cursor_lifecycle.py b/tests/test_005_connection_cursor_lifecycle.py index e91965dfe..27821dde7 100644 --- a/tests/test_005_connection_cursor_lifecycle.py +++ b/tests/test_005_connection_cursor_lifecycle.py @@ -27,6 +27,24 @@ from mssql_python import connect, InterfaceError +def test_closed_native_connection_typing_guards() -> None: + from mssql_python import Connection, Cursor, SQL_ATTR_ACCESS_MODE, SQL_DRIVER_NAME + + connection = Connection.__new__(Connection) + connection._closed = True + connection._conn = None + with pytest.raises(InterfaceError, match="closed"): + _ = connection.autocommit + with pytest.raises(InterfaceError, match="closed"): + connection.setautocommit(True) + with pytest.raises(InterfaceError, match="closed"): + connection.set_attr(SQL_ATTR_ACCESS_MODE, 0) + with pytest.raises(InterfaceError, match="closed"): + connection.getinfo(SQL_DRIVER_NAME) + with pytest.raises(InterfaceError, match="closed"): + Cursor(connection) + + def drop_table_if_exists(cursor, table_name): """Drop the table if it exists""" try: diff --git a/tests/test_typing.py b/tests/test_typing.py new file mode 100644 index 000000000..e57b7769f --- /dev/null +++ b/tests/test_typing.py @@ -0,0 +1,40 @@ +"""Type-check only mssql_python source and stubs without database operations.""" + +from pathlib import Path +import subprocess +import sys + +import pytest + +REPO_ROOT = Path(__file__).resolve().parents[1] +SOURCE_DIR = "mssql_python" +pytestmark = pytest.mark.typing + + +def test_typing(tmp_path: Path) -> None: + result = subprocess.run( + [ + sys.executable, + "-m", + "mypy", + "--config-file=", + "--strict", + "--explicit-package-bases", + "--exclude", + r"(^|/)build/", + "--no-incremental", + "--cache-dir", + str(tmp_path / "mypy"), + SOURCE_DIR, + ], + cwd=REPO_ROOT, + capture_output=True, + text=True, + timeout=300, + ) + if result.returncode != 0: + pytest.fail( + f"mypy failed for {SOURCE_DIR} (exit {result.returncode}):\n" + f"{result.stdout}\n{result.stderr}", + pytrace=False, + )