From 31ac5a7a2d0145420e3fd930b0eedd45be3f0136 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Thu, 17 Sep 2026 21:21:01 +0530 Subject: [PATCH 01/15] REFACTOR: Keep fetchmany column metadata native and call-local Remove native metadata dictionary roundtrips while preserving eager Unicode names and fresh per-call descriptions. Add behavior and profiling regression coverage. Performance acceptance remains unresolved after the bounded local study. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- CHANGELOG.md | 4 + mssql_python/pybind/ddbc_bindings.cpp | 117 +++++-- tests/test_040_fetch_native_metadata.py | 387 ++++++++++++++++++++++++ 3 files changed, 480 insertions(+), 28 deletions(-) create mode 100644 tests/test_040_fetch_native_metadata.py diff --git a/CHANGELOG.md b/CHANGELOG.md index aab046b3a..eee11e3d9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -57,6 +57,10 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), does not change the default provider or ship any Rust driver binaries. ### Changed +- `fetchmany()` keeps freshly described column types and sizes in call-local native + metadata instead of round-tripping them through Python dictionaries. Column names + retain eager Unicode conversion; public column descriptions and fetch behavior + are unchanged. - `mssql-python` now depends on `mssql-python-rs==0.1.0` for `mssql_py_core` instead of embedding files owned by that separately published distribution. - **GH-769 deprecation policy:** The misplaced `GetInfoConstants` members diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 35a3b6da4..3a7ed8a67 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3022,9 +3022,51 @@ SQLSMALLINT SQLNumResultCols_wrap(SqlHandlePtr statementHandle) { return columnCount; } -// Wrap SQLDescribeCol -SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMetadata) { - PERF_TIMER("SQLDescribeCol_wrap"); +namespace { + +struct FetchColumnMetadata { + py::object name; + SQLSMALLINT dataType; + SQLULEN columnSize; + SQLSMALLINT decimalDigits; + SQLSMALLINT nullable; +}; + +py::dict GetFetchColumnMetadata(const py::list& columns, size_t index) { + return columns[index].cast(); +} + +const FetchColumnMetadata& GetFetchColumnMetadata( + const std::vector& columns, size_t index) { + return columns.at(index); +} + +SQLSMALLINT GetFetchColumnType(const py::dict& column) { + return column["DataType"].cast(); +} + +SQLSMALLINT GetFetchColumnType(const FetchColumnMetadata& column) { + return column.dataType; +} + +SQLULEN GetFetchColumnSize(const py::dict& column) { + return column["ColumnSize"].cast(); +} + +SQLULEN GetFetchColumnSize(const FetchColumnMetadata& column) { + return column.columnSize; +} + +std::string GetFetchColumnName(const py::dict& column) { + return column["ColumnName"].cast(); +} + +std::string GetFetchColumnName(const FetchColumnMetadata& column) { + return column.name.cast(); +} + +template +SQLRETURN DescribeColumns(SqlHandlePtr StatementHandle, AppendColumn&& appendColumn) { LOG("SQLDescribeCol: Getting column descriptions for statement_handle=%p", (void*)StatementHandle->get()); if (!SQLDescribeCol_ptr) { @@ -3052,14 +3094,12 @@ SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMeta &ColumnSize, &DecimalDigits, &Nullable); if (SQL_SUCCEEDED(retcode)) { - // Append a named py::dict to ColumnMetadata - // TODO: Should we define a struct for this task instead of dict? - ColumnMetadata.append( - py::dict("ColumnName"_a = dupeSqlWCharAsUtf16Le( - ColumnName, std::min(static_cast(NameLength), - (sizeof(ColumnName) / sizeof(SQLWCHAR)) - 1)), - "DataType"_a = DataType, "ColumnSize"_a = ColumnSize, - "DecimalDigits"_a = DecimalDigits, "Nullable"_a = Nullable)); + // Own the name and preserve eager UTF-16 conversion, including codec errors. + auto name = py::cast(dupeSqlWCharAsUtf16Le( + ColumnName, std::min(static_cast(NameLength), + (sizeof(ColumnName) / sizeof(SQLWCHAR)) - 1))); + appendColumn(FetchColumnMetadata{ + std::move(name), DataType, ColumnSize, DecimalDigits, Nullable}); } else { return retcode; } @@ -3067,6 +3107,19 @@ SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMeta return SQL_SUCCESS; } +} // namespace + +// Wrap SQLDescribeCol +SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMetadata) { + PERF_TIMER("SQLDescribeCol_wrap"); + return DescribeColumns(StatementHandle, [&](FetchColumnMetadata column) { + ColumnMetadata.append( + py::dict("ColumnName"_a = column.name, "DataType"_a = column.dataType, + "ColumnSize"_a = column.columnSize, "DecimalDigits"_a = column.decimalDigits, + "Nullable"_a = column.nullable)); + }); +} + SQLRETURN SQLSpecialColumns_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT identifierType, const py::object& catalogObj, const py::object& schemaObj, const std::u16string& table, SQLSMALLINT scope, @@ -3999,16 +4052,17 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri // For column in the result set, binds a buffer to retrieve column data // TODO: Move to anonymous namespace, since it is not used outside this file -SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& columnNames, +template +SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames, SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR) { PERF_TIMER("SQLBindColums"); SQLRETURN ret = SQL_SUCCESS; const bool useWideChar = (charCtype == SQL_C_WCHAR); // Bind columns based on their data types for (SQLUSMALLINT col = 1; col <= numCols; col++) { - auto columnMeta = columnNames[col - 1].cast(); - SQLSMALLINT dataType = columnMeta["DataType"].cast(); - SQLULEN columnSize = columnMeta["ColumnSize"].cast(); + const auto& columnMeta = GetFetchColumnMetadata(columnNames, col - 1); + SQLSMALLINT dataType = GetFetchColumnType(columnMeta); + SQLULEN columnSize = GetFetchColumnSize(columnMeta); switch (dataType) { case SQL_CHAR: @@ -4140,7 +4194,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column buffers.indicators[col - 1].data()); break; default: - std::string columnName = columnMeta["ColumnName"].cast(); + std::string columnName = GetFetchColumnName(columnMeta); std::ostringstream errorString; errorString << "Unsupported data type for column - " << columnName.c_str() << ", Type - " << dataType << ", column ID - " << col; @@ -4149,7 +4203,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column break; } if (!SQL_SUCCEEDED(ret)) { - std::string columnName = columnMeta["ColumnName"].cast(); + std::string columnName = GetFetchColumnName(columnMeta); std::ostringstream errorString; errorString << "Failed to bind column - " << columnName.c_str() << ", Type - " << dataType << ", column ID - " << col; @@ -4163,7 +4217,8 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column // Fetch rows in batches // TODO: Move to anonymous namespace, since it is not used outside this file -SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& columnNames, +template +SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames, py::list& rows, SQLUSMALLINT numCols, SQLULEN& numRowsFetched, const std::vector& lobColumns, const std::string& charEncoding = "utf-16le", @@ -4213,9 +4268,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum { PERF_TIMER("FetchBatchData::cache_column_metadata"); for (SQLUSMALLINT col = 0; col < numCols; col++) { - const auto& columnMeta = columnNames[col].cast(); - columnInfos[col].dataType = columnMeta["DataType"].cast(); - columnInfos[col].columnSize = columnMeta["ColumnSize"].cast(); + const auto& columnMeta = GetFetchColumnMetadata(columnNames, col); + columnInfos[col].dataType = GetFetchColumnType(columnMeta); + columnInfos[col].columnSize = GetFetchColumnSize(columnMeta); columnInfos[col].isLob = std::find(lobColumns.begin(), lobColumns.end(), col + 1) != lobColumns.end(); columnInfos[col].processedColumnSize = columnInfos[col].columnSize; @@ -4504,8 +4559,8 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum break; } default: { - const auto& columnMeta = columnNames[col - 1].cast(); - std::string columnName = columnMeta["ColumnName"].cast(); + const auto& columnMeta = GetFetchColumnMetadata(columnNames, col - 1); + std::string columnName = GetFetchColumnName(columnMeta); std::ostringstream errorString; errorString << "Unsupported data type for column - " << columnName.c_str() << ", Type - " << dataType << ", column ID - " << col; @@ -4650,18 +4705,24 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle); // Retrieve column metadata - py::list columnNames; - ret = SQLDescribeCol_wrap(StatementHandle, columnNames); + std::vector columnNames; + ret = DescribeColumns(StatementHandle, [&](FetchColumnMetadata column) { + columnNames.push_back(std::move(column)); + }); if (!SQL_SUCCEEDED(ret)) { LOG("FetchMany_wrap: Failed to get column descriptions - SQLRETURN=%d", ret); return ret; } + if (numCols < 0 || columnNames.size() != static_cast(numCols)) { + LOG("FetchMany_wrap: Column metadata count does not match result column count"); + ThrowStdException("Column metadata count does not match result column count"); + } std::vector lobColumns; for (SQLSMALLINT i = 0; i < numCols; i++) { - auto colMeta = columnNames[i].cast(); - SQLSMALLINT dataType = colMeta["DataType"].cast(); - SQLULEN columnSize = colMeta["ColumnSize"].cast(); + const auto& colMeta = GetFetchColumnMetadata(columnNames, i); + SQLSMALLINT dataType = GetFetchColumnType(colMeta); + SQLULEN columnSize = GetFetchColumnSize(colMeta); if (IsLobOrVariantColumn(dataType, columnSize)) { lobColumns.push_back(i + 1); // 1-based diff --git a/tests/test_040_fetch_native_metadata.py b/tests/test_040_fetch_native_metadata.py new file mode 100644 index 000000000..eb46976e4 --- /dev/null +++ b/tests/test_040_fetch_native_metadata.py @@ -0,0 +1,387 @@ +"""Call-local fetch metadata must preserve public descriptions and fetch state.""" + +import datetime as dt +import os +from pathlib import Path +import subprocess +import sys +import textwrap +from decimal import Decimal +from uuid import UUID + +import pytest + +import mssql_python +from mssql_python import ddbc_bindings + + +@pytest.fixture +def metadata_cursor(conn_str): + with mssql_python.connect(conn_str) as connection: + with connection.cursor() as cursor: + yield cursor + + +def _query(columns, count=20): + values = ",".join(f"({i})" for i in range(1, 21)) + return ( + f"SELECT {','.join(columns)} FROM (VALUES {values}) AS source(id) " + f"WHERE id<={count} ORDER BY id" + ) + + +def _assert_rows(rows, expected): + assert [tuple(row) for row in rows] == expected + assert [[type(value) for value in row] for row in rows] == [ + [type(value) for value in row] for row in expected + ] + + +def _describe(cursor): + result = [] + assert ddbc_bindings.DDBCSQLDescribeCol(cursor.hstmt, result) == 0 + for column in result: + assert set(column) == {"ColumnName", "DataType", "ColumnSize", "DecimalDigits", "Nullable"} + assert type(column["ColumnName"]) is str + for key in ("DataType", "ColumnSize", "DecimalDigits", "Nullable"): + assert type(column[key]) is int + return result + + +@pytest.mark.parametrize("width", [3, 24]) +@pytest.mark.parametrize("size", [None, 1, 10, 1000, "varied"]) +@pytest.mark.parametrize("count", [0, 20]) +def test_fetchmany_shape_sizes_and_eof(metadata_cursor, width, size, count): + cursor = metadata_cursor + expressions = ["id", "CONVERT(NVARCHAR(30),N'row')", "CONVERT(FLOAT,id)*0.25"] + columns = [f"{expressions[i % 3]} AS c{i}" for i in range(width)] + cursor.execute(_query(columns, count)) + assert cursor.arraysize == 1 + description = cursor.description + assert all(len(column) == 7 for column in description) + assert [column[0] for column in description] == [f"c{i}" for i in range(width)] + metadata = _describe(cursor) + output = [] + iteration = 0 + while True: + fetch_size = (1, 10, 3, 1000)[iteration % 4] if size == "varied" else size + batch = cursor.fetchmany() if fetch_size is None else cursor.fetchmany(fetch_size) + assert cursor.description == description + if not batch: + break + output.extend(batch) + iteration += 1 + _assert_rows(output, [(i, "row", i * 0.25) * (width // 3) for i in range(1, count + 1)]) + assert _describe(cursor) == metadata + assert cursor.fetchmany(1) == [] + assert cursor.fetchone() is None + assert cursor.fetchall() == [] + + +@pytest.mark.parametrize("method", ["fetchmany", "fetchall", "arrow_batch"]) +def test_public_metadata_names_and_fields(metadata_cursor, method): + cursor = metadata_cursor + names = [ + "duplicate", + "duplicate", + "\u03a9\u540d", + "emoji_\U0001f600", + "bracket]name", + "x" * 128, + ] + columns = [f"CONVERT(INT,id) AS [{name.replace(']', ']]')}]" for name in names] + cursor.execute(_query(columns, 1)) + metadata = _describe(cursor) + assert metadata == [ + {"ColumnName": name, "DataType": 4, "ColumnSize": 10, "DecimalDigits": 0, "Nullable": 1} + for name in names + ] + assert [column[0] for column in cursor.description] == names + if method == "arrow_batch": + pytest.importorskip("pyarrow") + batch = cursor.arrow_batch(1) + assert batch.schema.names == names + assert [column.to_pylist() for column in batch.columns] == [[1]] * len(names) + else: + rows = cursor.fetchmany(1) if method == "fetchmany" else cursor.fetchall() + _assert_rows(rows, [(1,) * len(names)]) + assert _describe(cursor) == metadata + + +_TYPES = [ + ("INT", "7", 7), + ("SMALLINT", "-7", -7), + ("BIGINT", "2147483649", 2147483649), + ("TINYINT", "255", 255), + ("BIT", "1", True), + ("REAL", "1.5", 1.5), + ("FLOAT", "2.25", 2.25), + ("DECIMAL(20,4)", "123.4500", Decimal("123.4500")), + ("NUMERIC(28,8)", "-0.125", Decimal("-0.125")), + ("MONEY", "4.25", Decimal("4.25")), + ("DATE", "'2001-02-03'", dt.date(2001, 2, 3)), + ("TIME(7)", "'12:34:56.1234567'", dt.time(12, 34, 56, 123456)), + ("DATETIME2(7)", "'2001-02-03T12:34:56.1234567'", dt.datetime(2001, 2, 3, 12, 34, 56, 123456)), + ("DATETIME", "'2001-02-03T12:34:56'", dt.datetime(2001, 2, 3, 12, 34, 56)), + ( + "DATETIMEOFFSET(7)", + "'2001-02-03T12:34:56.1234567+05:30'", + dt.datetime(2001, 2, 3, 12, 34, 56, 123456, dt.timezone(dt.timedelta(minutes=330))), + ), + ( + "UNIQUEIDENTIFIER", + "'12345678-1234-5678-1234-567812345678'", + UUID("12345678-1234-5678-1234-567812345678"), + ), + ("VARCHAR(20)", "'ascii'", "ascii"), + ("CHAR(8)", "'ascii'", "ascii "), + ("NVARCHAR(30)", "N'\u03a9\U0001f600'", "\u03a9\U0001f600"), + ("NCHAR(5)", "N'\u03a9'", "\u03a9 "), + ("VARBINARY(10)", "0x00010200", b"\x00\x01\x02\x00"), + ("BINARY(4)", "0x00010203", b"\x00\x01\x02\x03"), + ("VARCHAR(1)", "''", ""), + ("NVARCHAR(1)", "N''", ""), +] + + +@pytest.mark.parametrize("size", [1, 10, 1000]) +def test_fetchmany_typed_nulls_and_values(metadata_cursor, size): + columns = [ + f"CASE WHEN id%3=0 THEN CAST(NULL AS {sqltype}) " + f"ELSE CAST({literal} AS {sqltype}) END AS c{i}" + for i, (sqltype, literal, _) in enumerate(_TYPES) + ] + cursor = metadata_cursor + cursor.execute(_query(columns)) + description = cursor.description + metadata = _describe(cursor) + assert len(metadata) == 24 + output = [] + while batch := cursor.fetchmany(size): + output.extend(batch) + values = tuple(value for _, _, value in _TYPES) + _assert_rows(output, [(None,) * 24 if i % 3 == 0 else values for i in range(1, 21)]) + assert cursor.description == description + assert _describe(cursor) == metadata + + +def test_reexecute_and_nextset_change_shape(metadata_cursor): + cursor = metadata_cursor + for _ in range(3): + cursor.execute("SELECT 1 AS first_name; SELECT N'new' AS second_name, 2 AS extra") + _assert_rows(cursor.fetchmany(1), [(1,)]) + assert cursor.nextset() + assert [col[0] for col in cursor.description] == ["second_name", "extra"] + _assert_rows(cursor.fetchmany(10), [("new", 2)]) + assert not cursor.nextset() + cursor.execute("SELECT CAST(3.5 AS DECIMAL(6,2)) AS replacement") + assert _describe(cursor)[0]["ColumnName"] == "replacement" + _assert_rows(cursor.fetchmany(), [(Decimal("3.50"),)]) + + +def test_converter_changes_on_execute_and_live_decoding(metadata_cursor): + cursor = metadata_cursor + connection = cursor.connection + cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 4)) + _assert_rows(cursor.fetchmany(1), [(1, "ascii")]) + calls = [] + + def convert(value): + calls.append(value) + return value + 100 + + connection.add_output_converter(mssql_python.SQL_INTEGER, convert) + connection.setdecoding(mssql_python.SQL_CHAR, "utf-8", mssql_python.SQL_CHAR) + cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 4)) + _assert_rows(cursor.fetchmany(1), [(101, "ascii")]) + assert calls == [1] + connection.setdecoding(mssql_python.SQL_CHAR, "latin1", mssql_python.SQL_CHAR) + _assert_rows(cursor.fetchmany(1), [(102, "ascii")]) + assert calls == [1, 2] + connection.remove_output_converter(mssql_python.SQL_INTEGER) + cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 1)) + _assert_rows(cursor.fetchmany(1), [(1, "ascii")]) + connection.setdecoding(mssql_python.SQL_CHAR) + cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 1)) + _assert_rows(cursor.fetchmany(10), [(1, "ascii")]) + + +@pytest.mark.parametrize("size", [1, 10]) +def test_fetchmany_lob_and_xml_typed_nulls(metadata_cursor, size): + cursor = metadata_cursor + columns = [ + "CASE WHEN id%2=0 THEN CAST(NULL AS NVARCHAR(MAX)) ELSE " + "REPLICATE(CAST(N'x' AS NVARCHAR(MAX)),9001) END AS txt", + "CASE WHEN id%2=0 THEN CAST(NULL AS VARBINARY(MAX)) ELSE " + "CAST(REPLICATE(CAST('a' AS VARCHAR(MAX)),10003) AS VARBINARY(MAX)) END AS bin", + "CASE WHEN id%2=0 THEN CAST(NULL AS XML) ELSE CAST('value' AS XML) END AS xml", + ] + cursor.execute(_query(columns, 4)) + metadata = _describe(cursor) + rows = [] + while batch := cursor.fetchmany(size): + rows.extend(batch) + values = ("x" * 9001, b"a" * 10003, "value") + _assert_rows(rows, [values, (None, None, None), values, (None, None, None)]) + assert _describe(cursor) == metadata + + +def _isolated(script, tmp_path): + environment = dict(os.environ) + root = str(Path(mssql_python.__file__).resolve().parent.parent) + environment["PYTHONPATH"] = os.pathsep.join([root, environment.get("PYTHONPATH", "")]) + result = subprocess.run( + [sys.executable, "-c", textwrap.dedent(script)], + cwd=tmp_path, + env=environment, + capture_output=True, + text=True, + timeout=45, + ) + assert result.returncode == 0, result.stdout + result.stderr + + +def test_interleaving_movement_and_variant_freshness(tmp_path): + _isolated( + """ + import os + import gc + import mssql_python as db + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + query = "SELECT id FROM (VALUES(1),(2),(3),(4),(5),(6),(7)) s(id) ORDER BY id" + for _ in range(4): + cursor.execute(query) + assert cursor.fetchmany(1)[0][0] == 1 + gc.collect() + assert cursor.fetchone()[0] == 2 + cursor.scroll(1) + assert cursor.fetchmany(1)[0][0] == 4 + cursor.skip(1) + assert [tuple(row) for row in cursor.fetchall()] == [(6,), (7,)] + cursor.execute("CREATE TABLE #metadata_variant(id INT, v SQL_VARIANT)") + cursor.execute( + "INSERT INTO #metadata_variant VALUES " + "(1,CAST('abc' AS VARCHAR(3)))," + "(2,CAST('abcdefgh' AS VARCHAR(8)))," + "(3,CAST(REPLICATE('x',30) AS VARCHAR(30)))" + ) + for method in ("fetchmany", "fetchall"): + cursor.execute("SELECT v FROM #metadata_variant ORDER BY id") + assert cursor.fetchone()[0] == "abc" + if method == "fetchmany": + assert cursor.fetchmany(1)[0][0] == "abcdefgh" + assert cursor.fetchmany(1)[0][0] == "x"*30 + else: + assert [r[0] for r in cursor.fetchall()] == ["abcdefgh", "x"*30] + """, + tmp_path, + ) + + +def test_closure_and_error_recovery(tmp_path): + _isolated( + """ + import os + import mssql_python as db + from mssql_python import Cursor, InterfaceError, ProgrammingError, DatabaseError + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute("SELECT 1 AS c") + assert cursor.fetchmany(0) == [] + assert cursor.fetchmany(-1) == [] + assert cursor.fetchmany(1)[0][0] == 1 + try: + cursor.execute("SELECT invalid_column FROM (VALUES(1)) t(c)") + except DatabaseError: + pass + else: + raise AssertionError("invalid query did not raise") + cursor.execute("SELECT 2 AS changed") + assert cursor.fetchmany(1)[0][0] == 2 + try: + cursor.fetchmany(1) + except ProgrammingError: + pass + else: + raise AssertionError("closed cursor did not raise") + connection = db.connect(os.environ["DB_CONNECTION_STRING"]) + cursor = Cursor(connection) + cursor.execute("SELECT 1") + connection.close() + try: + cursor.fetchmany(1) + except (InterfaceError, ProgrammingError): + pass + else: + raise AssertionError("closed connection did not raise") + cursor.close() + """, + tmp_path, + ) + + +def test_malformed_column_name_fails_before_fetch(tmp_path): + _isolated( + """ + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute( + "DECLARE @s NVARCHAR(200) = N'SELECT 1 AS [' + " + "CAST(0x00D8 AS NVARCHAR(1)) + N']'; EXEC(@s)" + ) + for operation in ( + lambda: native.DDBCSQLDescribeCol(cursor.hstmt, []), + lambda: cursor.fetchmany(1), + ): + try: + operation() + except UnicodeDecodeError: + pass + else: + raise AssertionError("malformed UTF-16 column name did not raise") + assert native.DDBCSQLFetch(cursor.hstmt) == 0 + assert native.DDBCSQLFetch(cursor.hstmt) == 100 + """, + tmp_path, + ) + + +@pytest.mark.skipif( + not hasattr(ddbc_bindings, "profiling"), reason="requires native profiling instrumentation" +) +def test_fetchmany_avoids_python_description_roundtrip(tmp_path): + _isolated( + """ + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute("SELECT id FROM (VALUES(1),(2)) s(id) ORDER BY id") + native.profiling.reset() + native.profiling.enable() + try: + metadata = [] + native.DDBCSQLDescribeCol(cursor.hstmt, metadata) + finally: + native.profiling.disable() + assert native.profiling.get_stats()["ddbc::SQLDescribeCol_wrap"]["calls"] == 1 + assert len(metadata) == 1 + native.profiling.reset() + native.profiling.enable() + try: + assert cursor.fetchmany(1)[0][0] == 1 + assert cursor.fetchmany(1)[0][0] == 2 + assert cursor.fetchmany(1) == [] + finally: + native.profiling.disable() + stats = native.profiling.get_stats() + assert stats["ddbc::FetchMany_wrap"]["calls"] == 3 + assert stats.get("ddbc::SQLDescribeCol_wrap", {}).get("calls", 0) == 0 + """, + tmp_path, + ) From 252e9b6905526da323cd9690646de41cbdf02c23 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Tue, 22 Sep 2026 14:52:42 +0530 Subject: [PATCH 02/15] PERF: Reuse stable native metadata within result sets Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- CHANGELOG.md | 10 +- mssql_python/pybind/connection/connection.cpp | 24 + mssql_python/pybind/connection/connection.h | 1 + mssql_python/pybind/ddbc_bindings.cpp | 220 +++++++-- mssql_python/pybind/ddbc_bindings.h | 2 + mssql_python/pybind/result_metadata.hpp | 76 +++ tests/test_040_fetch_native_metadata.py | 439 +++++++++++++++++- 7 files changed, 722 insertions(+), 50 deletions(-) create mode 100644 mssql_python/pybind/result_metadata.hpp diff --git a/CHANGELOG.md b/CHANGELOG.md index e4ff8eecb..2cac01d4d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -57,10 +57,12 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), does not change the default provider or ship any Rust driver binaries. ### Changed -- `fetchmany()` keeps freshly described column types and sizes in call-local native - metadata instead of round-tripping them through Python dictionaries. Column names - retain eager Unicode conversion; public column descriptions and fetch behavior - are unchanged. +- Fetches reuse owned native metadata for stable columns within a result set; + `fetchmany()` avoids the Python metadata-dictionary roundtrip. Re-execution, + result transitions and statement/connection cleanup invalidate this metadata. + Public descriptions stay fresh, name-validation timing is preserved, and + `sql_variant` columns retain per-row descriptions and per-value probes. + Fetch buffers, decoding settings and converted values are not cached. - DATE, TIME, and TIMESTAMP fetch conversion uses checked CPython constructors for the standard datetime types, while preserving cached substitute constructors, their positional arguments and exceptions, and fractional-second truncation. diff --git a/mssql_python/pybind/connection/connection.cpp b/mssql_python/pybind/connection/connection.cpp index 6960942fa..9c4304d2f 100644 --- a/mssql_python/pybind/connection/connection.cpp +++ b/mssql_python/pybind/connection/connection.cpp @@ -115,6 +115,7 @@ void Connection::connect(const py::dict& attrs_before) { void Connection::disconnect(bool rollbackBeforeDisconnect) { PERF_TIMER("Connection::disconnect"); + clearResultMetadata(); // Determine GIL state once, up front. disconnect() runs both from // pybind11-bound methods (GIL held) and from GIL-less destructor / shutdown // paths: Connection::~Connection() dropping the last shared_ptr, or teardown @@ -265,6 +266,24 @@ void Connection::checkError(SQLRETURN ret) const { } } +void Connection::clearResultMetadata() { + std::vector handles; + { + std::lock_guard lock(_childHandlesMutex); + handles.reserve(_childStatementHandles.size()); + for (const auto& weakHandle : _childStatementHandles) { + if (auto handle = weakHandle.lock()) { + handles.push_back(std::move(handle)); + } + } + } + // Releasing the last handle can acquire the connection cleanup gate. + // Keep that destruction outside the child-list lock. + for (const auto& handle : handles) { + handle->resultMetadata.clear(); + } +} + void Connection::commit() { PERF_TIMER("Connection::commit"); if (!_dbcHandle) { @@ -272,6 +291,7 @@ void Connection::commit() { } updateLastUsed(); LOG("Committing transaction"); + clearResultMetadata(); SQLRETURN ret; { // Release the GIL during the blocking SQLEndTran network round-trip. @@ -288,6 +308,7 @@ void Connection::rollback() { } updateLastUsed(); LOG("Rolling back transaction"); + clearResultMetadata(); SQLRETURN ret; { // Release the GIL during the blocking SQLEndTran network round-trip. @@ -302,6 +323,7 @@ void Connection::setAutocommit(bool enable) { if (!_dbcHandle) { ThrowStdException("Connection handle not allocated"); } + clearResultMetadata(); SQLINTEGER value = enable ? SQL_AUTOCOMMIT_ON : SQL_AUTOCOMMIT_OFF; LOG("Setting autocommit=%d", enable); SQLRETURN ret; @@ -395,6 +417,7 @@ SqlHandlePtr Connection::allocStatementHandle() { } SQLRETURN Connection::setAttribute(SQLINTEGER attribute, py::object value) { + clearResultMetadata(); LOG("Setting SQL attribute=%d", attribute); // SQLPOINTER ptr = nullptr; // SQLINTEGER length = 0; @@ -581,6 +604,7 @@ bool Connection::reset() { if (!_dbcHandle) { ThrowStdException("Connection handle not allocated"); } + clearResultMetadata(); LOG("Resetting connection via SQL_ATTR_RESET_CONNECTION"); // NOTE: SQL_ATTR_RESET_CONNECTION is a pool-checkin reset: it asks the // driver to wipe per-session state (temp tables, open cursors, SET diff --git a/mssql_python/pybind/connection/connection.h b/mssql_python/pybind/connection/connection.h index 0c66aed1e..f43613842 100644 --- a/mssql_python/pybind/connection/connection.h +++ b/mssql_python/pybind/connection/connection.h @@ -103,6 +103,7 @@ class Connection { void allocateDbcHandle(); void checkError(SQLRETURN ret) const; void applyAttrsBefore(const py::dict& attrs_before); + void clearResultMetadata(); std::u16string _connStr; bool _fromPool = false; diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 7fcbcf1f0..13dfc9719 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -1581,6 +1581,7 @@ void SqlHandle::markImplicitlyFreed() { _type); return; // Refuse to mark - let normal free() handle it } + resultMetadata.clear(); _implicitly_freed = true; } @@ -1597,6 +1598,7 @@ void SqlHandle::free() { SQLRETURN SqlHandle::freeHandle() { PERF_TIMER("SqlHandle::free"); + resultMetadata.clear(); bool pythonShuttingDown = is_python_finalizing(); bool skipDuringShutdown = _type == SQL_HANDLE_STMT || _type == SQL_HANDLE_DBC; #ifdef _WIN32 @@ -1637,6 +1639,7 @@ SQLRETURN SqlHandle::freeHandle() { } void SqlHandle::close_cursor() { + resultMetadata.clear(); if (is_python_finalizing()) { return; } @@ -1664,6 +1667,7 @@ void SqlHandle::close_cursor() { } void SqlHandle::cancel() { + resultMetadata.clear(); if (is_python_finalizing()) { return; } @@ -1698,6 +1702,7 @@ SQLRETURN SQLResetStmt_wrap(SqlHandlePtr statementHandle) { if (statementHandle->isImplicitlyFreed()) { return SQL_INVALID_HANDLE; } + statementHandle->resultMetadata.clear(); if (!SQLFreeStmt_ptr) { DriverLoader::getInstance().loadDriver(); } @@ -1719,6 +1724,7 @@ SQLRETURN SQLResetStmt_wrap(SqlHandlePtr statementHandle) { SQLRETURN SQLGetTypeInfo_Wrapper(SqlHandlePtr StatementHandle, SQLSMALLINT DataType) { PERF_TIMER("SQLGetTypeInfo_Wrapper"); + StatementHandle->resultMetadata.clear(); if (!SQLGetTypeInfo_ptr) { ThrowStdException("SQLGetTypeInfo function not loaded"); } @@ -1731,6 +1737,7 @@ SQLRETURN SQLGetTypeInfo_Wrapper(SqlHandlePtr StatementHandle, SQLSMALLINT DataT SQLRETURN SQLProcedures_wrap(SqlHandlePtr StatementHandle, const py::object& catalogObj, const py::object& schemaObj, const py::object& procedureObj) { PERF_TIMER("SQLProcedures_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLProcedures_ptr) { ThrowStdException("SQLProcedures function not loaded"); } @@ -1755,6 +1762,7 @@ SQLRETURN SQLForeignKeys_wrap(SqlHandlePtr StatementHandle, const py::object& pk const py::object& fkCatalogObj, const py::object& fkSchemaObj, const py::object& fkTableObj) { PERF_TIMER("SQLForeignKeys_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLForeignKeys_ptr) { ThrowStdException("SQLForeignKeys function not loaded"); } @@ -1787,6 +1795,7 @@ SQLRETURN SQLForeignKeys_wrap(SqlHandlePtr StatementHandle, const py::object& pk SQLRETURN SQLPrimaryKeys_wrap(SqlHandlePtr StatementHandle, const py::object& catalogObj, const py::object& schemaObj, const std::u16string& table) { PERF_TIMER("SQLPrimaryKeys_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLPrimaryKeys_ptr) { ThrowStdException("SQLPrimaryKeys function not loaded"); } @@ -1809,6 +1818,7 @@ SQLRETURN SQLStatistics_wrap(SqlHandlePtr StatementHandle, const py::object& cat const py::object& schemaObj, const std::u16string& table, SQLUSMALLINT unique, SQLUSMALLINT reserved) { PERF_TIMER("SQLStatistics_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLStatistics_ptr) { ThrowStdException("SQLStatistics function not loaded"); } @@ -1831,6 +1841,7 @@ SQLRETURN SQLColumns_wrap(SqlHandlePtr StatementHandle, const py::object& catalo const py::object& schemaObj, const py::object& tableObj, const py::object& columnObj) { PERF_TIMER("SQLColumns_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLColumns_ptr) { ThrowStdException("SQLColumns function not loaded"); } @@ -1946,6 +1957,7 @@ py::list SQLGetAllDiagRecords(SqlHandlePtr handle) { // Wrap SQLExecDirect SQLRETURN SQLExecDirect_wrap(SqlHandlePtr StatementHandle, const std::u16string& Query) { PERF_TIMER("SQLExecDirect_wrap"); + StatementHandle->resultMetadata.clear(); LOG("SQLExecDirect: Executing query directly - statement_handle=%p, " "query_length=%zu chars", (void*)StatementHandle->get(), Query.length()); @@ -1982,6 +1994,7 @@ SQLRETURN SQLTables_wrap(SqlHandlePtr StatementHandle, const std::u16string& cat const std::u16string& schema, const std::u16string& table, const std::u16string& tableType) { PERF_TIMER("SQLTables_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLTables_ptr) { LOG("SQLTables: Function pointer not initialized, loading driver"); DriverLoader::getInstance().loadDriver(); @@ -2028,6 +2041,7 @@ SQLRETURN SQLExecute_wrap(const SqlHandlePtr statementHandle, return SQL_INVALID_HANDLE; } + statementHandle->resultMetadata.clear(); SQLHANDLE hStmt = statementHandle->get(); // Configure forward-only / read-only cursor (matches slow path semantics). @@ -2828,6 +2842,7 @@ SQLRETURN SQLExecuteMany_wrap(const SqlHandlePtr statementHandle, const std::u16 std::vector& paramInfos, size_t paramSetSize, const py::dict& encodingSettings) { PERF_TIMER("SQLExecuteMany_wrap"); + statementHandle->resultMetadata.clear(); LOG("SQLExecuteMany: Starting batch execution - param_count=%zu, " "param_set_size=%zu", columnwise_params.size(), paramSetSize); @@ -3009,14 +3024,6 @@ SQLSMALLINT SQLNumResultCols_wrap(SqlHandlePtr statementHandle) { namespace { -struct FetchColumnMetadata { - py::object name; - SQLSMALLINT dataType; - SQLULEN columnSize; - SQLSMALLINT decimalDigits; - SQLSMALLINT nullable; -}; - py::dict GetFetchColumnMetadata(const py::list& columns, size_t index) { return columns[index].cast(); } @@ -3047,7 +3054,7 @@ std::string GetFetchColumnName(const py::dict& column) { } std::string GetFetchColumnName(const FetchColumnMetadata& column) { - return column.name.cast(); + return py::cast(column.name).cast(); } template @@ -3074,17 +3081,18 @@ SQLRETURN DescribeColumns(SqlHandlePtr StatementHandle, AppendColumn&& appendCol SQLSMALLINT DecimalDigits; SQLSMALLINT Nullable; - retcode = SQLDescribeCol_ptr(StatementHandle->get(), i, ColumnName, - sizeof(ColumnName) / sizeof(SQLWCHAR), &NameLength, &DataType, - &ColumnSize, &DecimalDigits, &Nullable); + { + PERF_TIMER("SQLDescribeCol::driver_call"); + retcode = SQLDescribeCol_ptr(StatementHandle->get(), i, ColumnName, + sizeof(ColumnName) / sizeof(SQLWCHAR), &NameLength, + &DataType, &ColumnSize, &DecimalDigits, &Nullable); + } if (SQL_SUCCEEDED(retcode)) { - // Own the name and preserve eager UTF-16 conversion, including codec errors. - auto name = py::cast(dupeSqlWCharAsUtf16Le( + auto name = dupeSqlWCharAsUtf16Le( ColumnName, std::min(static_cast(NameLength), - (sizeof(ColumnName) / sizeof(SQLWCHAR)) - 1))); - appendColumn(FetchColumnMetadata{ - std::move(name), DataType, ColumnSize, DecimalDigits, Nullable}); + (sizeof(ColumnName) / sizeof(SQLWCHAR)) - 1)); + appendColumn(std::move(name), DataType, ColumnSize, DecimalDigits, Nullable); } else { return retcode; } @@ -3092,17 +3100,57 @@ SQLRETURN DescribeColumns(SqlHandlePtr StatementHandle, AppendColumn&& appendCol return SQL_SUCCESS; } +SQLRETURN GetResultMetadata(const SqlHandlePtr& statement, SQLSMALLINT columnCount, + std::shared_ptr& metadata) { + const auto snapshot = statement->resultMetadata.snapshot(); + const bool matches = snapshot.metadata && columnCount >= 0 && + snapshot.metadata->columns.size() == static_cast(columnCount); + if (matches && snapshot.metadata->namesValidated) { + metadata = snapshot.metadata; + return SQL_SUCCESS; + } + auto pending = matches ? std::make_shared(*snapshot.metadata) + : std::make_shared(); + if (!matches) { + SQLRETURN ret = DescribeColumns( + statement, [&](std::u16string name, SQLSMALLINT type, SQLULEN size, + SQLSMALLINT digits, SQLSMALLINT nullable) { + // Preserve eager name validation before advancing the result set. + py::cast(name); + pending->columns.push_back( + {std::move(name), type, type == SQL_SS_VARIANT ? 0 : size, digits, nullable}); + }); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } else { + // Row-wise fetches originally read names without decoding them. A later + // many/all fetch must still validate those names before its first advance. + for (const auto& column : pending->columns) { + py::cast(column.name); + } + } + pending->namesValidated = true; + statement->resultMetadata.publish(snapshot.generation, pending); + metadata = std::move(pending); + return SQL_SUCCESS; +} + } // namespace // Wrap SQLDescribeCol SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMetadata) { PERF_TIMER("SQLDescribeCol_wrap"); - return DescribeColumns(StatementHandle, [&](FetchColumnMetadata column) { + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); + ret = DescribeColumns(StatementHandle, [&](std::u16string name, SQLSMALLINT type, + SQLULEN size, SQLSMALLINT digits, + SQLSMALLINT nullable) { ColumnMetadata.append( - py::dict("ColumnName"_a = column.name, "DataType"_a = column.dataType, - "ColumnSize"_a = column.columnSize, "DecimalDigits"_a = column.decimalDigits, - "Nullable"_a = column.nullable)); + py::dict("ColumnName"_a = name, "DataType"_a = type, "ColumnSize"_a = size, + "DecimalDigits"_a = digits, "Nullable"_a = nullable)); }); + return ret; } SQLRETURN SQLSpecialColumns_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT identifierType, @@ -3110,6 +3158,7 @@ SQLRETURN SQLSpecialColumns_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT ident const std::u16string& table, SQLSMALLINT scope, SQLSMALLINT nullable) { PERF_TIMER("SQLSpecialColumns_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLSpecialColumns_ptr) { ThrowStdException("SQLSpecialColumns function not loaded"); } @@ -3138,8 +3187,13 @@ SQLRETURN SQLFetch_wrap(SqlHandlePtr StatementHandle) { } // Release the GIL during the blocking ODBC call - py::gil_scoped_release release; - return SQLFetch_ptr(StatementHandle->get()); + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); + { + py::gil_scoped_release release; + ret = SQLFetch_ptr(StatementHandle->get()); + } + return ret; } // Non-static so it can be called from inline functions in header @@ -3342,24 +3396,58 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLRETURN ret = SQL_SUCCESS; SQLHSTMT hStmt = StatementHandle->get(); - // Cache decimal separator to avoid repeated system calls + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); + const auto snapshot = StatementHandle->resultMetadata.snapshot(); + // The separately exposed GetData entry point may request only a prefix. + const auto metadata = snapshot.metadata && snapshot.metadata->columns.size() >= colCount + ? snapshot.metadata + : nullptr; + auto pending = metadata ? nullptr : std::make_shared(); + bool complete = true; + if (pending) { + pending->columns.reserve(colCount); + } for (SQLSMALLINT i = 1; i <= colCount; ++i) { - SQLWCHAR columnName[256]; + SQLWCHAR uncachedColumnName[256]; + const SQLWCHAR* columnName = uncachedColumnName; SQLSMALLINT columnNameLen; SQLSMALLINT dataType; SQLULEN columnSize; SQLSMALLINT decimalDigits; SQLSMALLINT nullable; - ret = SQLDescribeCol_ptr(hStmt, i, columnName, sizeof(columnName) / sizeof(SQLWCHAR), - &columnNameLen, &dataType, &columnSize, &decimalDigits, &nullable); - if (!SQL_SUCCEEDED(ret)) { - LOG("SQLGetData: Error retrieving metadata for column %d - " - "SQLDescribeCol SQLRETURN=%d", - i, ret); - row.append(py::none()); - continue; + if (metadata && metadata->columns.at(i - 1).dataType != SQL_SS_VARIANT) { + const auto& column = metadata->columns.at(i - 1); + dataType = column.dataType; + columnSize = column.columnSize; + columnName = reinterpretU16stringAsSqlWChar(column.name); + ret = SQL_SUCCESS; + } else { + { + PERF_TIMER("SQLDescribeCol::driver_call"); + ret = SQLDescribeCol_ptr(hStmt, i, uncachedColumnName, + sizeof(uncachedColumnName) / sizeof(SQLWCHAR), + &columnNameLen, &dataType, &columnSize, &decimalDigits, + &nullable); + } + if (!SQL_SUCCEEDED(ret)) { + LOG("SQLGetData: Error retrieving metadata for column %d - " + "SQLDescribeCol SQLRETURN=%d", + i, ret); + complete = false; + row.append(py::none()); + continue; + } + if (pending) { + // Capture declared metadata before probing a variant's current value. + pending->columns.push_back({ + dupeSqlWCharAsUtf16Le( + uncachedColumnName, std::min(static_cast(columnNameLen), + std::size(uncachedColumnName) - 1)), + dataType, dataType == SQL_SS_VARIANT ? 0 : columnSize, decimalDigits, + nullable}); + } } // Preprocess sql_variant: detect underlying type to route to correct conversion logic @@ -3372,10 +3460,14 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p // SQLColAttribute(SQL_CA_SS_VARIANT_TYPE) to return the correct underlying C type. // Without this probe call, SQLColAttribute returns incorrect type codes. SQLLEN indicator; - ret = SQLGetData_ptr(hStmt, i, SQL_C_BINARY, NULL, 0, &indicator); + { + PERF_TIMER("sql_variant::null_probe"); + ret = SQLGetData_ptr(hStmt, i, SQL_C_BINARY, NULL, 0, &indicator); + } if (!SQL_SUCCEEDED(ret)) { LOG_ERROR("SQLGetData: Failed to probe sql_variant column %d - SQLRETURN=%d", i, ret); + complete = false; row.append(py::none()); continue; } @@ -3385,10 +3477,14 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } // Now retrieve the underlying C type SQLLEN variantCType = 0; - ret = - SQLColAttribute_ptr(hStmt, i, SQL_CA_SS_VARIANT_TYPE, NULL, 0, NULL, &variantCType); + { + PERF_TIMER("sql_variant::subtype"); + ret = SQLColAttribute_ptr(hStmt, i, SQL_CA_SS_VARIANT_TYPE, NULL, 0, NULL, + &variantCType); + } if (!SQL_SUCCEEDED(ret)) { LOG_ERROR("SQLGetData: Failed to get sql_variant underlying type for column %d", i); + complete = false; row.append(py::none()); continue; } @@ -4035,6 +4131,14 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p ThrowStdException(errorString.str()); break; } + if (!SQL_SUCCEEDED(ret)) { + complete = false; + } + } + if (!complete) { + StatementHandle->resultMetadata.clear(); + } else if (pending && pending->columns.size() == colCount) { + StatementHandle->resultMetadata.publish(snapshot.generation, std::move(pending)); } return ret; } @@ -4055,7 +4159,8 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri SQLFreeStmt_ptr(StatementHandle->get(), SQL_UNBIND); // Perform scroll operation - SQLRETURN ret; + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); { // Release the GIL during the blocking ODBC fetch py::gil_scoped_release release; @@ -4714,20 +4819,20 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch // Issue #531: upgrade SQL_C_CHAR + utf-8 to SQL_C_WCHAR on Windows so the // driver does lossless UTF-16 conversion instead of returning ACP bytes. charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding); - SQLRETURN ret; + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); SQLHSTMT hStmt = StatementHandle->get(); // Retrieve column count SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle); // Retrieve column metadata - std::vector columnNames; - ret = DescribeColumns(StatementHandle, [&](FetchColumnMetadata column) { - columnNames.push_back(std::move(column)); - }); + std::shared_ptr metadata; + ret = GetResultMetadata(StatementHandle, numCols, metadata); if (!SQL_SUCCEEDED(ret)) { LOG("FetchMany_wrap: Failed to get column descriptions - SQLRETURN=%d", ret); return ret; } + const auto& columnNames = metadata->columns; if (numCols < 0 || columnNames.size() != static_cast(numCols)) { LOG("FetchMany_wrap: Column metadata count does not match result column count"); ThrowStdException("Column metadata count does not match result column count"); @@ -4933,7 +5038,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, // An overly large fetch size doesn't seem to help performance int fetchSize = 64; - SQLRETURN ret; + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); SQLHSTMT hStmt = StatementHandle->get(); // Retrieve column count SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle); @@ -5841,12 +5947,14 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, // Issue #531: upgrade SQL_C_CHAR + utf-8 to SQL_C_WCHAR on Windows so the // driver does lossless UTF-16 conversion instead of returning ACP bytes. charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding); - SQLRETURN ret; + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); SQLHSTMT hStmt = StatementHandle->get(); // Retrieve column count SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle); // Retrieve column metadata + const auto metadataSnapshot = StatementHandle->resultMetadata.snapshot(); py::list columnNames; ret = SQLDescribeCol_wrap(StatementHandle, columnNames); if (!SQL_SUCCEEDED(ret)) { @@ -5872,6 +5980,25 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, LOG("FetchAll_wrap: LOB columns detected (%zu columns), using per-row " "SQLGetData path", lobColumns.size()); + if (numCols < 0 || columnNames.size() != static_cast(numCols)) { + LOG("FetchAll_wrap: Column metadata count does not match result column count"); + ThrowStdException("Column metadata count does not match result column count"); + } + // Keep the public-list setup for fetchall, but reuse its already-validated + // names/declared fields instead of describing stable columns on every row. + auto metadata = std::make_shared(); + metadata->namesValidated = true; + metadata->columns.reserve(numCols); + for (SQLSMALLINT i = 0; i < numCols; ++i) { + const auto column = GetFetchColumnMetadata(columnNames, i); + SQLSMALLINT type = GetFetchColumnType(column); + metadata->columns.push_back({ + column["ColumnName"].cast(), type, + type == SQL_SS_VARIANT ? 0 : GetFetchColumnSize(column), + column["DecimalDigits"].cast(), + column["Nullable"].cast()}); + } + StatementHandle->resultMetadata.publish(metadataSnapshot.generation, std::move(metadata)); while (true) { { // Release GIL during the blocking fetch @@ -5988,7 +6115,8 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row, // Issue #531: upgrade SQL_C_CHAR + utf-8 to SQL_C_WCHAR on Windows so the // driver does lossless UTF-16 conversion instead of returning ACP bytes. charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding); - SQLRETURN ret; + SQLRETURN ret = SQL_ERROR; + ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); SQLHSTMT hStmt = StatementHandle->get(); // Unbind any columns from previous fetch operations (e.g., fetchmany) @@ -6020,6 +6148,7 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row, // Wrap SQLMoreResults SQLRETURN SQLMoreResults_wrap(SqlHandlePtr StatementHandle) { PERF_TIMER("SQLMoreResults_wrap"); + StatementHandle->resultMetadata.clear(); LOG("SQLMoreResults_wrap: Check for more results"); if (!SQLMoreResults_ptr) { LOG("SQLMoreResults_wrap: Function pointer not initialized. Loading " @@ -6237,6 +6366,7 @@ PYBIND11_MODULE(ddbc_bindings, m) { m.def( "DDBCSQLSetStmtAttr", [](SqlHandlePtr stmt, SQLINTEGER attr, py::object value) { + stmt->resultMetadata.clear(); SQLPOINTER ptr_value; if (py::isinstance(value)) { // For integer attributes like SQL_ATTR_QUERY_TIMEOUT diff --git a/mssql_python/pybind/ddbc_bindings.h b/mssql_python/pybind/ddbc_bindings.h index 11c33d8d2..3706e6e9d 100644 --- a/mssql_python/pybind/ddbc_bindings.h +++ b/mssql_python/pybind/ddbc_bindings.h @@ -32,6 +32,7 @@ using py::literals::operator""_a; #include #include +#include "result_metadata.hpp" //------------------------------------------------------------------------------------------------- // SQL Server specific ODBC constants @@ -326,6 +327,7 @@ class SqlHandle { // thread-safe by spec (same assumption as the rest of the driver). std::unordered_map describeCache; void clearDescribeCache() { describeCache.clear(); } + ResultMetadataCache resultMetadata; private: // The caller must release the GIL before waiting for native cleanup. diff --git a/mssql_python/pybind/result_metadata.hpp b/mssql_python/pybind/result_metadata.hpp new file mode 100644 index 000000000..7daece3c7 --- /dev/null +++ b/mssql_python/pybind/result_metadata.hpp @@ -0,0 +1,76 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +struct FetchColumnMetadata { + std::u16string name; + SQLSMALLINT dataType; + SQLULEN columnSize; + SQLSMALLINT decimalDigits; + SQLSMALLINT nullable; +}; + +struct ResultMetadata { + std::vector columns; + bool namesValidated = false; +}; + +// Native-only ownership: cleanup/cancellation may run without the GIL. No ODBC, +// Python, or parent/child handle locks may be acquired while holding this mutex. +class ResultMetadataCache { + public: + struct Snapshot { + uint64_t generation; + std::shared_ptr metadata; + }; + + Snapshot snapshot() const { + std::lock_guard lock(mutex_); + return {generation_, metadata_}; + } + + void publish(uint64_t generation, std::shared_ptr metadata) { + std::lock_guard lock(mutex_); + if (generation == generation_) { + metadata_ = std::move(metadata); + } + } + + void clear() { + std::lock_guard lock(mutex_); + ++generation_; + metadata_.reset(); + } + + private: + mutable std::mutex mutex_; + uint64_t generation_ = 0; + std::shared_ptr metadata_; +}; + +class ResultMetadataFailureGuard { + public: + ResultMetadataFailureGuard(ResultMetadataCache& cache, const SQLRETURN& result) + : cache_(cache), result_(result), exceptions_(std::uncaught_exceptions()) {} + + ~ResultMetadataFailureGuard() { + if (std::uncaught_exceptions() > exceptions_ || + (!SQL_SUCCEEDED(result_) && result_ != SQL_NO_DATA)) { + cache_.clear(); + } + } + + private: + ResultMetadataCache& cache_; + const SQLRETURN& result_; + int exceptions_; +}; diff --git a/tests/test_040_fetch_native_metadata.py b/tests/test_040_fetch_native_metadata.py index eb46976e4..879ce4d92 100644 --- a/tests/test_040_fetch_native_metadata.py +++ b/tests/test_040_fetch_native_metadata.py @@ -1,4 +1,4 @@ -"""Call-local fetch metadata must preserve public descriptions and fetch state.""" +"""Native fetch metadata must preserve public descriptions and result-set state.""" import datetime as dt import os @@ -385,3 +385,440 @@ def test_fetchmany_avoids_python_description_roundtrip(tmp_path): """, tmp_path, ) + + +@pytest.mark.parametrize("method", ["one", "many"]) +def test_result_metadata_prepared_reexecution(metadata_cursor, method): + cursor = metadata_cursor + statement = cursor.hstmt + query = "SELECT CAST(? AS INT) AS n, CAST(? AS NVARCHAR(30)) AS text_value" + for value in range(4): + cursor.execute(query, (value, f"value-{value}")) + assert cursor.hstmt is statement + assert cursor.is_stmt_prepared[0] + profiling = hasattr(ddbc_bindings, "profiling") + if profiling: + ddbc_bindings.profiling.reset() + ddbc_bindings.profiling.enable() + try: + rows = [cursor.fetchone()] if method == "one" else cursor.fetchmany() + finally: + if profiling: + ddbc_bindings.profiling.disable() + _assert_rows(rows, [(value, f"value-{value}")]) + if profiling: + assert ( + ddbc_bindings.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] + == 2 + ) + assert cursor.fetchone() is None + cursor.execute("SELECT CAST(? AS DECIMAL(8,2)) AS amount", (Decimal("3.25"),)) + _assert_rows(cursor.fetchall(), [(Decimal("3.25"),)]) + + +def test_result_metadata_catalog_replacement(metadata_cursor): + cursor = metadata_cursor + for _ in range(2): + cursor.execute("SELECT 42 AS previous_column") + _assert_rows(cursor.fetchmany(), [(42,)]) + cursor.getTypeInfo(mssql_python.SQL_INTEGER) + description = cursor.description + assert len(description) > 1 + row = cursor.fetchone() + assert row is not None and len(row) == len(description) + assert row[1] == mssql_python.SQL_INTEGER + cursor.fetchall() + cursor.execute("SELECT N'replaced' AS new_column, 5 AS extra") + _assert_rows(cursor.fetchmany(), [("replaced", 5)]) + + +def test_result_metadata_native_reset_and_replacement(tmp_path): + _isolated( + """ + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + stmt = connection._conn.alloc_statement_handle() + try: + for _ in range(3): + assert native.DDBCSQLExecDirect(stmt, "SELECT 1 AS a") in (0, 1) + rows = [] + assert native.DDBCSQLFetchMany(stmt, rows, 1) in (0, 1) + assert rows == [[1]] + assert native.DDBCSQLResetStmt(stmt) in (0, 1) + assert native.DDBCSQLExecDirect( + stmt, "SELECT CAST(2 AS BIGINT) AS b, N'new' AS c" + ) in (0, 1) + row = [] + assert native.DDBCSQLFetchOne(stmt, row) in (0, 1) + assert row == [2, "new"] + stmt._close_cursor() + finally: + stmt.free() + """, + tmp_path, + ) + + +@pytest.mark.parametrize("operation", ["commit", "rollback", "autocommit"]) +def test_result_metadata_transaction_recovery(metadata_cursor, operation): + cursor = metadata_cursor + connection = cursor.connection + cursor.execute(_query(["id"], 3)) + _assert_rows(cursor.fetchmany(), [(1,)]) + if operation == "autocommit": + connection.autocommit = True + else: + getattr(connection, operation)() + cursor.execute("SELECT CAST(5.75 AS DECIMAL(8,2)) AS changed, N'text' AS extra") + _assert_rows(cursor.fetchall(), [(Decimal("5.75"), "text")]) + + +@pytest.mark.parametrize("operation", ["commit", "rollback", "autocommit"]) +def test_result_metadata_transaction_preserved_cursor(metadata_cursor, operation): + cursor = metadata_cursor + connection = cursor.connection + info = ( + mssql_python.SQL_CURSOR_ROLLBACK_BEHAVIOR + if operation == "rollback" + else mssql_python.SQL_CURSOR_COMMIT_BEHAVIOR + ) + if connection.getinfo(info) != 2: # SQL_CB_PRESERVE + pytest.skip("Driver does not preserve cursors; native helper coverage is required") + cursor.execute(_query(["id"], 3)) + _assert_rows([cursor.fetchone()], [(1,)]) + if operation == "autocommit": + connection.autocommit = True + else: + getattr(connection, operation)() + profiling = hasattr(ddbc_bindings, "profiling") + if profiling: + ddbc_bindings.profiling.reset() + ddbc_bindings.profiling.enable() + try: + _assert_rows([cursor.fetchone()], [(2,)]) + _assert_rows(cursor.fetchmany(), [(3,)]) + finally: + if profiling: + ddbc_bindings.profiling.disable() + if profiling: + assert ( + ddbc_bindings.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 1 + ) + + +def test_result_metadata_arrow_interleave(tmp_path): + pytest.importorskip("pyarrow") + _isolated( + """ + import gc + import os + import mssql_python as db + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + query = ("SELECT id, CAST(id AS BIGINT) AS big FROM " + "(VALUES(1),(2),(3),(4),(5),(6)) s(id) ORDER BY id") + for _ in range(3): + cursor.execute(query) + assert tuple(cursor.fetchone()) == (1, 1) + assert [tuple(r) for r in cursor.fetchmany(2)] == [(2, 2), (3, 3)] + batch = cursor.arrow_batch(1) + assert [c.to_pylist() for c in batch.columns] == [[4], [4]] + gc.collect() + assert tuple(cursor.fetchone()) == (5, 5) + assert [tuple(r) for r in cursor.fetchall()] == [(6, 6)] + assert cursor.fetchmany() == [] + """, + tmp_path, + ) + + +@pytest.mark.parametrize("method", ["one", "many", "all"]) +def test_result_metadata_variant_type_and_size_changes(metadata_cursor, method): + cursor = metadata_cursor + cursor.execute( + "CREATE TABLE #metadata_mixed_variant (id INT, v SQL_VARIANT, txt NVARCHAR(MAX))" + ) + cursor.execute( + "INSERT INTO #metadata_mixed_variant VALUES " + "(1,CAST(NULL AS SQL_VARIANT),N'first')," + "(2,CAST(CAST('abc' AS VARCHAR(3)) AS SQL_VARIANT),NULL)," + "(3,CAST(CAST('abcdefgh' AS VARCHAR(8)) AS SQL_VARIANT),N'third')," + "(4,CAST(CAST(17 AS INT) AS SQL_VARIANT),NULL)," + "(5,CAST(CAST(3.25 AS DECIMAL(8,2)) AS SQL_VARIANT),N'fifth')," + "(6,CAST(NULL AS SQL_VARIANT),NULL)," + "(7,CAST(CAST(0x010200 AS VARBINARY(3)) AS SQL_VARIANT),N'last')" + ) + cursor.execute("SELECT v, txt FROM #metadata_mixed_variant ORDER BY id") + if method == "one": + rows = list(cursor) + elif method == "many": + rows = [] + while batch := cursor.fetchmany(): + rows.extend(batch) + else: + rows = cursor.fetchall() + _assert_rows( + rows, + [ + (None, "first"), + ("abc", None), + ("abcdefgh", "third"), + (17, None), + (Decimal("3.25"), "fifth"), + (None, None), + (b"\x01\x02\x00", "last"), + ], + ) + + +def test_result_metadata_one_then_malformed_name_many_does_not_advance(tmp_path): + _isolated( + """ + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute( + "DECLARE @s NVARCHAR(300) = N'SELECT id AS [' + " + "CAST(0x00D8 AS NVARCHAR(1)) + " + "N'] FROM (VALUES(1),(2),(3)) s(id) ORDER BY id'; EXEC(@s)" + ) + # The low-level row path never decoded a supported column's name. + row = [] + assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) + assert row == [1] + for _ in range(2): + try: + native.DDBCSQLFetchMany(cursor.hstmt, [], 1) + except UnicodeDecodeError: + pass + else: + raise AssertionError("many accepted the malformed column name") + row = [] + profiling = hasattr(native, "profiling") + if profiling: + native.profiling.reset() + native.profiling.enable() + try: + assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) + finally: + if profiling: + native.profiling.disable() + assert row == [2] + if profiling: + assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 1 + cursor.execute("SELECT 4 AS valid_name, N'recovered' AS text_value") + assert tuple(cursor.fetchone()) == (4, "recovered") + assert cursor.fetchmany() == [] + """, + tmp_path, + ) + + +@pytest.mark.skipif( + not hasattr(ddbc_bindings, "profiling"), reason="requires actual ODBC call instrumentation" +) +@pytest.mark.parametrize("method", ["one", "many"]) +def test_result_metadata_actual_description_counts(tmp_path, method): + _isolated( + f""" + import os + from decimal import Decimal + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + columns = ",".join(f"id AS c{{i}}" for i in range(24)) + query = ("WITH n AS (SELECT TOP(10000) ROW_NUMBER() OVER " + "(ORDER BY a.object_id,b.object_id) AS id " + "FROM sys.all_objects a CROSS JOIN sys.all_objects b) " + f"SELECT {{columns}} FROM n ORDER BY id") + cursor.execute(query) + native.profiling.reset() + native.profiling.enable() + try: + for value in range(1, 10001): + row = cursor.fetchone() if {method!r} == "one" else cursor.fetchmany()[0] + assert tuple(row) == (value,) * 24 + assert cursor.fetchone() is None + assert cursor.fetchmany() == [] + finally: + native.profiling.disable() + stats = native.profiling.get_stats() + assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == 24, stats + assert stats.get("ddbc::SQLDescribeCol_wrap", {{}}).get("calls", 0) == 0 + native.profiling.reset() + native.profiling.enable() + try: + metadata = [] + assert native.DDBCSQLDescribeCol(cursor.hstmt, metadata) in (0, 1) + finally: + native.profiling.disable() + assert len(metadata) == 24 + assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 24 + cursor.execute( + "SELECT CAST(3 AS INT) AS changed, CAST(N'x' AS NVARCHAR(1)) AS text_value; " + "SELECT CAST(7.25 AS DECIMAL(8,2)) AS amount, " + "CAST(N'next long value' AS NVARCHAR(40)) AS name" + ) + assert tuple(cursor.fetchone()) == (3, "x") + assert cursor.nextset() + native.profiling.reset() + native.profiling.enable() + try: + assert tuple(cursor.fetchmany()[0]) == (Decimal("7.25"), "next long value") + assert cursor.fetchmany() == [] + finally: + native.profiling.disable() + assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 2 + """, + tmp_path, + ) + + +@pytest.mark.skipif( + not hasattr(ddbc_bindings, "profiling"), reason="requires actual ODBC call instrumentation" +) +@pytest.mark.parametrize("method", ["one", "many", "all"]) +def test_result_metadata_variant_descriptions_remain_per_row(tmp_path, method): + _isolated( + f""" + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute( + "SELECT id, v FROM (VALUES " + "(1,CAST(NULL AS SQL_VARIANT))," + "(2,CAST('abc' AS SQL_VARIANT))," + "(3,CAST(17 AS SQL_VARIANT))) s(id,v) ORDER BY id" + ) + native.profiling.reset() + native.profiling.enable() + try: + if {method!r} == "one": + rows = list(cursor) + elif {method!r} == "all": + rows = cursor.fetchall() + else: + rows = [] + while batch := cursor.fetchmany(): + rows.extend(batch) + finally: + native.profiling.disable() + assert [tuple(row) for row in rows] == [(1,None),(2,"abc"),(3,17)] + stats = native.profiling.get_stats() + expected = 4 if {method!r} == "one" else 5 + assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == expected, stats + assert stats["ddbc::sql_variant::null_probe"]["calls"] == 3, stats + assert stats["ddbc::sql_variant::subtype"]["calls"] == 2, stats + """, + tmp_path, + ) + + +@pytest.mark.parametrize("method", ["one", "many", "all"]) +def test_result_metadata_all_null_rows(tmp_path, method): + _isolated( + f""" + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute("SELECT CAST(NULL AS INT) AS scalar_null") + scalar_null = [] + scalar_status = native.DDBCSQLFetchOne(cursor.hstmt, scalar_null) + diagnostics = native.DDBCSQLGetAllDiagRecords(cursor.hstmt) + assert scalar_null == [None] + if scalar_status == 0: + assert diagnostics == [] + recovery_descriptions = 0 + else: + assert scalar_status == -1 + assert len(diagnostics) == 1 and "22002" in diagnostics[0][0], diagnostics + assert "Indicator variable required but not supplied" in diagnostics[0][1] + recovery_descriptions = 4 if {method!r} == "many" else 2 + values = ",".join(f"({{i}})" for i in range(1, 16)) + cursor.execute( + "SELECT CASE WHEN id%7=0 THEN NULL ELSE id END AS c0," + "CASE WHEN id%7=0 THEN CAST(NULL AS SQL_VARIANT) " + "WHEN id%3=0 THEN CAST(id AS SQL_VARIANT) " + "WHEN id%3=1 THEN CAST(N'row-'+CONVERT(NVARCHAR(12),id) AS SQL_VARIANT) " + "ELSE CAST(CONVERT(FLOAT,id)*0.25 AS SQL_VARIANT) END AS c1 " + f"FROM (VALUES{{values}}) s(id) ORDER BY id" + ) + profiling = hasattr(native, "profiling") + if profiling: + native.profiling.reset() + native.profiling.enable() + try: + if {method!r} == "one": + rows = list(cursor) + elif {method!r} == "all": + rows = cursor.fetchall() + else: + rows = [] + while batch := cursor.fetchmany(): + rows.extend(batch) + finally: + if profiling: + native.profiling.disable() + expected = [ + (None,None) if i%7==0 else (i,(i,f"row-{{i}}",i*0.25)[i%3]) + for i in range(1,16) + ] + assert [tuple(row) for row in rows] == expected + assert [[type(value) for value in row] for row in rows] == [ + [type(value) for value in row] for row in expected + ] + if profiling: + stats = native.profiling.get_stats() + # Only drivers/builds reporting a real scalar NULL error + # require the additional post-error cache repopulations. + expected_describes = (16 if {method!r} == "one" else 17) + recovery_descriptions + assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == expected_describes, stats + assert stats["ddbc::sql_variant::null_probe"]["calls"] == 15 + assert stats["ddbc::sql_variant::subtype"]["calls"] == 13 + """, + tmp_path, + ) + + +def test_result_metadata_odbc_error_invalidates(tmp_path): + _isolated( + """ + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute("SELECT id FROM (VALUES(1),(2),(3)) s(id) ORDER BY id") + row = [] + assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) + assert row == [1] + assert native.DDBCSQLGetData( + cursor.hstmt, 2, [], "utf-16le", "utf-16le", db.SQL_WCHAR + ) == -1 + diagnostics = native.DDBCSQLGetAllDiagRecords(cursor.hstmt) + assert any("07009" in state for state, _ in diagnostics), diagnostics + profiling = hasattr(native, "profiling") + if profiling: + native.profiling.reset() + native.profiling.enable() + try: + row = [] + assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) + assert row == [2] + finally: + if profiling: + native.profiling.disable() + if profiling: + assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 1 + """, + tmp_path, + ) From abd4f4b006f9bda5147d051fe4c047762eee1a10 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Tue, 22 Sep 2026 19:58:48 +0530 Subject: [PATCH 03/15] REFACTOR: Extract child result metadata invalidation helper Preserve reserve-before-locking-weak-handles and release retained handles outside the child-list mutex. Expose the unchanged native-only algorithm for direct invariant coverage. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/connection/connection.cpp | 16 +------------- mssql_python/pybind/result_metadata.hpp | 21 +++++++++++++++++++ 2 files changed, 22 insertions(+), 15 deletions(-) diff --git a/mssql_python/pybind/connection/connection.cpp b/mssql_python/pybind/connection/connection.cpp index 9c4304d2f..de4eb45c0 100644 --- a/mssql_python/pybind/connection/connection.cpp +++ b/mssql_python/pybind/connection/connection.cpp @@ -267,21 +267,7 @@ void Connection::checkError(SQLRETURN ret) const { } void Connection::clearResultMetadata() { - std::vector handles; - { - std::lock_guard lock(_childHandlesMutex); - handles.reserve(_childStatementHandles.size()); - for (const auto& weakHandle : _childStatementHandles) { - if (auto handle = weakHandle.lock()) { - handles.push_back(std::move(handle)); - } - } - } - // Releasing the last handle can acquire the connection cleanup gate. - // Keep that destruction outside the child-list lock. - for (const auto& handle : handles) { - handle->resultMetadata.clear(); - } + ClearChildResultMetadata(_childHandlesMutex, _childStatementHandles); } void Connection::commit() { diff --git a/mssql_python/pybind/result_metadata.hpp b/mssql_python/pybind/result_metadata.hpp index 7daece3c7..0e08e28bd 100644 --- a/mssql_python/pybind/result_metadata.hpp +++ b/mssql_python/pybind/result_metadata.hpp @@ -9,6 +9,7 @@ #include #include #include +#include #include struct FetchColumnMetadata { @@ -57,6 +58,26 @@ class ResultMetadataCache { std::shared_ptr metadata_; }; +template +void ClearChildResultMetadata(std::mutex& childHandlesMutex, + const std::vector>& childHandles) { + std::vector> handles; + { + std::lock_guard lock(childHandlesMutex); + handles.reserve(childHandles.size()); + for (const auto& weakHandle : childHandles) { + if (auto handle = weakHandle.lock()) { + handles.push_back(std::move(handle)); + } + } + } + // Releasing the last handle can acquire the connection cleanup gate. + // Keep that destruction outside the child-list lock. + for (const auto& handle : handles) { + handle->resultMetadata.clear(); + } +} + class ResultMetadataFailureGuard { public: ResultMetadataFailureGuard(ResultMetadataCache& cache, const SQLRETURN& result) From c5fe1425854ee73ed4fac554e7919b3ee05bf4f1 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Tue, 22 Sep 2026 19:59:03 +0530 Subject: [PATCH 04/15] CHORE: Harden native metadata regression coverage Require successful scalar NULL handling and exact describe counts. Add production-header cache, failure, concurrency, allocation and lifetime tests with active Release assertions, plus Windows/Linux/macOS CTest CI and test guidance. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .github/prompts/run-tests.prompt.md | 21 ++ .github/workflows/native-metadata-tests.yml | 49 +++++ tests/native/CMakeLists.txt | 31 +++ tests/native/allocation_failure.cpp | 32 +++ tests/native/result_metadata_tests.cpp | 208 ++++++++++++++++++++ tests/test_040_fetch_native_metadata.py | 16 +- 6 files changed, 345 insertions(+), 12 deletions(-) create mode 100644 .github/workflows/native-metadata-tests.yml create mode 100644 tests/native/CMakeLists.txt create mode 100644 tests/native/allocation_failure.cpp create mode 100644 tests/native/result_metadata_tests.cpp diff --git a/.github/prompts/run-tests.prompt.md b/.github/prompts/run-tests.prompt.md index da8bcfa88..2d6eb3a9b 100644 --- a/.github/prompts/run-tests.prompt.md +++ b/.github/prompts/run-tests.prompt.md @@ -122,6 +122,27 @@ Help the developer run tests to validate their changes. Follow this process base ## STEP 1: Choose What to Test +### Native metadata invariants (no database) + +The standalone CMake tests in `tests/native` exercise the production metadata +cache and child-handle invalidation helper without importing the Python package +or connecting to SQL Server. They require a C++17 compiler, CMake, and ODBC +headers (Windows SDK, `unixodbc-dev` on Linux, or `unixodbc` on macOS). +The Native Metadata Tests workflow runs them on Windows, Linux, and macOS. + +```bash +cmake -S tests/native -B build/native-metadata -DCMAKE_BUILD_TYPE=Release +cmake --build build/native-metadata --config Release --parallel 2 +ctest --test-dir build/native-metadata -C Release --output-on-failure +``` + +Assertions remain enabled in Release. Cases cover stale-generation rejection, +held snapshots, failure/EOF guards, concurrent invalidation, child isolation, +reserve failure before strong-reference acquisition, and last-owner destruction +outside the child-list lock. These native checks supplement, not replace, the +live transaction tests, which skip when the driver does not preserve cursors. +Native-only tests do not require the Python-test prerequisites above. + ### Test Categories | Category | Description | When to Use | diff --git a/.github/workflows/native-metadata-tests.yml b/.github/workflows/native-metadata-tests.yml new file mode 100644 index 000000000..a5887fbfa --- /dev/null +++ b/.github/workflows/native-metadata-tests.yml @@ -0,0 +1,49 @@ +name: Native Metadata Tests + +on: + pull_request: + types: [opened, reopened, synchronize, ready_for_review] + paths: + - 'mssql_python/pybind/result_metadata.hpp' + - 'mssql_python/pybind/connection/connection.cpp' + - 'tests/native/**' + - '.github/workflows/native-metadata-tests.yml' + push: + branches: [main] + paths: + - 'mssql_python/pybind/result_metadata.hpp' + - 'mssql_python/pybind/connection/connection.cpp' + - 'tests/native/**' + - '.github/workflows/native-metadata-tests.yml' + +permissions: + contents: read + +jobs: + native-metadata: + name: Native metadata (${{ matrix.os }}) + runs-on: ${{ matrix.os }} + timeout-minutes: 10 + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, windows-latest, macos-latest] + steps: + - name: Checkout + uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + with: + persist-credentials: false + - name: Install ODBC headers (Linux) + if: runner.os == 'Linux' + run: | + sudo apt-get update + sudo apt-get install -y unixodbc-dev + - name: Install ODBC headers (macOS) + if: runner.os == 'macOS' + run: brew install unixodbc + - name: Configure + run: cmake -S tests/native -B build/native-metadata -DCMAKE_BUILD_TYPE=Release + - name: Build + run: cmake --build build/native-metadata --config Release --parallel 2 + - name: Test + run: ctest --test-dir build/native-metadata -C Release --output-on-failure diff --git a/tests/native/CMakeLists.txt b/tests/native/CMakeLists.txt new file mode 100644 index 000000000..83c8f3861 --- /dev/null +++ b/tests/native/CMakeLists.txt @@ -0,0 +1,31 @@ +cmake_minimum_required(VERSION 3.15) +project(mssql_python_native_tests LANGUAGES CXX) + +enable_testing() +find_package(Threads REQUIRED) + +add_executable(result_metadata_tests result_metadata_tests.cpp allocation_failure.cpp) +target_compile_features(result_metadata_tests PRIVATE cxx_std_17) +target_include_directories(result_metadata_tests PRIVATE ../../mssql_python/pybind) +target_link_libraries(result_metadata_tests PRIVATE Threads::Threads) + +if(WIN32) + target_compile_definitions(result_metadata_tests PRIVATE WIN32_LEAN_AND_MEAN NOMINMAX) +else() + find_path(ODBC_INCLUDE_DIR sql.h PATHS /opt/homebrew/include /usr/local/include) + if(NOT ODBC_INCLUDE_DIR) + message(FATAL_ERROR "ODBC headers are required: install unixodbc-dev or unixodbc.") + endif() + target_include_directories(result_metadata_tests PRIVATE "${ODBC_INCLUDE_DIR}") +endif() + +if(MSVC) + target_compile_options(result_metadata_tests PRIVATE /W4 /WX /UNDEBUG) +else() + target_compile_options(result_metadata_tests PRIVATE -Wall -Wextra -Werror -UNDEBUG) +endif() + +foreach(case_name IN ITEMS snapshots failures concurrent children allocation last_owner) + add_test(NAME result_metadata_${case_name} COMMAND result_metadata_tests ${case_name}) + set_tests_properties(result_metadata_${case_name} PROPERTIES TIMEOUT 20) +endforeach() diff --git a/tests/native/allocation_failure.cpp b/tests/native/allocation_failure.cpp new file mode 100644 index 000000000..bf9daab61 --- /dev/null +++ b/tests/native/allocation_failure.cpp @@ -0,0 +1,32 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#include +#include +#include + +struct TestHandle; +extern std::weak_ptr observedHandle; +extern bool failAllocation; +extern long ownersAtFailure; + +// Keep replacement allocation functions opaque to optimized test call sites. +void* operator new(std::size_t size) { + if (failAllocation) { + failAllocation = false; + ownersAtFailure = observedHandle.use_count(); + throw std::bad_alloc(); + } + if (void* memory = std::malloc(size ? size : 1)) { + return memory; + } + throw std::bad_alloc(); +} + +void operator delete(void* memory) noexcept { + std::free(memory); +} + +void operator delete(void* memory, std::size_t) noexcept { + std::free(memory); +} diff --git a/tests/native/result_metadata_tests.cpp b/tests/native/result_metadata_tests.cpp new file mode 100644 index 000000000..963608564 --- /dev/null +++ b/tests/native/result_metadata_tests.cpp @@ -0,0 +1,208 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#ifdef _WIN32 +#include +#endif +#include "result_metadata.hpp" + +#include +#include +#include +#include +#include +#include + +#ifdef NDEBUG +#error Native metadata tests require assertions, including Release builds. +#endif + +struct TestHandle { + ResultMetadataCache resultMetadata; + std::mutex* childMutex = nullptr; + + ~TestHandle() { + if (childMutex) { + bool acquired = false; + std::thread observer([&] { + acquired = childMutex->try_lock(); + if (acquired) { + childMutex->unlock(); + } + }); + observer.join(); + assert(acquired); + } + } +}; + +std::weak_ptr observedHandle; +bool failAllocation = false; +long ownersAtFailure = -1; + +static std::shared_ptr MakeMetadata() { + auto metadata = std::make_shared(); + metadata->columns.push_back({u"owned", SQL_INTEGER, 10, 0, 1}); + return metadata; +} + +static void Populate(ResultMetadataCache& cache) { + const auto snapshot = cache.snapshot(); + cache.publish(snapshot.generation, MakeMetadata()); +} + +static void TestSnapshots() { + ResultMetadataCache cache; + const auto initial = cache.snapshot(); + assert(!initial.metadata); + auto metadata = MakeMetadata(); + std::weak_ptr weak = metadata; + cache.publish(initial.generation, metadata); + auto held = cache.snapshot(); + assert(held.metadata == metadata); + cache.clear(); + assert(!cache.snapshot().metadata); + assert(cache.snapshot().generation != initial.generation); + + Populate(cache); + const auto replacement = cache.snapshot(); + cache.publish(initial.generation, metadata); + assert(cache.snapshot().metadata == replacement.metadata); + assert(held.metadata->columns.at(0).name == u"owned"); + metadata.reset(); + assert(!weak.expired()); + held.metadata.reset(); + assert(weak.expired()); +} + +static void TestFailures() { + ResultMetadataCache cache; + const SQLRETURN results[] = {SQL_SUCCESS, SQL_SUCCESS_WITH_INFO, SQL_NO_DATA, + SQL_ERROR, SQL_INVALID_HANDLE}; + for (SQLRETURN result : results) { + Populate(cache); + const auto before = cache.snapshot(); + { + ResultMetadataFailureGuard guard(cache, result); + } + const auto after = cache.snapshot(); + if (SQL_SUCCEEDED(result) || result == SQL_NO_DATA) { + assert(after.metadata == before.metadata); + assert(after.generation == before.generation); + } else { + assert(!after.metadata); + assert(after.generation != before.generation); + } + } + Populate(cache); + SQLRETURN result = SQL_SUCCESS; + try { + ResultMetadataFailureGuard guard(cache, result); + throw std::runtime_error("conversion failure"); + } catch (const std::runtime_error&) { + assert(!cache.snapshot().metadata); + } +} + +static void TestConcurrentInvalidation() { + ResultMetadataCache cache; + const auto metadata = MakeMetadata(); + std::thread invalidator([&] { + for (int i = 0; i < 1000; ++i) { + cache.clear(); + } + }); + for (int i = 0; i < 1000; ++i) { + const auto snapshot = cache.snapshot(); + cache.publish(snapshot.generation, metadata); + if (snapshot.metadata) { + assert(snapshot.metadata->columns.at(0).name == u"owned"); + } + } + invalidator.join(); + cache.clear(); + assert(!cache.snapshot().metadata); +} + +static void TestChildren() { + std::mutex childMutex; + auto first = std::make_shared(); + auto second = std::make_shared(); + auto unrelated = std::make_shared(); + std::vector> children{first, {}, second}; + Populate(first->resultMetadata); + Populate(second->resultMetadata); + Populate(unrelated->resultMetadata); + const auto held = first->resultMetadata.snapshot(); + ClearChildResultMetadata(childMutex, children); + assert(!first->resultMetadata.snapshot().metadata); + assert(!second->resultMetadata.snapshot().metadata); + assert(unrelated->resultMetadata.snapshot().metadata); + assert(held.metadata->columns.at(0).name == u"owned"); + ClearChildResultMetadata(childMutex, children); +} + +static void TestAllocationFailure() { + std::mutex childMutex; + auto owner = std::make_shared(); + observedHandle = owner; + std::vector> children{owner}; + Populate(owner->resultMetadata); + const auto before = owner->resultMetadata.snapshot(); + failAllocation = true; + try { + ClearChildResultMetadata(childMutex, children); + assert(false); + } catch (const std::bad_alloc&) { + assert(ownersAtFailure == 1); + assert(owner->resultMetadata.snapshot().metadata == before.metadata); + assert(childMutex.try_lock()); + childMutex.unlock(); + } + assert(!failAllocation); + ClearChildResultMetadata(childMutex, children); + assert(!owner->resultMetadata.snapshot().metadata); +} + +static void TestLastOwner() { + std::mutex childMutex; + auto owner = std::make_shared(); + owner->childMutex = &childMutex; + const std::weak_ptr weak = owner; + std::vector> children{owner}; + // Drop the external owner during invalidation, leaving only the helper's snapshot. + auto metadata = std::shared_ptr(new ResultMetadata, [&](auto* value) { + owner.reset(); + delete value; + }); + const auto generation = owner->resultMetadata.snapshot().generation; + owner->resultMetadata.publish(generation, std::move(metadata)); + ClearChildResultMetadata(childMutex, children); + assert(!owner && weak.expired()); +} + +int main(int argc, char** argv) { + if (argc != 2) { + std::fputs("Expected one native metadata test case\n", stderr); + return 2; + } + const char* name = argv[1]; + if (std::strcmp(name, "snapshots") == 0) { + TestSnapshots(); + } else if (std::strcmp(name, "failures") == 0) { + TestFailures(); + } else if (std::strcmp(name, "concurrent") == 0) { + TestConcurrentInvalidation(); + } else if (std::strcmp(name, "children") == 0) { + TestChildren(); + } else if (std::strcmp(name, "allocation") == 0) { + TestAllocationFailure(); + } else if (std::strcmp(name, "last_owner") == 0) { + TestLastOwner(); + } else { + std::fprintf(stderr, "Unknown native metadata test case: %s\n", name); + return 2; + } + std::printf("%s passed\n", name); + return 0; +} diff --git a/tests/test_040_fetch_native_metadata.py b/tests/test_040_fetch_native_metadata.py index 879ce4d92..102ad6f35 100644 --- a/tests/test_040_fetch_native_metadata.py +++ b/tests/test_040_fetch_native_metadata.py @@ -485,7 +485,7 @@ def test_result_metadata_transaction_preserved_cursor(metadata_cursor, operation else mssql_python.SQL_CURSOR_COMMIT_BEHAVIOR ) if connection.getinfo(info) != 2: # SQL_CB_PRESERVE - pytest.skip("Driver does not preserve cursors; native helper coverage is required") + pytest.skip("Driver does not preserve cursors; cache/helper coverage is in tests/native") cursor.execute(_query(["id"], 3)) _assert_rows([cursor.fetchone()], [(1,)]) if operation == "autocommit": @@ -734,15 +734,9 @@ def test_result_metadata_all_null_rows(tmp_path, method): scalar_null = [] scalar_status = native.DDBCSQLFetchOne(cursor.hstmt, scalar_null) diagnostics = native.DDBCSQLGetAllDiagRecords(cursor.hstmt) + assert scalar_status == 0, (scalar_status, diagnostics) assert scalar_null == [None] - if scalar_status == 0: - assert diagnostics == [] - recovery_descriptions = 0 - else: - assert scalar_status == -1 - assert len(diagnostics) == 1 and "22002" in diagnostics[0][0], diagnostics - assert "Indicator variable required but not supplied" in diagnostics[0][1] - recovery_descriptions = 4 if {method!r} == "many" else 2 + assert diagnostics == [] values = ",".join(f"({{i}})" for i in range(1, 16)) cursor.execute( "SELECT CASE WHEN id%7=0 THEN NULL ELSE id END AS c0," @@ -778,9 +772,7 @@ def test_result_metadata_all_null_rows(tmp_path, method): ] if profiling: stats = native.profiling.get_stats() - # Only drivers/builds reporting a real scalar NULL error - # require the additional post-error cache repopulations. - expected_describes = (16 if {method!r} == "one" else 17) + recovery_descriptions + expected_describes = 16 if {method!r} == "one" else 17 assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == expected_describes, stats assert stats["ddbc::sql_variant::null_probe"]["calls"] == 15 assert stats["ddbc::sql_variant::subtype"]["calls"] == 13 From 416e54d6ad7e39dd759d97088194a801fe8c3bee Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Tue, 22 Sep 2026 20:30:19 +0530 Subject: [PATCH 05/15] CHORE: Address native metadata test review findings Use C++ streams for test-runner output, document the line-scoped allocator rule exception, and trigger native invariant tests for production integration header and binding changes. Keep production fetch code and the allocation-failure probes unchanged. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .github/workflows/native-metadata-tests.yml | 6 ++++++ tests/native/allocation_failure.cpp | 3 ++- tests/native/result_metadata_tests.cpp | 8 ++++---- 3 files changed, 12 insertions(+), 5 deletions(-) diff --git a/.github/workflows/native-metadata-tests.yml b/.github/workflows/native-metadata-tests.yml index a5887fbfa..d0ea213be 100644 --- a/.github/workflows/native-metadata-tests.yml +++ b/.github/workflows/native-metadata-tests.yml @@ -5,6 +5,9 @@ on: types: [opened, reopened, synchronize, ready_for_review] paths: - 'mssql_python/pybind/result_metadata.hpp' + - 'mssql_python/pybind/ddbc_bindings.cpp' + - 'mssql_python/pybind/ddbc_bindings.h' + - 'mssql_python/pybind/connection/connection.h' - 'mssql_python/pybind/connection/connection.cpp' - 'tests/native/**' - '.github/workflows/native-metadata-tests.yml' @@ -12,6 +15,9 @@ on: branches: [main] paths: - 'mssql_python/pybind/result_metadata.hpp' + - 'mssql_python/pybind/ddbc_bindings.cpp' + - 'mssql_python/pybind/ddbc_bindings.h' + - 'mssql_python/pybind/connection/connection.h' - 'mssql_python/pybind/connection/connection.cpp' - 'tests/native/**' - '.github/workflows/native-metadata-tests.yml' diff --git a/tests/native/allocation_failure.cpp b/tests/native/allocation_failure.cpp index bf9daab61..72cf10cca 100644 --- a/tests/native/allocation_failure.cpp +++ b/tests/native/allocation_failure.cpp @@ -17,7 +17,8 @@ void* operator new(std::size_t size) { ownersAtFailure = observedHandle.use_count(); throw std::bad_alloc(); } - if (void* memory = std::malloc(size ? size : 1)) { + // operator new must not recurse; size is a byte count, with no arithmetic. + if (void* memory = std::malloc(size ? size : 1)) { // DevSkim: ignore DS161085 return memory; } throw std::bad_alloc(); diff --git a/tests/native/result_metadata_tests.cpp b/tests/native/result_metadata_tests.cpp index 963608564..699608a9e 100644 --- a/tests/native/result_metadata_tests.cpp +++ b/tests/native/result_metadata_tests.cpp @@ -7,8 +7,8 @@ #include "result_metadata.hpp" #include -#include #include +#include #include #include #include @@ -183,7 +183,7 @@ static void TestLastOwner() { int main(int argc, char** argv) { if (argc != 2) { - std::fputs("Expected one native metadata test case\n", stderr); + std::cerr << "Expected one native metadata test case\n"; return 2; } const char* name = argv[1]; @@ -200,9 +200,9 @@ int main(int argc, char** argv) { } else if (std::strcmp(name, "last_owner") == 0) { TestLastOwner(); } else { - std::fprintf(stderr, "Unknown native metadata test case: %s\n", name); + std::cerr << "Unknown native metadata test case: " << name << '\n'; return 2; } - std::printf("%s passed\n", name); + std::cout << name << " passed\n"; return 0; } From 7c7e6d2cbf238845b383152d23d25072480198ab Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Tue, 22 Sep 2026 21:50:01 +0530 Subject: [PATCH 06/15] PERF: Reuse native fetch buffers and ODBC bindings Preserve the bounded fetch-buffer follow-up on the original #796 dependency. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/connection/connection.cpp | 9 +- mssql_python/pybind/connection/connection.h | 2 +- mssql_python/pybind/ddbc_bindings.cpp | 411 +++++++++++++--- mssql_python/pybind/ddbc_bindings.h | 68 +-- mssql_python/pybind/fetch_bindings.hpp | 243 ++++++++++ tests/native/fetch_bindings_test.cpp | 397 +++++++++++++++ tests/test_041_fetch_buffer_reuse.py | 454 ++++++++++++++++++ 7 files changed, 1461 insertions(+), 123 deletions(-) create mode 100644 mssql_python/pybind/fetch_bindings.hpp create mode 100644 tests/native/fetch_bindings_test.cpp create mode 100644 tests/test_041_fetch_buffer_reuse.py diff --git a/mssql_python/pybind/connection/connection.cpp b/mssql_python/pybind/connection/connection.cpp index 9c4304d2f..8b1ac52af 100644 --- a/mssql_python/pybind/connection/connection.cpp +++ b/mssql_python/pybind/connection/connection.cpp @@ -115,7 +115,7 @@ void Connection::connect(const py::dict& attrs_before) { void Connection::disconnect(bool rollbackBeforeDisconnect) { PERF_TIMER("Connection::disconnect"); - clearResultMetadata(); + clearResultMetadata(false); // Determine GIL state once, up front. disconnect() runs both from // pybind11-bound methods (GIL held) and from GIL-less destructor / shutdown // paths: Connection::~Connection() dropping the last shared_ptr, or teardown @@ -169,10 +169,10 @@ void Connection::disconnect(bool rollbackBeforeDisconnect) { // Also cover children whose weak_ptr expired as their destructor // began waiting for this gate: they cannot appear in the snapshot. _cleanupState->disconnected = true; - std::lock_guard lock(_childHandlesMutex); for (const auto& handle : childHandles) { handle->markImplicitlyFreed(); } + std::lock_guard lock(_childHandlesMutex); _childStatementHandles.clear(); _allocationsSinceCompaction = 0; } @@ -266,7 +266,7 @@ void Connection::checkError(SQLRETURN ret) const { } } -void Connection::clearResultMetadata() { +void Connection::clearResultMetadata(bool detachFetchBindings) { std::vector handles; { std::lock_guard lock(_childHandlesMutex); @@ -281,6 +281,9 @@ void Connection::clearResultMetadata() { // Keep that destruction outside the child-list lock. for (const auto& handle : handles) { handle->resultMetadata.clear(); + if (detachFetchBindings && handle->fetchBindings.hasPlan()) { + handle->requireDetachedFetchBindings(); + } } } diff --git a/mssql_python/pybind/connection/connection.h b/mssql_python/pybind/connection/connection.h index f43613842..f1f327e75 100644 --- a/mssql_python/pybind/connection/connection.h +++ b/mssql_python/pybind/connection/connection.h @@ -103,7 +103,7 @@ class Connection { void allocateDbcHandle(); void checkError(SQLRETURN ret) const; void applyAttrsBefore(const py::dict& attrs_before); - void clearResultMetadata(); + void clearResultMetadata(bool detachFetchBindings = true); std::u16string _connStr; bool _fromPool = false; diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 13dfc9719..852a667f7 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -1542,14 +1542,71 @@ void DriverLoader::loadDriver() { } } +namespace { + +SQLRETURN BindFetchColumn(SQLHSTMT stmt, SQLUSMALLINT column, SQLSMALLINT type, + SQLPOINTER data, SQLLEN length, SQLLEN* indicators) { + PERF_TIMER("fetch_bindings::SQLBindCol"); + return SQLBindCol_ptr(stmt, column, type, data, length, indicators); +} + +SQLRETURN SetFetchAttribute(SQLHSTMT stmt, SQLINTEGER attribute, SQLPOINTER value, + SQLINTEGER length) { + if (attribute == SQL_ATTR_ROW_ARRAY_SIZE) { + PERF_TIMER("fetch_bindings::SQLSetStmtAttr::ROW_ARRAY_SIZE"); + return SQLSetStmtAttr_ptr(stmt, attribute, value, length); + } + PERF_TIMER("fetch_bindings::SQLSetStmtAttr::ROWS_FETCHED_PTR"); + return SQLSetStmtAttr_ptr(stmt, attribute, value, length); +} + +SQLRETURN GetFetchAttribute(SQLHSTMT stmt, SQLINTEGER attribute, SQLPOINTER value, + SQLINTEGER length, SQLINTEGER* returnedLength) { + PERF_TIMER("fetch_bindings::SQLGetStmtAttr"); + return SQLGetStmtAttr_ptr(stmt, attribute, value, length, returnedLength); +} + +SQLRETURN UnbindFetchColumns(SQLHSTMT stmt) { + PERF_TIMER("fetch_bindings::SQL_UNBIND"); + return SQLFreeStmt_ptr(stmt, SQL_UNBIND); +} + +inline SQLRETURN BeginResultTransition(const SqlHandlePtr& stmt) { + stmt->resultMetadata.clear(); + return stmt->detachFetchBindings(); +} + +void ThrowFetchCleanupError(SQLSMALLINT type, SQLHANDLE handle, SQLRETURN ret, + const char* operation) { + const auto error = SQLReadError(type, handle, ret); + std::string message = std::string(operation) + ": " + error.ddbcErrorMsg; + if (error.sqlState.size() == 5) { + message = "SQLSTATE:" + error.sqlState + ":" + message; + } + ThrowStdException(message); +} + +} // namespace + // SqlHandle definition SqlHandle::SqlHandle(SQLSMALLINT type, SQLHANDLE rawHandle, std::shared_ptr cleanupState) : _type(type), _handle(rawHandle), _cleanupState(std::move(cleanupState)) {} SqlHandle::~SqlHandle() { - if (_handle) { - free(); + try { + if (_handle) { + SQLRETURN ret = freeHandle(); + if (!SQL_SUCCEEDED(ret)) { + // A failed free leaves a live handle. Detach if possible before + // the last plan owner applies its native-only emergency policy. + SQLRETURN detached = detachFetchBindings(); + std::fprintf(stderr, "mssql-python: native handle cleanup failed (%d), " + "fetch buffer detach returned %d\n", ret, detached); + } + } + } catch (...) { + std::fputs("mssql-python: unexpected failure during native handle cleanup\n", stderr); } } @@ -1570,7 +1627,7 @@ SQLSMALLINT SqlHandle::type() const { void SqlHandle::markImplicitlyFreed() { // SAFETY: Only STMT handles should be marked as implicitly freed. - // When a DBC handle is freed, the ODBC driver automatically frees all child STMT handles. + // Successful SQLDisconnect frees the connection's child statements. // Other handle types (ENV, DBC, DESC) are NOT automatically freed by parents. // Calling this on wrong handle types will cause silent handle leaks. if (_type != SQL_HANDLE_STMT) { @@ -1582,6 +1639,7 @@ void SqlHandle::markImplicitlyFreed() { return; // Refuse to mark - let normal free() handle it } resultMetadata.clear(); + fetchBindings.nativeReleased(); _implicitly_freed = true; } @@ -1593,7 +1651,62 @@ void SqlHandle::markImplicitlyFreed() { * If you need destruction logs, use explicit close() methods instead. */ void SqlHandle::free() { - freeHandle(); + const bool hadFetchBindings = fetchBindings.hasPlan(); + SQLRETURN ret = freeHandle(); + if (hadFetchBindings && !SQL_SUCCEEDED(ret)) { + ThrowFetchCleanupError(_type, _handle, ret, "Freeing statement with retained fetch buffers"); + } +} + +SQLRETURN SqlHandle::detachFetchBindingsNative() { + auto plan = fetchBindings.snapshot(); + if (!plan) { + return SQL_SUCCESS; + } + if (_implicitly_freed || (_cleanupState && _cleanupState->disconnected)) { + fetchBindings.nativeReleased(); + return SQL_SUCCESS; + } + if (!_handle || !SQLFreeStmt_ptr || !SQLSetStmtAttr_ptr) { + return SQL_INVALID_HANDLE; + } + SQLRETURN ret = plan->detach(_handle, UnbindFetchColumns, SetFetchAttribute); + if (SQL_SUCCEEDED(ret)) { + fetchBindings.remove(plan); + } + return ret; +} + +SQLRETURN SqlHandle::detachPresentFetchBindings() { + auto plan = fetchBindings.snapshot(); + if (!plan) { + return SQL_SUCCESS; + } + if (is_python_finalizing()) { + return SQL_ERROR; + } + auto detachNative = [this]() { + auto cleanupLock = lockForCleanup(); + return detachFetchBindingsNative(); + }; + SQLRETURN ret; + if (PyGILState_Check()) { + py::gil_scoped_release release; + ret = detachNative(); + } else { + ret = detachNative(); + } + if (!SQL_SUCCEEDED(ret)) { + resultMetadata.clear(); + } + return ret; +} + +void SqlHandle::requireDetachedFetchBindings() { + SQLRETURN ret = detachFetchBindings(); + if (!SQL_SUCCEEDED(ret)) { + ThrowFetchCleanupError(_type, _handle, ret, "Detaching retained fetch buffers"); + } } SQLRETURN SqlHandle::freeHandle() { @@ -1621,11 +1734,13 @@ SQLRETURN SqlHandle::freeHandle() { describeCache.clear(); if (_implicitly_freed || (_cleanupState && _cleanupState->disconnected)) { _handle = nullptr; + fetchBindings.nativeReleased(); return SQL_SUCCESS; } SQLRETURN ret = SQLFreeHandle_ptr(_type, _handle); if (SQL_SUCCEEDED(ret)) { _handle = nullptr; + fetchBindings.nativeReleased(); } return ret; }; @@ -1652,6 +1767,12 @@ void SqlHandle::close_cursor() { if (!SQLFreeStmt_ptr) { ThrowStdException("SQLFreeStmt function not loaded"); } + if (fetchBindings.hasPlan()) { + SQLRETURN detached = detachFetchBindingsNative(); + if (!SQL_SUCCEEDED(detached)) { + return detached; + } + } return SQLFreeStmt_ptr(_handle, SQL_CLOSE); }; SQLRETURN ret; @@ -1662,7 +1783,7 @@ void SqlHandle::close_cursor() { ret = closeNative(); } if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { - ThrowStdException("SQLFreeStmt(SQL_CLOSE) failed"); + ThrowFetchCleanupError(_type, _handle, ret, "SQLFreeStmt(SQL_CLOSE)/fetch cleanup failed"); } } @@ -1702,7 +1823,9 @@ SQLRETURN SQLResetStmt_wrap(SqlHandlePtr statementHandle) { if (statementHandle->isImplicitlyFreed()) { return SQL_INVALID_HANDLE; } - statementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(statementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLFreeStmt_ptr) { DriverLoader::getInstance().loadDriver(); } @@ -1724,7 +1847,9 @@ SQLRETURN SQLResetStmt_wrap(SqlHandlePtr statementHandle) { SQLRETURN SQLGetTypeInfo_Wrapper(SqlHandlePtr StatementHandle, SQLSMALLINT DataType) { PERF_TIMER("SQLGetTypeInfo_Wrapper"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLGetTypeInfo_ptr) { ThrowStdException("SQLGetTypeInfo function not loaded"); } @@ -1737,7 +1862,9 @@ SQLRETURN SQLGetTypeInfo_Wrapper(SqlHandlePtr StatementHandle, SQLSMALLINT DataT SQLRETURN SQLProcedures_wrap(SqlHandlePtr StatementHandle, const py::object& catalogObj, const py::object& schemaObj, const py::object& procedureObj) { PERF_TIMER("SQLProcedures_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLProcedures_ptr) { ThrowStdException("SQLProcedures function not loaded"); } @@ -1762,7 +1889,9 @@ SQLRETURN SQLForeignKeys_wrap(SqlHandlePtr StatementHandle, const py::object& pk const py::object& fkCatalogObj, const py::object& fkSchemaObj, const py::object& fkTableObj) { PERF_TIMER("SQLForeignKeys_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLForeignKeys_ptr) { ThrowStdException("SQLForeignKeys function not loaded"); } @@ -1795,7 +1924,9 @@ SQLRETURN SQLForeignKeys_wrap(SqlHandlePtr StatementHandle, const py::object& pk SQLRETURN SQLPrimaryKeys_wrap(SqlHandlePtr StatementHandle, const py::object& catalogObj, const py::object& schemaObj, const std::u16string& table) { PERF_TIMER("SQLPrimaryKeys_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLPrimaryKeys_ptr) { ThrowStdException("SQLPrimaryKeys function not loaded"); } @@ -1818,7 +1949,9 @@ SQLRETURN SQLStatistics_wrap(SqlHandlePtr StatementHandle, const py::object& cat const py::object& schemaObj, const std::u16string& table, SQLUSMALLINT unique, SQLUSMALLINT reserved) { PERF_TIMER("SQLStatistics_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLStatistics_ptr) { ThrowStdException("SQLStatistics function not loaded"); } @@ -1841,7 +1974,9 @@ SQLRETURN SQLColumns_wrap(SqlHandlePtr StatementHandle, const py::object& catalo const py::object& schemaObj, const py::object& tableObj, const py::object& columnObj) { PERF_TIMER("SQLColumns_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLColumns_ptr) { ThrowStdException("SQLColumns function not loaded"); } @@ -1957,7 +2092,9 @@ py::list SQLGetAllDiagRecords(SqlHandlePtr handle) { // Wrap SQLExecDirect SQLRETURN SQLExecDirect_wrap(SqlHandlePtr StatementHandle, const std::u16string& Query) { PERF_TIMER("SQLExecDirect_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } LOG("SQLExecDirect: Executing query directly - statement_handle=%p, " "query_length=%zu chars", (void*)StatementHandle->get(), Query.length()); @@ -1994,7 +2131,9 @@ SQLRETURN SQLTables_wrap(SqlHandlePtr StatementHandle, const std::u16string& cat const std::u16string& schema, const std::u16string& table, const std::u16string& tableType) { PERF_TIMER("SQLTables_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLTables_ptr) { LOG("SQLTables: Function pointer not initialized, loading driver"); DriverLoader::getInstance().loadDriver(); @@ -2041,7 +2180,9 @@ SQLRETURN SQLExecute_wrap(const SqlHandlePtr statementHandle, return SQL_INVALID_HANDLE; } - statementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(statementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } SQLHANDLE hStmt = statementHandle->get(); // Configure forward-only / read-only cursor (matches slow path semantics). @@ -2842,7 +2983,9 @@ SQLRETURN SQLExecuteMany_wrap(const SqlHandlePtr statementHandle, const std::u16 std::vector& paramInfos, size_t paramSetSize, const py::dict& encodingSettings) { PERF_TIMER("SQLExecuteMany_wrap"); - statementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(statementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } LOG("SQLExecuteMany: Starting batch execution - param_count=%zu, " "param_set_size=%zu", columnwise_params.size(), paramSetSize); @@ -3101,8 +3244,8 @@ SQLRETURN DescribeColumns(SqlHandlePtr StatementHandle, AppendColumn&& appendCol } SQLRETURN GetResultMetadata(const SqlHandlePtr& statement, SQLSMALLINT columnCount, + const ResultMetadataCache::Snapshot& snapshot, std::shared_ptr& metadata) { - const auto snapshot = statement->resultMetadata.snapshot(); const bool matches = snapshot.metadata && columnCount >= 0 && snapshot.metadata->columns.size() == static_cast(columnCount); if (matches && snapshot.metadata->namesValidated) { @@ -3158,7 +3301,9 @@ SQLRETURN SQLSpecialColumns_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT ident const std::u16string& table, SQLSMALLINT scope, SQLSMALLINT nullable) { PERF_TIMER("SQLSpecialColumns_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } if (!SQLSpecialColumns_ptr) { ThrowStdException("SQLSpecialColumns function not loaded"); } @@ -3180,6 +3325,9 @@ SQLRETURN SQLSpecialColumns_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT ident // Wrap SQLFetch to retrieve rows SQLRETURN SQLFetch_wrap(SqlHandlePtr StatementHandle) { PERF_TIMER("SQLFetch_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(); !SQL_SUCCEEDED(ret)) { + return ret; + } LOG("SQLFetch: Fetching next row for statement_handle=%p", (void*)StatementHandle->get()); if (!SQLFetch_ptr) { LOG("SQLFetch: Function pointer not initialized, loading driver"); @@ -3382,6 +3530,9 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p const std::string& wcharEncoding = "utf-16le", int charCtype = SQL_C_WCHAR) { PERF_TIMER("SQLGetData_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(); !SQL_SUCCEEDED(ret)) { + return ret; + } // Note: wcharEncoding parameter is reserved for future use // Currently WCHAR data always uses UTF-16LE for Windows compatibility (void)wcharEncoding; // Suppress unused parameter warning @@ -4154,9 +4305,15 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri DriverLoader::getInstance().loadDriver(); // Load the driver } - // Unbind any columns from previous fetch operations to avoid memory - // corruption - SQLFreeStmt_ptr(StatementHandle->get(), SQL_UNBIND); + bool hadFetchPlan; + if (SQLRETURN ret = StatementHandle->detachFetchBindings(&hadFetchPlan); !SQL_SUCCEEDED(ret)) { + ThrowFetchCleanupError(SQL_HANDLE_STMT, StatementHandle->get(), ret, + "Detaching retained fetch buffers before scroll"); + return ret; + } + if (!hadFetchPlan) { + UnbindFetchColumns(StatementHandle->get()); + } // Perform scroll operation SQLRETURN ret = SQL_ERROR; @@ -4181,12 +4338,34 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri // For column in the result set, binds a buffer to retrieve column data // TODO: Move to anonymous namespace, since it is not used outside this file -template +template +void ResizeFetchBuffer(std::vector& buffer, size_t count) { +#ifdef ENABLE_PROFILING + if (count > buffer.capacity()) { + PERF_TIMER("fetch_bindings::column_buffer_allocation"); + buffer.resize(count); + return; + } +#endif + buffer.resize(count); +} + +template SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames, - SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR) { + SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR, + std::vector* bindings = nullptr) { PERF_TIMER("SQLBindColums"); SQLRETURN ret = SQL_SUCCESS; const bool useWideChar = (charCtype == SQL_C_WCHAR); + auto bindColumn = [bindings](SQLHSTMT stmt, SQLUSMALLINT column, SQLSMALLINT cType, + SQLPOINTER data, SQLLEN length, SQLLEN* indicators) -> SQLRETURN { + if constexpr (PrepareOnly) { + bindings->push_back({column, cType, data, length, indicators}); + return SQL_SUCCESS; + } else { + return BindFetchColumn(stmt, column, cType, data, length, indicators); + } + }; // Bind columns based on their data types for (SQLUSMALLINT col = 1; col <= numCols; col++) { const auto& columnMeta = GetFetchColumnMetadata(columnNames, col - 1); @@ -4202,8 +4381,8 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // Bind VARCHAR columns as SQL_C_WCHAR so the ODBC driver // returns UTF-16 data, avoiding code-page decode issues. uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; - buffers.wcharBuffers[col - 1].resize(fetchSize * fetchBufferSize); - ret = SQLBindCol_ptr( + ResizeFetchBuffer(buffers.wcharBuffers[col - 1], fetchSize * fetchBufferSize); + ret = bindColumn( hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), fetchBufferSize * sizeof(SQLWCHAR), buffers.indicators[col - 1].data()); } else { @@ -4213,8 +4392,8 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& #else uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; #endif - buffers.charBuffers[col - 1].resize(fetchSize * fetchBufferSize); - ret = SQLBindCol_ptr( + ResizeFetchBuffer(buffers.charBuffers[col - 1], fetchSize * fetchBufferSize); + ret = bindColumn( hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), fetchBufferSize * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); } @@ -4227,81 +4406,81 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // suffice HandleZeroColumnSizeAtFetch(columnSize); uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; - buffers.wcharBuffers[col - 1].resize(fetchSize * fetchBufferSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.wcharBuffers[col - 1], fetchSize * fetchBufferSize); + ret = bindColumn(hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), fetchBufferSize * sizeof(SQLWCHAR), buffers.indicators[col - 1].data()); break; } case SQL_INTEGER: - buffers.intBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_SLONG, buffers.intBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.intBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_SLONG, buffers.intBuffers[col - 1].data(), sizeof(SQLINTEGER), buffers.indicators[col - 1].data()); break; case SQL_SMALLINT: - buffers.smallIntBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_SSHORT, + ResizeFetchBuffer(buffers.smallIntBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_SSHORT, buffers.smallIntBuffers[col - 1].data(), sizeof(SQLSMALLINT), buffers.indicators[col - 1].data()); break; case SQL_TINYINT: - buffers.charBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_TINYINT, buffers.charBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.charBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_TINYINT, buffers.charBuffers[col - 1].data(), sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; case SQL_BIT: - buffers.charBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_BIT, buffers.charBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.charBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_BIT, buffers.charBuffers[col - 1].data(), sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; case SQL_REAL: - buffers.realBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_FLOAT, buffers.realBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.realBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_FLOAT, buffers.realBuffers[col - 1].data(), sizeof(SQLREAL), buffers.indicators[col - 1].data()); break; case SQL_DECIMAL: case SQL_NUMERIC: - buffers.charBuffers[col - 1].resize(fetchSize * MAX_DIGITS_IN_NUMERIC); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.charBuffers[col - 1], fetchSize * MAX_DIGITS_IN_NUMERIC); + ret = bindColumn(hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), MAX_DIGITS_IN_NUMERIC * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; case SQL_DOUBLE: case SQL_FLOAT: - buffers.doubleBuffers[col - 1].resize(fetchSize); + ResizeFetchBuffer(buffers.doubleBuffers[col - 1], fetchSize); ret = - SQLBindCol_ptr(hStmt, col, SQL_C_DOUBLE, buffers.doubleBuffers[col - 1].data(), + bindColumn(hStmt, col, SQL_C_DOUBLE, buffers.doubleBuffers[col - 1].data(), sizeof(SQLDOUBLE), buffers.indicators[col - 1].data()); break; case SQL_TIMESTAMP: case SQL_TYPE_TIMESTAMP: case SQL_DATETIME: - buffers.timestampBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr( + ResizeFetchBuffer(buffers.timestampBuffers[col - 1], fetchSize); + ret = bindColumn( hStmt, col, SQL_C_TYPE_TIMESTAMP, buffers.timestampBuffers[col - 1].data(), sizeof(SQL_TIMESTAMP_STRUCT), buffers.indicators[col - 1].data()); break; case SQL_BIGINT: - buffers.bigIntBuffers[col - 1].resize(fetchSize); + ResizeFetchBuffer(buffers.bigIntBuffers[col - 1], fetchSize); ret = - SQLBindCol_ptr(hStmt, col, SQL_C_SBIGINT, buffers.bigIntBuffers[col - 1].data(), + bindColumn(hStmt, col, SQL_C_SBIGINT, buffers.bigIntBuffers[col - 1].data(), sizeof(SQLBIGINT), buffers.indicators[col - 1].data()); break; case SQL_TYPE_DATE: - buffers.dateBuffers[col - 1].resize(fetchSize); + ResizeFetchBuffer(buffers.dateBuffers[col - 1], fetchSize); ret = - SQLBindCol_ptr(hStmt, col, SQL_C_TYPE_DATE, buffers.dateBuffers[col - 1].data(), + bindColumn(hStmt, col, SQL_C_TYPE_DATE, buffers.dateBuffers[col - 1].data(), sizeof(SQL_DATE_STRUCT), buffers.indicators[col - 1].data()); break; case SQL_SS_TIME2: - buffers.timeBuffers[col - 1].resize(fetchSize); + ResizeFetchBuffer(buffers.timeBuffers[col - 1], fetchSize); ret = - SQLBindCol_ptr(hStmt, col, SQL_C_SS_TIME2, buffers.timeBuffers[col - 1].data(), + bindColumn(hStmt, col, SQL_C_SS_TIME2, buffers.timeBuffers[col - 1].data(), sizeof(SQL_SS_TIME2_STRUCT), buffers.indicators[col - 1].data()); break; case SQL_GUID: - buffers.guidBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_GUID, buffers.guidBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.guidBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_GUID, buffers.guidBuffers[col - 1].data(), sizeof(SQLGUID), buffers.indicators[col - 1].data()); break; case SQL_SS_UDT: @@ -4311,13 +4490,13 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // TODO: handle variable length data correctly. This logic wont // suffice HandleZeroColumnSizeAtFetch(columnSize); - buffers.charBuffers[col - 1].resize(fetchSize * columnSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_BINARY, buffers.charBuffers[col - 1].data(), + ResizeFetchBuffer(buffers.charBuffers[col - 1], fetchSize * columnSize); + ret = bindColumn(hStmt, col, SQL_C_BINARY, buffers.charBuffers[col - 1].data(), columnSize, buffers.indicators[col - 1].data()); break; case SQL_SS_TIMESTAMPOFFSET: - buffers.datetimeoffsetBuffers[col - 1].resize(fetchSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_SS_TIMESTAMPOFFSET, + ResizeFetchBuffer(buffers.datetimeoffsetBuffers[col - 1], fetchSize); + ret = bindColumn(hStmt, col, SQL_C_SS_TIMESTAMPOFFSET, buffers.datetimeoffsetBuffers[col - 1].data(), sizeof(DateTimeOffset) * fetchSize, buffers.indicators[col - 1].data()); @@ -4346,12 +4525,12 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // Fetch rows in batches // TODO: Move to anonymous namespace, since it is not used outside this file -template +template SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames, py::list& rows, SQLUSMALLINT numCols, SQLULEN& numRowsFetched, const std::vector& lobColumns, const std::string& charEncoding = "utf-16le", - int charCtype = SQL_C_WCHAR) { + int charCtype = SQL_C_WCHAR, SQLULEN rowCapacity = 0) { PERF_TIMER("FetchBatchData"); LOG("FetchBatchData: Fetching data in batches"); SQLRETURN ret; @@ -4371,6 +4550,11 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& ret); return ret; } + if constexpr (CheckCapacity) { + if (numRowsFetched > rowCapacity) { + ThrowStdException("ODBC returned more rows than the bound fetch buffer capacity"); + } + } // Pre-cache column metadata to avoid repeated dictionary lookups. // The vectors below are consumed later by construct_rows, so they are // declared at function scope; only the population work is wrapped in the @@ -4821,13 +5005,26 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding); SQLRETURN ret = SQL_ERROR; ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); + if (fetchSize <= 0) { + ThrowStdException("Native fetchmany requires a positive fetch size"); + } + auto plan = StatementHandle->fetchBindings.snapshot(); + const auto metadataSnapshot = StatementHandle->resultMetadata.snapshot(); + if (plan && !plan->matches(metadataSnapshot, fetchSize, charEncoding, wcharEncoding, charCtype)) { + ret = StatementHandle->detachFetchBindings(); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + plan.reset(); + } SQLHSTMT hStmt = StatementHandle->get(); // Retrieve column count SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle); // Retrieve column metadata std::shared_ptr metadata; - ret = GetResultMetadata(StatementHandle, numCols, metadata); + const uint64_t metadataGeneration = metadataSnapshot.generation; + ret = GetResultMetadata(StatementHandle, numCols, metadataSnapshot, metadata); if (!SQL_SUCCEEDED(ret)) { LOG("FetchMany_wrap: Failed to get column descriptions - SQLRETURN=%d", ret); return ret; @@ -4853,6 +5050,12 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch SQLULEN numRowsFetched = 0; // If we have LOBs → fall back to row-by-row fetch + SQLGetData_wrap if (!lobColumns.empty()) { + if (plan) { + ret = StatementHandle->detachFetchBindings(); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } LOG("FetchMany_wrap: LOB columns detected (%zu columns), using per-row " "SQLGetData path", lobColumns.size()); @@ -4876,6 +5079,53 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch return SQL_SUCCESS; } + if (StatementHandle->fetchBindings.eligible()) { + const ResultMetadataCache::Snapshot snapshot{metadataGeneration, metadata}; + if (plan && plan->metadata != metadata) { + ret = StatementHandle->detachFetchBindings(); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + plan.reset(); + } + if (!plan) { + { + PERF_TIMER("fetch_bindings::plan_allocation"); + plan = std::shared_ptr( + new FetchBindingPlan(snapshot, fetchSize, charEncoding, wcharEncoding, charCtype), + FetchBindingPlan::Deleter{}); + } + ret = SQLBindColums(hStmt, plan->buffers, columnNames, numCols, fetchSize, charCtype, + &plan->bindings); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + StatementHandle->fetchBindings.install(plan); + ret = plan->attach(hStmt, BindFetchColumn, SetFetchAttribute, GetFetchAttribute); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } + plan->resetValues(); + ret = FetchBatchData(hStmt, plan->buffers, columnNames, rows, numCols, plan->rowsFetched, + lobColumns, charEncoding, charCtype, fetchSize); + if (ret == SQL_NO_DATA || + (SQL_SUCCEEDED(ret) && + StatementHandle->resultMetadata.snapshot().generation != metadataGeneration)) { + SQLRETURN detached = StatementHandle->detachFetchBindings(); + if (!SQL_SUCCEEDED(detached)) { + ret = detached; + } + } + return ret; + } + + if (plan) { + ret = StatementHandle->detachFetchBindings(); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } // Initialize column buffers ColumnBuffers buffers(numCols, fetchSize); @@ -5032,6 +5282,9 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, int arrowBatchSize, int charCtype) { PERF_TIMER("FetchArrowBatch_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(); !SQL_SUCCEEDED(ret)) { + return ret; + } // Fetch narrow char data as SQL_C_CHAR if on Linux/macOS and configured by the user charCtype = EffectiveCharCtypeForFetch(charCtype, "utf-8"); @@ -5944,6 +6197,9 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, const std::string& wcharEncoding = "utf-16le", int charCtype = SQL_C_WCHAR) { PERF_TIMER("FetchAll_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(); !SQL_SUCCEEDED(ret)) { + return ret; + } // Issue #531: upgrade SQL_C_CHAR + utf-8 to SQL_C_WCHAR on Windows so the // driver does lossless UTF-16 conversion instead of returning ACP bytes. charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding); @@ -6119,10 +6375,14 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row, ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); SQLHSTMT hStmt = StatementHandle->get(); - // Unbind any columns from previous fetch operations (e.g., fetchmany) - // to avoid conflicts with SQLGetData. SQLGetData cannot be used on - // columns that are already bound. - SQLFreeStmt_ptr(hStmt, SQL_UNBIND); + bool hadFetchPlan; + ret = StatementHandle->detachFetchBindings(&hadFetchPlan); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + if (!hadFetchPlan) { + UnbindFetchColumns(hStmt); + } // Assume hStmt is already allocated and a query has been executed { @@ -6148,7 +6408,9 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row, // Wrap SQLMoreResults SQLRETURN SQLMoreResults_wrap(SqlHandlePtr StatementHandle) { PERF_TIMER("SQLMoreResults_wrap"); - StatementHandle->resultMetadata.clear(); + if (SQLRETURN ret = BeginResultTransition(StatementHandle); !SQL_SUCCEEDED(ret)) { + return ret; + } LOG("SQLMoreResults_wrap: Check for more results"); if (!SQLMoreResults_ptr) { LOG("SQLMoreResults_wrap: Function pointer not initialized. Loading " @@ -6365,8 +6627,23 @@ PYBIND11_MODULE(ddbc_bindings, m) { "Set the decimal separator character"); m.def( "DDBCSQLSetStmtAttr", - [](SqlHandlePtr stmt, SQLINTEGER attr, py::object value) { - stmt->resultMetadata.clear(); + [](SqlHandlePtr stmt, SQLINTEGER attr, py::object value) -> SQLRETURN { + if (SQLRETURN ret = BeginResultTransition(stmt); !SQL_SUCCEEDED(ret)) { + return ret; + } + switch (attr) { + case SQL_ATTR_ROW_ARRAY_SIZE: + case SQL_ATTR_ROW_BIND_TYPE: + case SQL_ATTR_ROW_BIND_OFFSET_PTR: + case SQL_ATTR_ROWS_FETCHED_PTR: + case SQL_ATTR_ROW_STATUS_PTR: + case SQL_ATTR_APP_ROW_DESC: + case SQL_ATTR_USE_BOOKMARKS: + stmt->fetchBindings.disableReuse(); + break; + default: + break; + } SQLPOINTER ptr_value; if (py::isinstance(value)) { // For integer attributes like SQL_ATTR_QUERY_TIMEOUT diff --git a/mssql_python/pybind/ddbc_bindings.h b/mssql_python/pybind/ddbc_bindings.h index 3706e6e9d..1c5e4d61d 100644 --- a/mssql_python/pybind/ddbc_bindings.h +++ b/mssql_python/pybind/ddbc_bindings.h @@ -33,6 +33,7 @@ using py::literals::operator""_a; #include #include #include "result_metadata.hpp" +#include "fetch_bindings.hpp" //------------------------------------------------------------------------------------------------- // SQL Server specific ODBC constants @@ -296,6 +297,14 @@ class SqlHandle { SQLSMALLINT type() const; void free(); SQLRETURN freeHandle(); + SQLRETURN detachFetchBindings(bool* hadPlan = nullptr) { + const bool present = fetchBindings.hasPlan(); + if (hadPlan) { + *hadPlan = present; + } + return present ? detachPresentFetchBindings() : SQL_SUCCESS; + } + void requireDetachedFetchBindings(); void close_cursor(); // Cancel an in-progress statement (SQLCancel). Safe to call from a // thread other than the one running the fetch — this is the *only* @@ -306,18 +315,15 @@ class SqlHandle { void cancel(); bool isImplicitlyFreed() const { return _implicitly_freed; } - // Mark this handle as implicitly freed (freed by parent handle) - // This prevents double-free attempts when the ODBC driver automatically - // frees child handles (e.g., STMT handles when DBC handle is freed) + // Record proven native statement release by a successful parent disconnect. + // This is not a logical close: retained driver pointers become releasable. // // SAFETY CONSTRAINTS: // - ONLY call this on SQL_HANDLE_STMT handles - // - ONLY call this when the parent DBC handle is about to be freed + // - ONLY call after SQLDisconnect has actually succeeded // - Calling on other handle types (ENV, DBC, DESC) will cause HANDLE LEAKS - // - The ODBC spec only guarantees automatic freeing of STMT handles by DBC parents // - // Current usage: Connection::disconnect() marks all tracked STMT handles - // before freeing the DBC handle. + // Connection::disconnect() calls this before freeing the DBC wrapper. void markImplicitlyFreed(); // GH-610: Per-handle SQLDescribeParam result cache. @@ -328,10 +334,13 @@ class SqlHandle { std::unordered_map describeCache; void clearDescribeCache() { describeCache.clear(); } ResultMetadataCache resultMetadata; + FetchBindingSlot fetchBindings; private: // The caller must release the GIL before waiting for native cleanup. std::unique_lock lockForCleanup() const; + SQLRETURN detachPresentFetchBindings(); + SQLRETURN detachFetchBindingsNative(); SQLSMALLINT _type; SQLHANDLE _handle; bool _implicitly_freed = false; // Tracks if handle was freed by parent @@ -398,51 +407,6 @@ void DDBCSetDecimalSeparator(const std::string& separator); // (Used internally by ddbc_bindings.cpp - not part of public API) //------------------------------------------------------------------------------------------------- -// Struct to hold the SQL Server TIME2 structure (SQL_C_SS_TIME2) -struct SQL_SS_TIME2_STRUCT { - SQLUSMALLINT hour; - SQLUSMALLINT minute; - SQLUSMALLINT second; - SQLUINTEGER fraction; // Nanoseconds -}; - -// Struct to hold the DateTimeOffset structure -struct DateTimeOffset { - SQLSMALLINT year; - SQLUSMALLINT month; - SQLUSMALLINT day; - SQLUSMALLINT hour; - SQLUSMALLINT minute; - SQLUSMALLINT second; - SQLUINTEGER fraction; // Nanoseconds - SQLSMALLINT timezone_hour; // Offset hours from UTC - SQLSMALLINT timezone_minute; // Offset minutes from UTC -}; - -// Struct to hold data buffers and indicators for each column -struct ColumnBuffers { - std::vector> charBuffers; - std::vector> wcharBuffers; - std::vector> intBuffers; - std::vector> smallIntBuffers; - std::vector> realBuffers; - std::vector> doubleBuffers; - std::vector> timestampBuffers; - std::vector> bigIntBuffers; - std::vector> dateBuffers; - std::vector> timeBuffers; - std::vector> guidBuffers; - std::vector> indicators; - std::vector> datetimeoffsetBuffers; - - ColumnBuffers(SQLSMALLINT numCols, int fetchSize) - : charBuffers(numCols), wcharBuffers(numCols), intBuffers(numCols), - smallIntBuffers(numCols), realBuffers(numCols), doubleBuffers(numCols), - timestampBuffers(numCols), bigIntBuffers(numCols), dateBuffers(numCols), - timeBuffers(numCols), guidBuffers(numCols), datetimeoffsetBuffers(numCols), - indicators(numCols, std::vector(fetchSize)) {} -}; - // Performance: Column processor function type for fast type conversion // Using function pointers eliminates switch statement overhead in the hot loop typedef void (*ColumnProcessor)(PyObject* row, ColumnBuffers& buffers, const void* colInfo, diff --git a/mssql_python/pybind/fetch_bindings.hpp b/mssql_python/pybind/fetch_bindings.hpp new file mode 100644 index 000000000..3768e729f --- /dev/null +++ b/mssql_python/pybind/fetch_bindings.hpp @@ -0,0 +1,243 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#ifdef _WIN32 +#include +#endif +#include +#include +#include "result_metadata.hpp" + +struct SQL_SS_TIME2_STRUCT { + SQLUSMALLINT hour; + SQLUSMALLINT minute; + SQLUSMALLINT second; + SQLUINTEGER fraction; // Nanoseconds. +}; + +struct DateTimeOffset { + SQLSMALLINT year; + SQLUSMALLINT month; + SQLUSMALLINT day; + SQLUSMALLINT hour; + SQLUSMALLINT minute; + SQLUSMALLINT second; + SQLUINTEGER fraction; // Nanoseconds. + SQLSMALLINT timezone_hour; + SQLSMALLINT timezone_minute; +}; + +struct ColumnBuffers { + std::vector> charBuffers; + std::vector> wcharBuffers; + std::vector> intBuffers; + std::vector> smallIntBuffers; + std::vector> realBuffers; + std::vector> doubleBuffers; + std::vector> timestampBuffers; + std::vector> bigIntBuffers; + std::vector> dateBuffers; + std::vector> timeBuffers; + std::vector> guidBuffers; + std::vector> indicators; + std::vector> datetimeoffsetBuffers; + + ColumnBuffers(SQLSMALLINT numCols, int fetchSize) + : charBuffers(numCols), wcharBuffers(numCols), intBuffers(numCols), + smallIntBuffers(numCols), realBuffers(numCols), doubleBuffers(numCols), + timestampBuffers(numCols), bigIntBuffers(numCols), dateBuffers(numCols), + timeBuffers(numCols), guidBuffers(numCols), + indicators(numCols, std::vector(fetchSize)), + datetimeoffsetBuffers(numCols) {} +}; + +struct FetchColumnBinding { + SQLUSMALLINT column; + SQLSMALLINT cType; + SQLPOINTER data; + SQLLEN bufferLength; + SQLLEN* indicators; +}; + +// Only the statement's fetch operation mutates a plan. Cancellation invalidates +// the metadata generation instead; a shared lease protects lifetime, not mutation. +class FetchBindingPlan { + public: + FetchBindingPlan(ResultMetadataCache::Snapshot snapshot, int size, + std::string charEncoding, std::string wcharEncoding, int charCtype) + : metadata(std::move(snapshot.metadata)), generation(snapshot.generation), + fetchSize(size), charEncoding(std::move(charEncoding)), + wcharEncoding(std::move(wcharEncoding)), charCtype(charCtype), + buffers(static_cast(metadata->columns.size()), size) { + bindings.reserve(metadata->columns.size()); + } + + bool matches(const ResultMetadataCache::Snapshot& snapshot, int size, + const std::string& charCodec, const std::string& wcharCodec, int cType) const { + return reusable && driverMayReference.load() && generation == snapshot.generation && + metadata == snapshot.metadata && fetchSize == size && + charEncoding == charCodec && wcharEncoding == wcharCodec && charCtype == cType; + } + + template + SQLRETURN attach(SQLHSTMT stmt, Bind bind, Set set, Get get) { + reusable = false; + needsReset = true; + SQLRETURN ret = set(stmt, SQL_ATTR_ROW_ARRAY_SIZE, + reinterpret_cast(static_cast(fetchSize)), 0); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + SQLULEN activeSize = 0; + ret = get(stmt, SQL_ATTR_ROW_ARRAY_SIZE, &activeSize, 0, nullptr); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + if (activeSize != static_cast(fetchSize)) { + throw std::runtime_error("ODBC changed the requested fetch row-array size"); + } + driverMayReference = true; + ret = set(stmt, SQL_ATTR_ROWS_FETCHED_PTR, &rowsFetched, 0); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + for (const auto& column : bindings) { + ret = bind(stmt, column.column, column.cType, column.data, column.bufferLength, + column.indicators); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } + reusable = true; + return ret; + } + + template + SQLRETURN detach(SQLHSTMT stmt, Unbind unbind, Set set) { + reusable = false; + if (!needsReset) { + return SQL_SUCCESS; + } + SQLRETURN ret = unbind(stmt); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + ret = set(stmt, SQL_ATTR_ROWS_FETCHED_PTR, nullptr, 0); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + driverMayReference = false; + ret = set(stmt, SQL_ATTR_ROW_ARRAY_SIZE, reinterpret_cast(1), 0); + if (SQL_SUCCEEDED(ret)) { + needsReset = false; + } + return ret; + } + + void resetValues() { + rowsFetched = 0; + for (auto& column : buffers.indicators) { + std::fill(column.begin(), column.end(), SQL_NULL_DATA); + } + } + + void nativeReleased() noexcept { driverMayReference = false; } + + struct Deleter { + void operator()(FetchBindingPlan* plan) const noexcept { + if (plan->driverMayReference.load()) { + // Final owner only: freeing this allocation could leave driver + // pointers dangling after failed native cleanup or finalization. + std::fputs("mssql-python: retaining fetch buffers after unconfirmed native " + "cleanup until process exit\n", stderr); + return; + } + delete plan; + } + }; + + const std::shared_ptr metadata; + const uint64_t generation; + const int fetchSize; + const std::string charEncoding; + const std::string wcharEncoding; + const int charCtype; + ColumnBuffers buffers; + std::vector bindings; + SQLULEN rowsFetched = 0; + + private: + bool reusable = false; + bool needsReset = false; + std::atomic driverMayReference{false}; +}; + +class FetchBindingSlot { + public: + bool hasPlan() const noexcept { return hasPlan_.load(std::memory_order_acquire); } + + std::shared_ptr snapshot() const { + if (!hasPlan()) { + return {}; + } + std::lock_guard lock(mutex_); + return plan_; + } + + void install(const std::shared_ptr& plan) { + std::lock_guard lock(mutex_); + if (plan_) { + throw std::logic_error("Fetch bindings must be detached before replacement"); + } + plan_ = plan; + hasPlan_.store(true, std::memory_order_release); + } + + void remove(const std::shared_ptr& expected) { + std::shared_ptr retired; + { + std::lock_guard lock(mutex_); + if (plan_ == expected) { + retired = std::move(plan_); + hasPlan_.store(false, std::memory_order_release); + } + } + } + + void nativeReleased() { + if (!hasPlan()) { + return; + } + std::shared_ptr retired; + { + std::lock_guard lock(mutex_); + retired = std::move(plan_); + hasPlan_.store(false, std::memory_order_release); + } + if (retired) { + retired->nativeReleased(); + } + } + + bool eligible() const { return eligible_.load(); } + + void disableReuse() { eligible_ = false; } + + private: + mutable std::mutex mutex_; + std::shared_ptr plan_; + std::atomic hasPlan_{false}; + std::atomic eligible_{true}; +}; diff --git a/tests/native/fetch_bindings_test.cpp b/tests/native/fetch_bindings_test.cpp new file mode 100644 index 000000000..883d8a6a4 --- /dev/null +++ b/tests/native/fetch_bindings_test.cpp @@ -0,0 +1,397 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +// Compile as C++17 with mssql_python/pybind on the include path. No driver or +// Python initialization: these tests exercise the production binding helper. +#ifdef NDEBUG +#error "fetch_bindings_test requires assertions enabled" +#endif + +#include "fetch_bindings.hpp" +#include +#include +#include +#include + +namespace { +size_t allocationCalls = 0; +size_t liveAllocations = 0; +long failAllocationAfter = -1; +} + +void* operator new(std::size_t size) { + if (failAllocationAfter == 0) { + throw std::bad_alloc(); + } + if (failAllocationAfter > 0) { + --failAllocationAfter; + } + void* pointer = std::malloc(size ? size : 1); + if (!pointer) { + throw std::bad_alloc(); + } + ++allocationCalls; + ++liveAllocations; + return pointer; +} + +void operator delete(void* pointer) noexcept { + if (pointer) { + --liveAllocations; + } + std::free(pointer); +} + +void operator delete(void* pointer, std::size_t) noexcept { ::operator delete(pointer); } + +namespace { + +struct Driver { + int calls = 0; + int failCall = -1; + int diagnostic = 0; + int bindCalls = 0; + int unbindCalls = 0; + int attributeCalls = 0; + SQLULEN rowArraySize = 1; + SQLULEN* rowsFetched = nullptr; + bool substituteSize = false; + std::array bound{}; + + SQLRETURN result() { + ++calls; + diagnostic = calls == failCall ? 9000 + calls : 0; + return calls == failCall ? SQL_ERROR : SQL_SUCCESS; + } + + SQLRETURN set(SQLHSTMT, SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER) { + ++attributeCalls; + if (attribute == SQL_ATTR_ROW_ARRAY_SIZE) { + rowArraySize = static_cast(reinterpret_cast(value)); + } else { + assert(attribute == SQL_ATTR_ROWS_FETCHED_PTR); + rowsFetched = static_cast(value); + } + // Deliberately retain the supplied address even when setup fails. + return result(); + } + + SQLRETURN get(SQLHSTMT, SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER, SQLINTEGER*) { + assert(attribute == SQL_ATTR_ROW_ARRAY_SIZE); + *static_cast(value) = rowArraySize + (substituteSize ? 1 : 0); + return result(); + } + + SQLRETURN bind(SQLHSTMT, SQLUSMALLINT column, SQLSMALLINT type, SQLPOINTER data, + SQLLEN length, SQLLEN* indicators) { + ++bindCalls; + bound.at(column - 1) = {column, type, data, length, indicators}; + return result(); + } + + SQLRETURN unbind(SQLHSTMT) { + ++unbindCalls; + SQLRETURN ret = result(); + if (SQL_SUCCEEDED(ret)) { + bound = {}; + } + return ret; + } + + SQLRETURN attach(FetchBindingPlan& plan) { + return plan.attach( + nullptr, + [this](auto... args) { return bind(args...); }, + [this](auto... args) { return set(args...); }, + [this](auto... args) { return get(args...); }); + } + + SQLRETURN detach(FetchBindingPlan& plan) { + return plan.detach( + nullptr, [this](auto stmt) { return unbind(stmt); }, + [this](auto... args) { return set(args...); }); + } + + void fetch(SQLULEN count) { + assert(rowsFetched); + assert(count <= rowArraySize); + *rowsFetched = count; + for (SQLULEN i = 0; i < count; ++i) { + static_cast(bound[0].data)[i] = static_cast(i + 41); + bound[0].indicators[i] = sizeof(SQLINTEGER); + static_cast(bound[1].data)[i * 17] = 'x'; + bound[1].indicators[i] = sizeof(SQLWCHAR); + } + } +}; + +ResultMetadataCache::Snapshot metadata() { + auto value = std::make_shared(); + value->namesValidated = true; + value->columns = {{u"id", SQL_INTEGER, 10, 0, SQL_NULLABLE}, + {u"name", SQL_WVARCHAR, 16, 0, SQL_NULLABLE}}; + return {42, std::move(value)}; +} + +std::shared_ptr makePlan(const ResultMetadataCache::Snapshot& snapshot, + int size = 2) { + auto plan = std::shared_ptr( + new FetchBindingPlan(snapshot, size, "utf-16le", "utf-16le", SQL_C_WCHAR), + FetchBindingPlan::Deleter{}); + plan->buffers.intBuffers[0].resize(size); + plan->buffers.wcharBuffers[1].resize(size * 17); + plan->bindings.push_back({1, SQL_C_SLONG, plan->buffers.intBuffers[0].data(), + sizeof(SQLINTEGER), plan->buffers.indicators[0].data()}); + plan->bindings.push_back({2, SQL_C_WCHAR, plan->buffers.wcharBuffers[1].data(), + 17 * sizeof(SQLWCHAR), plan->buffers.indicators[1].data()}); + return plan; +} + +void allocationFailuresBeforeBinding() { + auto snapshot = metadata(); + const size_t before = liveAllocations; + bool reachedSuccess = false; + for (long failure = 0; failure < 100; ++failure) { + failAllocationAfter = failure; + try { + auto plan = makePlan(snapshot); + failAllocationAfter = -1; + reachedSuccess = true; + } catch (const std::bad_alloc&) { + failAllocationAfter = -1; + } + assert(liveAllocations == before); + if (reachedSuccess) { + break; + } + } + assert(reachedSuccess); +} + +void compatibleHitsDoNotAllocateOrBind() { + auto snapshot = metadata(); + FetchBindingSlot slot; + auto plan = makePlan(snapshot); + slot.install(plan); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + const auto* data = driver.bound[0].data; + const auto* indicators = driver.bound[0].indicators; + const auto* fetched = driver.rowsFetched; + const size_t allocationsBefore = allocationCalls; + for (int i = 0; i < 10000; ++i) { + auto lease = slot.snapshot(); + assert(lease->matches(snapshot, 2, "utf-16le", "utf-16le", SQL_C_WCHAR)); + lease->resetValues(); + assert(lease->rowsFetched == 0); + assert(lease->buffers.indicators[0][1] == SQL_NULL_DATA); + driver.fetch(i % 2 + 1); + assert(driver.bound[0].data == data); + assert(driver.bound[0].indicators == indicators); + assert(driver.rowsFetched == fetched); + } + assert(allocationCalls == allocationsBefore); + assert(driver.bindCalls == 2); + assert(driver.attributeCalls == 2); + assert(driver.unbindCalls == 0); + assert(!plan->matches(snapshot, 1, "utf-16le", "utf-16le", SQL_C_WCHAR)); + assert(!plan->matches(snapshot, 2, "utf-8", "utf-16le", SQL_C_WCHAR)); + assert(!plan->matches(snapshot, 2, "utf-16le", "utf-8", SQL_C_WCHAR)); + assert(!plan->matches(snapshot, 2, "utf-16le", "utf-16le", SQL_C_CHAR)); + auto changed = snapshot; + ++changed.generation; + assert(!plan->matches(changed, 2, "utf-16le", "utf-16le", SQL_C_WCHAR)); + changed = metadata(); + assert(!plan->matches(changed, 2, "utf-16le", "utf-16le", SQL_C_WCHAR)); + assert(driver.detach(*plan) == SQL_SUCCESS); + slot.remove(plan); + assert(driver.rowsFetched == nullptr); + assert(driver.rowArraySize == 1); + assert(!slot.snapshot()); +} + +void partialSetupPreservesOwnershipAndDiagnostics() { + for (int fail = 1; fail <= 5; ++fail) { + auto snapshot = metadata(); + FetchBindingSlot slot; + auto plan = makePlan(snapshot); + slot.install(plan); + Driver driver; + driver.failCall = fail; + assert(driver.attach(*plan) == SQL_ERROR); + assert(driver.calls == fail); + assert(driver.diagnostic == 9000 + fail); + assert(slot.snapshot() == plan); + assert(!plan->matches(snapshot, 2, "utf-16le", "utf-16le", SQL_C_WCHAR)); + driver.failCall = -1; + assert(driver.detach(*plan) == SQL_SUCCESS); + slot.remove(plan); + } +} + +void partialDetachNeverReleasesStorage() { + for (int fail = 1; fail <= 3; ++fail) { + auto plan = makePlan(metadata()); + FetchBindingSlot slot; + slot.install(plan); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + driver.failCall = driver.calls + fail; + assert(driver.detach(*plan) == SQL_ERROR); + assert(driver.diagnostic == 9000 + driver.failCall); + assert(slot.snapshot() == plan); + if (driver.bound[0].data && driver.rowsFetched) { + driver.fetch(1); + assert(plan->buffers.intBuffers[0][0] == 41); + } + driver.failCall = -1; + assert(driver.detach(*plan) == SQL_SUCCESS); + slot.remove(plan); + } +} + +void releaseWaitsForConversionLease() { + auto plan = makePlan(metadata()); + FetchBindingSlot slot; + slot.install(plan); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + driver.fetch(1); + std::weak_ptr weak = plan; + // Simulate the notification sent ONLY after native free/disconnect succeeds. + slot.nativeReleased(); + assert(!slot.snapshot()); + assert(!weak.expired()); + assert(plan->buffers.intBuffers[0][0] == 41); + plan.reset(); + assert(weak.expired()); +} + +void finalOwnerRetainsUnconfirmedDriverPointers() { + auto plan = makePlan(metadata()); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + auto* retained = plan.get(); + failAllocationAfter = 0; + plan.reset(); + failAllocationAfter = -1; + driver.fetch(1); + assert(retained->buffers.intBuffers[0][0] == 41); + // The test owns the deliberately abandoned raw allocation solely to reclaim + // it after proving that a driver can still access its original pointers. + retained->nativeReleased(); + delete retained; +} + +void substitutedSizeCannotFetch() { + auto plan = makePlan(metadata()); + Driver driver; + driver.substituteSize = true; + bool threw = false; + try { + driver.attach(*plan); + } catch (const std::runtime_error&) { + threw = true; + } + assert(threw); + assert(driver.bindCalls == 0); + assert(driver.rowsFetched == nullptr); + assert(driver.detach(*plan) == SQL_SUCCESS); +} + +void noEmergencyRetentionWithoutDriverPointers() { + const auto snapshot = metadata(); + const size_t before = liveAllocations; + for (int fail = 1; fail <= 2; ++fail) { + auto plan = makePlan(snapshot); + Driver driver; + driver.failCall = fail; + assert(driver.attach(*plan) == SQL_ERROR); + assert(driver.rowsFetched == nullptr); + assert(driver.bindCalls == 0); + plan.reset(); + assert(liveAllocations == before); + } + auto plan = makePlan(snapshot); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + driver.failCall = driver.calls + 3; + assert(driver.detach(*plan) == SQL_ERROR); + assert(driver.rowsFetched == nullptr); + assert(driver.bound[0].data == nullptr); + // Failed size restoration blocks the next operation but cannot justify + // abandoning storage after all retained addresses have been cleared. + plan.reset(); + assert(liveAllocations == before); +} + +void errorAndCancellationInvalidateWithoutFreeingBuffers() { + for (bool exception : {false, true}) { + ResultMetadataCache cache; + auto snapshot = metadata(); + cache.publish(0, snapshot.metadata); + snapshot = cache.snapshot(); + auto plan = makePlan(snapshot); + FetchBindingSlot slot; + slot.install(plan); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + SQLRETURN ret = SQL_SUCCESS; + try { + ResultMetadataFailureGuard failure(cache, ret); + if (exception) { + throw std::runtime_error("conversion failed"); + } + ret = SQL_ERROR; + } catch (const std::runtime_error&) { + } + assert(cache.snapshot().generation != snapshot.generation); + assert(slot.snapshot() == plan); + assert(!plan->matches(cache.snapshot(), 2, "utf-16le", "utf-16le", SQL_C_WCHAR)); + driver.fetch(1); + assert(plan->buffers.intBuffers[0][0] == 41); + assert(driver.detach(*plan) == SQL_SUCCESS); + slot.remove(plan); + } + ResultMetadataCache cache; + auto snapshot = metadata(); + cache.publish(0, snapshot.metadata); + snapshot = cache.snapshot(); + auto plan = makePlan(snapshot); + Driver driver; + assert(driver.attach(*plan) == SQL_SUCCESS); + cache.clear(); + driver.fetch(1); + assert(plan->buffers.intBuffers[0][0] == 41); + assert(!plan->matches(cache.snapshot(), 2, "utf-16le", "utf-16le", SQL_C_WCHAR)); + assert(driver.detach(*plan) == SQL_SUCCESS); +} + +void emptySlotsDoNotAllocate() { + FetchBindingSlot slot; + const size_t before = allocationCalls; + failAllocationAfter = 0; + for (int i = 0; i < 10000; ++i) { + assert(!slot.hasPlan()); + assert(!slot.snapshot()); + slot.nativeReleased(); + } + failAllocationAfter = -1; + assert(allocationCalls == before); +} + +} // namespace + +int main() { + allocationFailuresBeforeBinding(); + compatibleHitsDoNotAllocateOrBind(); + partialSetupPreservesOwnershipAndDiagnostics(); + partialDetachNeverReleasesStorage(); + releaseWaitsForConversionLease(); + finalOwnerRetainsUnconfirmedDriverPointers(); + substitutedSizeCannotFetch(); + noEmergencyRetentionWithoutDriverPointers(); + errorAndCancellationInvalidateWithoutFreeingBuffers(); + emptySlotsDoNotAllocate(); + assert(liveAllocations == 0); +} diff --git a/tests/test_041_fetch_buffer_reuse.py b/tests/test_041_fetch_buffer_reuse.py new file mode 100644 index 000000000..30e0fbb8f --- /dev/null +++ b/tests/test_041_fetch_buffer_reuse.py @@ -0,0 +1,454 @@ +"""Bounded fetchmany storage must survive reuse, transitions, and failed conversion. + +Attribute/unbind counters cover retained-plan helpers, not every legacy ODBC call. +Their absence on another route is not evidence of zero driver work. +""" + +import os +from pathlib import Path +import subprocess +import sys +import textwrap +from decimal import Decimal + +import pytest + +import mssql_python as db +from mssql_python import ddbc_bindings as native + + +@pytest.fixture +def reuse_cursor(conn_str): + with db.connect(conn_str) as connection: + with connection.cursor() as cursor: + yield cursor + + +def _query(count, width=2): + columns = ",".join(f"CAST(id AS INT) AS c{i}" for i in range(width)) + return ( + f"WITH n AS (SELECT TOP({count}) ROW_NUMBER() OVER " + "(ORDER BY a.object_id,b.object_id) AS id " + f"FROM sys.all_objects a CROSS JOIN sys.all_objects b) SELECT {columns} FROM n ORDER BY id" + ) + + +def _tuples(rows): + return [tuple(row) for row in rows] + + +def _isolated(script, tmp_path): + environment = dict(os.environ) + root = str(Path(db.__file__).resolve().parent.parent) + environment["PYTHONPATH"] = os.pathsep.join([root, environment.get("PYTHONPATH", "")]) + result = subprocess.run( + [sys.executable, "-c", textwrap.dedent(script)], + cwd=tmp_path, + env=environment, + capture_output=True, + text=True, + timeout=60, + ) + assert result.returncode == 0, result.stdout + result.stderr + + +@pytest.mark.parametrize("size", [1, 2, 1000]) +def test_repeated_size_partial_batch_and_eof(reuse_cursor, size): + cursor = reuse_cursor + count = size * 2 + 1 + cursor.execute(_query(count)) + assert _tuples(cursor.fetchmany(size)) == [(i, i) for i in range(1, size + 1)] + assert cursor.fetchmany(0) == [] + assert cursor.fetchmany(-1) == [] + assert _tuples(cursor.fetchmany(size)) == [(i, i) for i in range(size + 1, size * 2 + 1)] + assert _tuples(cursor.fetchmany(size)) == [(count, count)] + assert cursor.fetchmany(size) == [] + assert cursor.fetchmany(size) == [] + assert cursor.fetchone() is None + assert cursor.fetchall() == [] + + +def test_size_changes_and_arraysize_do_not_skip_rows(reuse_cursor): + cursor = reuse_cursor + cursor.execute(_query(2015)) + consumed = 0 + for size in (1, 1, 2, 2, 1000, 1000, 2, 1, 1): + batch = cursor.fetchmany(size) + assert _tuples(batch) == [(i, i) for i in range(consumed + 1, consumed + len(batch) + 1)] + consumed += len(batch) + cursor.arraysize = 2 + while batch := cursor.fetchmany(): + assert _tuples(batch) == [(i, i) for i in range(consumed + 1, consumed + len(batch) + 1)] + consumed += len(batch) + assert consumed == 2015 + + +def test_nulls_and_variable_values_overwrite_old_rows(reuse_cursor): + cursor = reuse_cursor + cursor.execute( + "SELECT id, CASE WHEN id%2=0 THEN CAST(NULL AS INT) ELSE id END AS n," + "CASE WHEN id%2=0 THEN CAST(NULL AS NVARCHAR(30)) " + "ELSE REPLICATE(N'x',id) END AS s," + "CASE WHEN id%2=0 THEN CAST(NULL AS VARBINARY(10)) ELSE 0x0102 END AS b," + "CASE WHEN id%2=0 THEN CAST(NULL AS DECIMAL(10,2)) " + "ELSE CAST(id AS DECIMAL(10,2)) END AS d " + "FROM (VALUES(1),(2),(3),(4),(5)) source(id) ORDER BY id" + ) + rows = [] + while batch := cursor.fetchmany(2): + rows.extend(batch) + assert _tuples(rows) == [ + (i, None, None, None, None) if i % 2 == 0 else (i, i, "x" * i, b"\x01\x02", Decimal(i)) + for i in range(1, 6) + ] + + +@pytest.mark.parametrize("mode", ["one", "scroll", "direct", "all", "arrow", "arrow_schema"]) +def test_many_transitions_preserve_cursor_position(reuse_cursor, mode): + cursor = reuse_cursor + if mode.startswith("arrow"): + pytest.importorskip("pyarrow") + cursor.execute(_query(9)) + assert _tuples(cursor.fetchmany(2)) == [(1, 1), (2, 2)] + assert _tuples(cursor.fetchmany(2)) == [(3, 3), (4, 4)] + if mode == "one": + assert tuple(cursor.fetchone()) == (5, 5) + elif mode == "scroll": + cursor.scroll(1) + elif mode == "direct": + assert native.DDBCSQLFetch(cursor.hstmt) in (0, 1) + row = [] + assert native.DDBCSQLGetData( + cursor.hstmt, 2, row, "utf-16le", "utf-16le", db.SQL_WCHAR + ) in (0, 1) + assert row == [5, 5] + elif mode == "all": + assert _tuples(cursor.fetchall()) == [(i, i) for i in range(5, 10)] + return + elif mode == "arrow": + assert cursor.arrow_batch(1).to_pydict() == {"c0": [5], "c1": [5]} + else: + assert cursor.arrow_batch(0).num_rows == 0 + assert _tuples(cursor.fetchmany(1)) == [(5, 5)] + assert _tuples(cursor.fetchmany(2)) == [(6, 6), (7, 7)] + assert _tuples(cursor.fetchall()) == [(8, 8), (9, 9)] + + +@pytest.mark.parametrize("distance", [1, 3]) +def test_scroll_cleanup_exception_preserves_public_position(reuse_cursor, monkeypatch, distance): + """Mock the bridge exception, not an actual ODBC cleanup failure.""" + cursor = reuse_cursor + cursor.execute(_query(8)) + assert _tuples(cursor.fetchmany(2)) == [(1, 1), (2, 2)] + position = (cursor._rownumber, cursor._next_row_index, cursor.rowcount) + statement = cursor.hstmt + failure = RuntimeError("SQLSTATE:HY010:Detaching retained fetch buffers before scroll: unbind") + calls = [] + + def fail_cleanup(*arguments): + calls.append(arguments) + raise failure + + with monkeypatch.context() as patch: + patch.setattr(native, "DDBCSQLFetchScroll", fail_cleanup) + with pytest.raises(IndexError, match="SQLSTATE:HY010:") as raised: + cursor.scroll(distance) + assert raised.value.__cause__ is failure + + assert len(calls) == 1 + assert cursor.hstmt is statement + assert (cursor._rownumber, cursor._next_row_index, cursor.rowcount) == position + assert _tuples(cursor.fetchmany(1)) == [(3, 3)] + + +def test_nextset_same_width_changes_type_and_size(reuse_cursor): + cursor = reuse_cursor + cursor.execute( + "SELECT id, CAST('a' AS VARCHAR(1)) AS text_value " + "FROM (VALUES(1),(2),(3)) source(id) ORDER BY id;" + "SELECT CAST(id+0.25 AS DECIMAL(9,2)) AS amount," + "CAST(REPLICATE(N'z',30) AS NVARCHAR(40)) AS long_value " + "FROM (VALUES(4),(5),(6)) source(id) ORDER BY id" + ) + assert _tuples(cursor.fetchmany(1)) == [(1, "a")] + assert _tuples(cursor.fetchmany(1)) == [(2, "a")] + assert cursor.nextset() + assert _tuples(cursor.fetchmany(1)) == [(Decimal("4.25"), "z" * 30)] + assert _tuples(cursor.fetchmany(1)) == [(Decimal("5.25"), "z" * 30)] + assert _tuples(cursor.fetchmany(1)) == [(Decimal("6.25"), "z" * 30)] + assert cursor.fetchmany(1) == [] + + +@pytest.mark.parametrize("reset_cursor", [True, False]) +def test_prepared_handle_reexecution_replaces_plan(reuse_cursor, reset_cursor): + cursor = reuse_cursor + query = "SELECT CAST(? AS INT)+id AS n FROM (VALUES(1),(2),(3)) source(id) ORDER BY id" + cursor.execute(query, 10) + statement = cursor.hstmt + assert _tuples(cursor.fetchmany(1)) == [(11,)] + assert _tuples(cursor.fetchmany(1)) == [(12,)] + cursor.execute(query, 20, reset_cursor=reset_cursor) + assert cursor.hstmt is statement + assert _tuples(cursor.fetchmany(1)) == [(21,)] + assert _tuples(cursor.fetchmany(1)) == [(22,)] + assert _tuples(cursor.fetchall()) == [(23,)] + + +def test_converter_failure_preserves_fallback_and_position(reuse_cursor): + cursor = reuse_cursor + cursor.execute(_query(4, 1)) + assert _tuples(cursor.fetchmany(1)) == [(1,)] + calls = [] + + def fail(value): + calls.append(value) + raise LookupError(f"converter rejected {value}") + + cursor.connection.add_output_converter(db.SQL_INTEGER, fail) + assert _tuples(cursor.fetchmany(1)) == [(2,)] + assert calls == [2] + cursor.connection.remove_output_converter(db.SQL_INTEGER) + assert _tuples(cursor.fetchmany(1)) == [(3,)] + assert _tuples(cursor.fetchmany(1)) == [(4,)] + + +def test_native_conversion_failure_recovers_without_transition(tmp_path): + """A cached constructor raises inside native conversion on only the middle row.""" + _isolated( + """ + import datetime + import os + + original_date = datetime.date + calls = [] + + class ConstructorFailure(Exception): + pass + + failure = ConstructorFailure("middle-row date conversion") + + class ValueCheckedDate(original_date): + def __new__(cls, year, month, day): + calls.append(year) + if year == 2002: + raise failure + return original_date.__new__(cls, year, month, day) + + datetime.date = ValueCheckedDate + import mssql_python as db + + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute( + "SELECT CAST(value AS DATE) FROM " + "(VALUES(1,'2001-01-01'),(2,'2002-01-01'),(3,'2003-01-01')) " + "source(id,value) ORDER BY id" + ) + statement = cursor.hstmt + calls.clear() + assert cursor.fetchmany(1)[0][0].year == 2001 + try: + cursor.fetchmany(1) + except ConstructorFailure as error: + assert error is failure + else: + raise AssertionError("Native constructor failure was swallowed") + assert cursor.hstmt is statement + assert cursor.fetchmany(1)[0][0].year == 2003 + assert calls == [2001, 2002, 2003] + assert cursor.fetchmany(1) == [] + """, + tmp_path, + ) + + +def test_decoding_change_on_live_bound_cursor(reuse_cursor): + cursor = reuse_cursor + cursor.execute( + "SELECT CAST('value' AS VARCHAR(20)) AS s " + "FROM (VALUES(1),(2),(3),(4)) source(id) ORDER BY id" + ) + assert _tuples(cursor.fetchmany(1)) == [("value",)] + assert _tuples(cursor.fetchmany(1)) == [("value",)] + cursor.connection.setdecoding(db.SQL_CHAR, encoding="utf-8", ctype=db.SQL_CHAR) + assert _tuples(cursor.fetchmany(1)) == [("value",)] + assert _tuples(cursor.fetchmany(1)) == [("value",)] + + +def test_external_statement_attribute_detaches_before_update(reuse_cursor): + cursor = reuse_cursor + cursor.execute(_query(7)) + assert _tuples(cursor.fetchmany(2)) == [(1, 1), (2, 2)] + # SQL_ATTR_QUERY_TIMEOUT does not change layout, but is an external mutation. + assert native.DDBCSQLSetStmtAttr(cursor.hstmt, 0, 0) in (0, 1) + assert _tuples(cursor.fetchmany(2)) == [(3, 3), (4, 4)] + assert _tuples(cursor.fetchmany(2)) == [(5, 5), (6, 6)] + assert _tuples(cursor.fetchmany(2)) == [(7, 7)] + + +@pytest.mark.parametrize("operation", ["commit", "rollback", "autocommit"]) +def test_transaction_transition_with_live_bindings(reuse_cursor, operation): + cursor = reuse_cursor + connection = cursor.connection + info = ( + db.SQL_CURSOR_ROLLBACK_BEHAVIOR + if operation == "rollback" + else db.SQL_CURSOR_COMMIT_BEHAVIOR + ) + if connection.getinfo(info) != 2: + pytest.skip("SQL_CB_PRESERVE is required for the live-cursor transition") + cursor.execute(_query(5)) + assert _tuples(cursor.fetchmany(2)) == [(1, 1), (2, 2)] + if operation == "autocommit": + connection.autocommit = True + else: + getattr(connection, operation)() + assert _tuples(cursor.fetchmany(2)) == [(3, 3), (4, 4)] + assert _tuples(cursor.fetchmany(2)) == [(5, 5)] + + +@pytest.mark.parametrize("close_mode", ["statement", "connection", "gc"]) +def test_bound_statement_lifetime_isolated(tmp_path, close_mode): + _isolated( + f""" + import gc + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + db.pooling(enabled=False) + for _ in range(5): + connection = db.connect(os.environ["DB_CONNECTION_STRING"], autocommit=True) + statement = connection._conn.alloc_statement_handle() + assert native.DDBCSQLExecDirect( + statement, "SELECT id FROM (VALUES(1),(2),(3),(4)) s(id) ORDER BY id" + ) in (0, 1) + for expected in (1, 2): + rows = [] + assert native.DDBCSQLFetchMany(statement, rows, 1) in (0, 1) + assert rows == [[expected]] + if {close_mode!r} == "statement": + statement.free() + connection.close() + elif {close_mode!r} == "connection": + connection._conn.close() + statement.free() + connection._conn = None + connection.close() + else: + del statement + gc.collect() + connection.close() + """, + tmp_path, + ) + + +@pytest.mark.skipif(not hasattr(native, "profiling"), reason="requires native operation counters") +@pytest.mark.parametrize("size,width", [(1, 24), (2, 3), (1000, 24)]) +def test_mechanism_actual_binding_and_allocation_counts(tmp_path, size, width): + count = 10000 if size == 1 else size * 2 + 1 + query = _query(count, width) + _isolated( + f""" + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute({query!r}) + native.profiling.reset() + native.profiling.enable() + try: + consumed = 0 + while batch := cursor.fetchmany({size}): + assert [tuple(row) for row in batch] == [ + (i,) * {width} for i in range(consumed+1, consumed+len(batch)+1) + ] + consumed += len(batch) + assert consumed == {count} + finally: + native.profiling.disable() + stats = native.profiling.get_stats() + def calls(name): + return stats.get("ddbc::" + name, {{}}).get("calls", 0) + assert calls("fetch_bindings::plan_allocation") == 1, stats + assert calls("fetch_bindings::column_buffer_allocation") == {width}, stats + assert calls("fetch_bindings::SQLBindCol") == {width}, stats + assert calls("fetch_bindings::SQL_UNBIND") == 1, stats + assert calls("fetch_bindings::SQLSetStmtAttr::ROW_ARRAY_SIZE") == 2, stats + assert calls("fetch_bindings::SQLSetStmtAttr::ROWS_FETCHED_PTR") == 2, stats + assert calls("fetch_bindings::SQLGetStmtAttr") == 1, stats + expected_fetches = {((count + size - 1) // size) + 1} + assert calls("FetchBatchData::SQLFetchScroll_call") == expected_fetches, stats + assert calls("SQLDescribeCol::driver_call") == {width}, stats + """, + tmp_path, + ) + + +@pytest.mark.skipif(not hasattr(native, "profiling"), reason="requires native operation counters") +def test_mechanism_size_and_mode_changes_rebind_only_when_needed(tmp_path): + _isolated( + f""" + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute({_query(12)!r}) + native.profiling.reset() + native.profiling.enable() + try: + for size, first in ((1,1),(1,2),(2,3),(2,5)): + assert [tuple(r) for r in cursor.fetchmany(size)] == [ + (i,i) for i in range(first,first+size) + ] + assert tuple(cursor.fetchone()) == (7,7) + assert [tuple(r) for r in cursor.fetchmany(2)] == [(8,8),(9,9)] + assert [tuple(r) for r in cursor.fetchmany(2)] == [(10,10),(11,11)] + finally: + native.profiling.disable() + stats = native.profiling.get_stats() + assert stats["ddbc::fetch_bindings::plan_allocation"]["calls"] == 3, stats + assert stats["ddbc::fetch_bindings::SQLBindCol"]["calls"] == 6, stats + assert stats["ddbc::fetch_bindings::SQL_UNBIND"]["calls"] == 2, stats + assert stats["ddbc::fetch_bindings::SQLSetStmtAttr::ROW_ARRAY_SIZE"]["calls"] == 5 + assert stats["ddbc::fetch_bindings::SQLSetStmtAttr::ROWS_FETCHED_PTR"]["calls"] == 5 + """, + tmp_path, + ) + + +@pytest.mark.skipif(not hasattr(native, "profiling"), reason="requires native operation counters") +@pytest.mark.parametrize("kind", ["lob", "variant"]) +def test_mechanism_fallback_does_not_retain_a_bound_plan(tmp_path, kind): + expression = ( + "CAST(REPLICATE(N'x',30) AS NVARCHAR(MAX))" + if kind == "lob" + else "CAST(CASE WHEN id=2 THEN NULL ELSE id END AS SQL_VARIANT)" + ) + query = f"SELECT {expression} FROM (VALUES(1),(2),(3)) source(id) ORDER BY id" + expected = [("x" * 30,)] * 3 if kind == "lob" else [(1,), (None,), (3,)] + _isolated( + f""" + import os + import mssql_python as db + from mssql_python import ddbc_bindings as native + with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: + with connection.cursor() as cursor: + cursor.execute({query!r}) + native.profiling.reset() + native.profiling.enable() + try: + rows = [] + while batch := cursor.fetchmany(2): + rows.extend(tuple(row) for row in batch) + assert rows == {expected!r} + finally: + native.profiling.disable() + stats = native.profiling.get_stats() + assert stats.get("ddbc::fetch_bindings::plan_allocation", {{}}).get("calls", 0) == 0 + assert stats.get("ddbc::fetch_bindings::SQLBindCol", {{}}).get("calls", 0) == 0 + """, + tmp_path, + ) From 65081f060ee878bebb1ad0a595b93ddbf2418d72 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Thu, 24 Sep 2026 11:03:11 +0530 Subject: [PATCH 07/15] REFACTOR: Minimize metadata optimization scope Remove PR-only tests, native test build and CI wiring, and restore the base test guide. Inline the test-only invalidation helper and single-caller metadata setup; stop retaining unused precision/nullability fields while preserving the cache lifetime, generation, Unicode, variant and failure contracts. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .github/prompts/run-tests.prompt.md | 21 - .github/workflows/native-metadata-tests.yml | 55 -- mssql_python/pybind/connection/connection.cpp | 16 +- mssql_python/pybind/ddbc_bindings.cpp | 89 +- mssql_python/pybind/result_metadata.hpp | 22 - tests/native/CMakeLists.txt | 31 - tests/native/allocation_failure.cpp | 33 - tests/native/result_metadata_tests.cpp | 208 ----- tests/test_040_fetch_native_metadata.py | 816 ------------------ 9 files changed, 51 insertions(+), 1240 deletions(-) delete mode 100644 .github/workflows/native-metadata-tests.yml delete mode 100644 tests/native/CMakeLists.txt delete mode 100644 tests/native/allocation_failure.cpp delete mode 100644 tests/native/result_metadata_tests.cpp delete mode 100644 tests/test_040_fetch_native_metadata.py diff --git a/.github/prompts/run-tests.prompt.md b/.github/prompts/run-tests.prompt.md index 2d6eb3a9b..da8bcfa88 100644 --- a/.github/prompts/run-tests.prompt.md +++ b/.github/prompts/run-tests.prompt.md @@ -122,27 +122,6 @@ Help the developer run tests to validate their changes. Follow this process base ## STEP 1: Choose What to Test -### Native metadata invariants (no database) - -The standalone CMake tests in `tests/native` exercise the production metadata -cache and child-handle invalidation helper without importing the Python package -or connecting to SQL Server. They require a C++17 compiler, CMake, and ODBC -headers (Windows SDK, `unixodbc-dev` on Linux, or `unixodbc` on macOS). -The Native Metadata Tests workflow runs them on Windows, Linux, and macOS. - -```bash -cmake -S tests/native -B build/native-metadata -DCMAKE_BUILD_TYPE=Release -cmake --build build/native-metadata --config Release --parallel 2 -ctest --test-dir build/native-metadata -C Release --output-on-failure -``` - -Assertions remain enabled in Release. Cases cover stale-generation rejection, -held snapshots, failure/EOF guards, concurrent invalidation, child isolation, -reserve failure before strong-reference acquisition, and last-owner destruction -outside the child-list lock. These native checks supplement, not replace, the -live transaction tests, which skip when the driver does not preserve cursors. -Native-only tests do not require the Python-test prerequisites above. - ### Test Categories | Category | Description | When to Use | diff --git a/.github/workflows/native-metadata-tests.yml b/.github/workflows/native-metadata-tests.yml deleted file mode 100644 index d0ea213be..000000000 --- a/.github/workflows/native-metadata-tests.yml +++ /dev/null @@ -1,55 +0,0 @@ -name: Native Metadata Tests - -on: - pull_request: - types: [opened, reopened, synchronize, ready_for_review] - paths: - - 'mssql_python/pybind/result_metadata.hpp' - - 'mssql_python/pybind/ddbc_bindings.cpp' - - 'mssql_python/pybind/ddbc_bindings.h' - - 'mssql_python/pybind/connection/connection.h' - - 'mssql_python/pybind/connection/connection.cpp' - - 'tests/native/**' - - '.github/workflows/native-metadata-tests.yml' - push: - branches: [main] - paths: - - 'mssql_python/pybind/result_metadata.hpp' - - 'mssql_python/pybind/ddbc_bindings.cpp' - - 'mssql_python/pybind/ddbc_bindings.h' - - 'mssql_python/pybind/connection/connection.h' - - 'mssql_python/pybind/connection/connection.cpp' - - 'tests/native/**' - - '.github/workflows/native-metadata-tests.yml' - -permissions: - contents: read - -jobs: - native-metadata: - name: Native metadata (${{ matrix.os }}) - runs-on: ${{ matrix.os }} - timeout-minutes: 10 - strategy: - fail-fast: false - matrix: - os: [ubuntu-latest, windows-latest, macos-latest] - steps: - - name: Checkout - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 - with: - persist-credentials: false - - name: Install ODBC headers (Linux) - if: runner.os == 'Linux' - run: | - sudo apt-get update - sudo apt-get install -y unixodbc-dev - - name: Install ODBC headers (macOS) - if: runner.os == 'macOS' - run: brew install unixodbc - - name: Configure - run: cmake -S tests/native -B build/native-metadata -DCMAKE_BUILD_TYPE=Release - - name: Build - run: cmake --build build/native-metadata --config Release --parallel 2 - - name: Test - run: ctest --test-dir build/native-metadata -C Release --output-on-failure diff --git a/mssql_python/pybind/connection/connection.cpp b/mssql_python/pybind/connection/connection.cpp index de4eb45c0..9c4304d2f 100644 --- a/mssql_python/pybind/connection/connection.cpp +++ b/mssql_python/pybind/connection/connection.cpp @@ -267,7 +267,21 @@ void Connection::checkError(SQLRETURN ret) const { } void Connection::clearResultMetadata() { - ClearChildResultMetadata(_childHandlesMutex, _childStatementHandles); + std::vector handles; + { + std::lock_guard lock(_childHandlesMutex); + handles.reserve(_childStatementHandles.size()); + for (const auto& weakHandle : _childStatementHandles) { + if (auto handle = weakHandle.lock()) { + handles.push_back(std::move(handle)); + } + } + } + // Releasing the last handle can acquire the connection cleanup gate. + // Keep that destruction outside the child-list lock. + for (const auto& handle : handles) { + handle->resultMetadata.clear(); + } } void Connection::commit() { diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index cf2dc5c0f..5068f3c19 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3101,42 +3101,6 @@ SQLRETURN DescribeColumns(SqlHandlePtr StatementHandle, AppendColumn&& appendCol return SQL_SUCCESS; } -SQLRETURN GetResultMetadata(const SqlHandlePtr& statement, SQLSMALLINT columnCount, - std::shared_ptr& metadata) { - const auto snapshot = statement->resultMetadata.snapshot(); - const bool matches = snapshot.metadata && columnCount >= 0 && - snapshot.metadata->columns.size() == static_cast(columnCount); - if (matches && snapshot.metadata->namesValidated) { - metadata = snapshot.metadata; - return SQL_SUCCESS; - } - auto pending = matches ? std::make_shared(*snapshot.metadata) - : std::make_shared(); - if (!matches) { - SQLRETURN ret = DescribeColumns( - statement, [&](std::u16string name, SQLSMALLINT type, SQLULEN size, - SQLSMALLINT digits, SQLSMALLINT nullable) { - // Preserve eager name validation before advancing the result set. - py::cast(name); - pending->columns.push_back( - {std::move(name), type, type == SQL_SS_VARIANT ? 0 : size, digits, nullable}); - }); - if (!SQL_SUCCEEDED(ret)) { - return ret; - } - } else { - // Row-wise fetches originally read names without decoding them. A later - // many/all fetch must still validate those names before its first advance. - for (const auto& column : pending->columns) { - py::cast(column.name); - } - } - pending->namesValidated = true; - statement->resultMetadata.publish(snapshot.generation, pending); - metadata = std::move(pending); - return SQL_SUCCESS; -} - } // namespace // Wrap SQLDescribeCol @@ -3446,8 +3410,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p dupeSqlWCharAsUtf16Le( uncachedColumnName, std::min(static_cast(columnNameLen), std::size(uncachedColumnName) - 1)), - dataType, dataType == SQL_SS_VARIANT ? 0 : columnSize, decimalDigits, - nullable}); + dataType, dataType == SQL_SS_VARIANT ? 0 : columnSize}); } } @@ -4827,12 +4790,37 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle); // Retrieve column metadata - std::shared_ptr metadata; - ret = GetResultMetadata(StatementHandle, numCols, metadata); - if (!SQL_SUCCEEDED(ret)) { - LOG("FetchMany_wrap: Failed to get column descriptions - SQLRETURN=%d", ret); - return ret; + auto snapshot = StatementHandle->resultMetadata.snapshot(); + auto metadata = std::move(snapshot.metadata); + const bool matches = metadata && numCols >= 0 && + metadata->columns.size() == static_cast(numCols); + if (!matches || !metadata->namesValidated) { + auto pending = matches ? std::make_shared(*metadata) + : std::make_shared(); + if (!matches) { + ret = DescribeColumns( + StatementHandle, [&](std::u16string name, SQLSMALLINT type, SQLULEN size, + SQLSMALLINT, SQLSMALLINT) { + // Preserve eager name validation before advancing the result set. + py::cast(name); + pending->columns.push_back( + {std::move(name), type, type == SQL_SS_VARIANT ? 0 : size}); + }); + if (!SQL_SUCCEEDED(ret)) { + LOG("FetchMany_wrap: Failed to get column descriptions - SQLRETURN=%d", ret); + return ret; + } + } else { + // Row-wise fetches read names without decoding them. Validate before advancing. + for (const auto& column : pending->columns) { + py::cast(column.name); + } + } + pending->namesValidated = true; + StatementHandle->resultMetadata.publish(snapshot.generation, pending); + metadata = std::move(pending); } + ret = SQL_SUCCESS; const auto& columnNames = metadata->columns; if (numCols < 0 || columnNames.size() != static_cast(numCols)) { LOG("FetchMany_wrap: Column metadata count does not match result column count"); @@ -4841,11 +4829,8 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch std::vector lobColumns; for (SQLSMALLINT i = 0; i < numCols; i++) { - const auto& colMeta = GetFetchColumnMetadata(columnNames, i); - SQLSMALLINT dataType = GetFetchColumnType(colMeta); - SQLULEN columnSize = GetFetchColumnSize(colMeta); - - if (IsLobOrVariantColumn(dataType, columnSize)) { + const auto& column = columnNames.at(i); + if (IsLobOrVariantColumn(column.dataType, column.columnSize)) { lobColumns.push_back(i + 1); // 1-based } } @@ -5991,13 +5976,11 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, metadata->namesValidated = true; metadata->columns.reserve(numCols); for (SQLSMALLINT i = 0; i < numCols; ++i) { - const auto column = GetFetchColumnMetadata(columnNames, i); - SQLSMALLINT type = GetFetchColumnType(column); + const auto column = columnNames[i].cast(); + SQLSMALLINT type = column["DataType"].cast(); metadata->columns.push_back({ column["ColumnName"].cast(), type, - type == SQL_SS_VARIANT ? 0 : GetFetchColumnSize(column), - column["DecimalDigits"].cast(), - column["Nullable"].cast()}); + type == SQL_SS_VARIANT ? 0 : column["ColumnSize"].cast()}); } StatementHandle->resultMetadata.publish(metadataSnapshot.generation, std::move(metadata)); while (true) { diff --git a/mssql_python/pybind/result_metadata.hpp b/mssql_python/pybind/result_metadata.hpp index 0e08e28bd..3a9912860 100644 --- a/mssql_python/pybind/result_metadata.hpp +++ b/mssql_python/pybind/result_metadata.hpp @@ -16,8 +16,6 @@ struct FetchColumnMetadata { std::u16string name; SQLSMALLINT dataType; SQLULEN columnSize; - SQLSMALLINT decimalDigits; - SQLSMALLINT nullable; }; struct ResultMetadata { @@ -58,26 +56,6 @@ class ResultMetadataCache { std::shared_ptr metadata_; }; -template -void ClearChildResultMetadata(std::mutex& childHandlesMutex, - const std::vector>& childHandles) { - std::vector> handles; - { - std::lock_guard lock(childHandlesMutex); - handles.reserve(childHandles.size()); - for (const auto& weakHandle : childHandles) { - if (auto handle = weakHandle.lock()) { - handles.push_back(std::move(handle)); - } - } - } - // Releasing the last handle can acquire the connection cleanup gate. - // Keep that destruction outside the child-list lock. - for (const auto& handle : handles) { - handle->resultMetadata.clear(); - } -} - class ResultMetadataFailureGuard { public: ResultMetadataFailureGuard(ResultMetadataCache& cache, const SQLRETURN& result) diff --git a/tests/native/CMakeLists.txt b/tests/native/CMakeLists.txt deleted file mode 100644 index 83c8f3861..000000000 --- a/tests/native/CMakeLists.txt +++ /dev/null @@ -1,31 +0,0 @@ -cmake_minimum_required(VERSION 3.15) -project(mssql_python_native_tests LANGUAGES CXX) - -enable_testing() -find_package(Threads REQUIRED) - -add_executable(result_metadata_tests result_metadata_tests.cpp allocation_failure.cpp) -target_compile_features(result_metadata_tests PRIVATE cxx_std_17) -target_include_directories(result_metadata_tests PRIVATE ../../mssql_python/pybind) -target_link_libraries(result_metadata_tests PRIVATE Threads::Threads) - -if(WIN32) - target_compile_definitions(result_metadata_tests PRIVATE WIN32_LEAN_AND_MEAN NOMINMAX) -else() - find_path(ODBC_INCLUDE_DIR sql.h PATHS /opt/homebrew/include /usr/local/include) - if(NOT ODBC_INCLUDE_DIR) - message(FATAL_ERROR "ODBC headers are required: install unixodbc-dev or unixodbc.") - endif() - target_include_directories(result_metadata_tests PRIVATE "${ODBC_INCLUDE_DIR}") -endif() - -if(MSVC) - target_compile_options(result_metadata_tests PRIVATE /W4 /WX /UNDEBUG) -else() - target_compile_options(result_metadata_tests PRIVATE -Wall -Wextra -Werror -UNDEBUG) -endif() - -foreach(case_name IN ITEMS snapshots failures concurrent children allocation last_owner) - add_test(NAME result_metadata_${case_name} COMMAND result_metadata_tests ${case_name}) - set_tests_properties(result_metadata_${case_name} PROPERTIES TIMEOUT 20) -endforeach() diff --git a/tests/native/allocation_failure.cpp b/tests/native/allocation_failure.cpp deleted file mode 100644 index 72cf10cca..000000000 --- a/tests/native/allocation_failure.cpp +++ /dev/null @@ -1,33 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -#include -#include -#include - -struct TestHandle; -extern std::weak_ptr observedHandle; -extern bool failAllocation; -extern long ownersAtFailure; - -// Keep replacement allocation functions opaque to optimized test call sites. -void* operator new(std::size_t size) { - if (failAllocation) { - failAllocation = false; - ownersAtFailure = observedHandle.use_count(); - throw std::bad_alloc(); - } - // operator new must not recurse; size is a byte count, with no arithmetic. - if (void* memory = std::malloc(size ? size : 1)) { // DevSkim: ignore DS161085 - return memory; - } - throw std::bad_alloc(); -} - -void operator delete(void* memory) noexcept { - std::free(memory); -} - -void operator delete(void* memory, std::size_t) noexcept { - std::free(memory); -} diff --git a/tests/native/result_metadata_tests.cpp b/tests/native/result_metadata_tests.cpp deleted file mode 100644 index 699608a9e..000000000 --- a/tests/native/result_metadata_tests.cpp +++ /dev/null @@ -1,208 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -#ifdef _WIN32 -#include -#endif -#include "result_metadata.hpp" - -#include -#include -#include -#include -#include -#include - -#ifdef NDEBUG -#error Native metadata tests require assertions, including Release builds. -#endif - -struct TestHandle { - ResultMetadataCache resultMetadata; - std::mutex* childMutex = nullptr; - - ~TestHandle() { - if (childMutex) { - bool acquired = false; - std::thread observer([&] { - acquired = childMutex->try_lock(); - if (acquired) { - childMutex->unlock(); - } - }); - observer.join(); - assert(acquired); - } - } -}; - -std::weak_ptr observedHandle; -bool failAllocation = false; -long ownersAtFailure = -1; - -static std::shared_ptr MakeMetadata() { - auto metadata = std::make_shared(); - metadata->columns.push_back({u"owned", SQL_INTEGER, 10, 0, 1}); - return metadata; -} - -static void Populate(ResultMetadataCache& cache) { - const auto snapshot = cache.snapshot(); - cache.publish(snapshot.generation, MakeMetadata()); -} - -static void TestSnapshots() { - ResultMetadataCache cache; - const auto initial = cache.snapshot(); - assert(!initial.metadata); - auto metadata = MakeMetadata(); - std::weak_ptr weak = metadata; - cache.publish(initial.generation, metadata); - auto held = cache.snapshot(); - assert(held.metadata == metadata); - cache.clear(); - assert(!cache.snapshot().metadata); - assert(cache.snapshot().generation != initial.generation); - - Populate(cache); - const auto replacement = cache.snapshot(); - cache.publish(initial.generation, metadata); - assert(cache.snapshot().metadata == replacement.metadata); - assert(held.metadata->columns.at(0).name == u"owned"); - metadata.reset(); - assert(!weak.expired()); - held.metadata.reset(); - assert(weak.expired()); -} - -static void TestFailures() { - ResultMetadataCache cache; - const SQLRETURN results[] = {SQL_SUCCESS, SQL_SUCCESS_WITH_INFO, SQL_NO_DATA, - SQL_ERROR, SQL_INVALID_HANDLE}; - for (SQLRETURN result : results) { - Populate(cache); - const auto before = cache.snapshot(); - { - ResultMetadataFailureGuard guard(cache, result); - } - const auto after = cache.snapshot(); - if (SQL_SUCCEEDED(result) || result == SQL_NO_DATA) { - assert(after.metadata == before.metadata); - assert(after.generation == before.generation); - } else { - assert(!after.metadata); - assert(after.generation != before.generation); - } - } - Populate(cache); - SQLRETURN result = SQL_SUCCESS; - try { - ResultMetadataFailureGuard guard(cache, result); - throw std::runtime_error("conversion failure"); - } catch (const std::runtime_error&) { - assert(!cache.snapshot().metadata); - } -} - -static void TestConcurrentInvalidation() { - ResultMetadataCache cache; - const auto metadata = MakeMetadata(); - std::thread invalidator([&] { - for (int i = 0; i < 1000; ++i) { - cache.clear(); - } - }); - for (int i = 0; i < 1000; ++i) { - const auto snapshot = cache.snapshot(); - cache.publish(snapshot.generation, metadata); - if (snapshot.metadata) { - assert(snapshot.metadata->columns.at(0).name == u"owned"); - } - } - invalidator.join(); - cache.clear(); - assert(!cache.snapshot().metadata); -} - -static void TestChildren() { - std::mutex childMutex; - auto first = std::make_shared(); - auto second = std::make_shared(); - auto unrelated = std::make_shared(); - std::vector> children{first, {}, second}; - Populate(first->resultMetadata); - Populate(second->resultMetadata); - Populate(unrelated->resultMetadata); - const auto held = first->resultMetadata.snapshot(); - ClearChildResultMetadata(childMutex, children); - assert(!first->resultMetadata.snapshot().metadata); - assert(!second->resultMetadata.snapshot().metadata); - assert(unrelated->resultMetadata.snapshot().metadata); - assert(held.metadata->columns.at(0).name == u"owned"); - ClearChildResultMetadata(childMutex, children); -} - -static void TestAllocationFailure() { - std::mutex childMutex; - auto owner = std::make_shared(); - observedHandle = owner; - std::vector> children{owner}; - Populate(owner->resultMetadata); - const auto before = owner->resultMetadata.snapshot(); - failAllocation = true; - try { - ClearChildResultMetadata(childMutex, children); - assert(false); - } catch (const std::bad_alloc&) { - assert(ownersAtFailure == 1); - assert(owner->resultMetadata.snapshot().metadata == before.metadata); - assert(childMutex.try_lock()); - childMutex.unlock(); - } - assert(!failAllocation); - ClearChildResultMetadata(childMutex, children); - assert(!owner->resultMetadata.snapshot().metadata); -} - -static void TestLastOwner() { - std::mutex childMutex; - auto owner = std::make_shared(); - owner->childMutex = &childMutex; - const std::weak_ptr weak = owner; - std::vector> children{owner}; - // Drop the external owner during invalidation, leaving only the helper's snapshot. - auto metadata = std::shared_ptr(new ResultMetadata, [&](auto* value) { - owner.reset(); - delete value; - }); - const auto generation = owner->resultMetadata.snapshot().generation; - owner->resultMetadata.publish(generation, std::move(metadata)); - ClearChildResultMetadata(childMutex, children); - assert(!owner && weak.expired()); -} - -int main(int argc, char** argv) { - if (argc != 2) { - std::cerr << "Expected one native metadata test case\n"; - return 2; - } - const char* name = argv[1]; - if (std::strcmp(name, "snapshots") == 0) { - TestSnapshots(); - } else if (std::strcmp(name, "failures") == 0) { - TestFailures(); - } else if (std::strcmp(name, "concurrent") == 0) { - TestConcurrentInvalidation(); - } else if (std::strcmp(name, "children") == 0) { - TestChildren(); - } else if (std::strcmp(name, "allocation") == 0) { - TestAllocationFailure(); - } else if (std::strcmp(name, "last_owner") == 0) { - TestLastOwner(); - } else { - std::cerr << "Unknown native metadata test case: " << name << '\n'; - return 2; - } - std::cout << name << " passed\n"; - return 0; -} diff --git a/tests/test_040_fetch_native_metadata.py b/tests/test_040_fetch_native_metadata.py deleted file mode 100644 index 102ad6f35..000000000 --- a/tests/test_040_fetch_native_metadata.py +++ /dev/null @@ -1,816 +0,0 @@ -"""Native fetch metadata must preserve public descriptions and result-set state.""" - -import datetime as dt -import os -from pathlib import Path -import subprocess -import sys -import textwrap -from decimal import Decimal -from uuid import UUID - -import pytest - -import mssql_python -from mssql_python import ddbc_bindings - - -@pytest.fixture -def metadata_cursor(conn_str): - with mssql_python.connect(conn_str) as connection: - with connection.cursor() as cursor: - yield cursor - - -def _query(columns, count=20): - values = ",".join(f"({i})" for i in range(1, 21)) - return ( - f"SELECT {','.join(columns)} FROM (VALUES {values}) AS source(id) " - f"WHERE id<={count} ORDER BY id" - ) - - -def _assert_rows(rows, expected): - assert [tuple(row) for row in rows] == expected - assert [[type(value) for value in row] for row in rows] == [ - [type(value) for value in row] for row in expected - ] - - -def _describe(cursor): - result = [] - assert ddbc_bindings.DDBCSQLDescribeCol(cursor.hstmt, result) == 0 - for column in result: - assert set(column) == {"ColumnName", "DataType", "ColumnSize", "DecimalDigits", "Nullable"} - assert type(column["ColumnName"]) is str - for key in ("DataType", "ColumnSize", "DecimalDigits", "Nullable"): - assert type(column[key]) is int - return result - - -@pytest.mark.parametrize("width", [3, 24]) -@pytest.mark.parametrize("size", [None, 1, 10, 1000, "varied"]) -@pytest.mark.parametrize("count", [0, 20]) -def test_fetchmany_shape_sizes_and_eof(metadata_cursor, width, size, count): - cursor = metadata_cursor - expressions = ["id", "CONVERT(NVARCHAR(30),N'row')", "CONVERT(FLOAT,id)*0.25"] - columns = [f"{expressions[i % 3]} AS c{i}" for i in range(width)] - cursor.execute(_query(columns, count)) - assert cursor.arraysize == 1 - description = cursor.description - assert all(len(column) == 7 for column in description) - assert [column[0] for column in description] == [f"c{i}" for i in range(width)] - metadata = _describe(cursor) - output = [] - iteration = 0 - while True: - fetch_size = (1, 10, 3, 1000)[iteration % 4] if size == "varied" else size - batch = cursor.fetchmany() if fetch_size is None else cursor.fetchmany(fetch_size) - assert cursor.description == description - if not batch: - break - output.extend(batch) - iteration += 1 - _assert_rows(output, [(i, "row", i * 0.25) * (width // 3) for i in range(1, count + 1)]) - assert _describe(cursor) == metadata - assert cursor.fetchmany(1) == [] - assert cursor.fetchone() is None - assert cursor.fetchall() == [] - - -@pytest.mark.parametrize("method", ["fetchmany", "fetchall", "arrow_batch"]) -def test_public_metadata_names_and_fields(metadata_cursor, method): - cursor = metadata_cursor - names = [ - "duplicate", - "duplicate", - "\u03a9\u540d", - "emoji_\U0001f600", - "bracket]name", - "x" * 128, - ] - columns = [f"CONVERT(INT,id) AS [{name.replace(']', ']]')}]" for name in names] - cursor.execute(_query(columns, 1)) - metadata = _describe(cursor) - assert metadata == [ - {"ColumnName": name, "DataType": 4, "ColumnSize": 10, "DecimalDigits": 0, "Nullable": 1} - for name in names - ] - assert [column[0] for column in cursor.description] == names - if method == "arrow_batch": - pytest.importorskip("pyarrow") - batch = cursor.arrow_batch(1) - assert batch.schema.names == names - assert [column.to_pylist() for column in batch.columns] == [[1]] * len(names) - else: - rows = cursor.fetchmany(1) if method == "fetchmany" else cursor.fetchall() - _assert_rows(rows, [(1,) * len(names)]) - assert _describe(cursor) == metadata - - -_TYPES = [ - ("INT", "7", 7), - ("SMALLINT", "-7", -7), - ("BIGINT", "2147483649", 2147483649), - ("TINYINT", "255", 255), - ("BIT", "1", True), - ("REAL", "1.5", 1.5), - ("FLOAT", "2.25", 2.25), - ("DECIMAL(20,4)", "123.4500", Decimal("123.4500")), - ("NUMERIC(28,8)", "-0.125", Decimal("-0.125")), - ("MONEY", "4.25", Decimal("4.25")), - ("DATE", "'2001-02-03'", dt.date(2001, 2, 3)), - ("TIME(7)", "'12:34:56.1234567'", dt.time(12, 34, 56, 123456)), - ("DATETIME2(7)", "'2001-02-03T12:34:56.1234567'", dt.datetime(2001, 2, 3, 12, 34, 56, 123456)), - ("DATETIME", "'2001-02-03T12:34:56'", dt.datetime(2001, 2, 3, 12, 34, 56)), - ( - "DATETIMEOFFSET(7)", - "'2001-02-03T12:34:56.1234567+05:30'", - dt.datetime(2001, 2, 3, 12, 34, 56, 123456, dt.timezone(dt.timedelta(minutes=330))), - ), - ( - "UNIQUEIDENTIFIER", - "'12345678-1234-5678-1234-567812345678'", - UUID("12345678-1234-5678-1234-567812345678"), - ), - ("VARCHAR(20)", "'ascii'", "ascii"), - ("CHAR(8)", "'ascii'", "ascii "), - ("NVARCHAR(30)", "N'\u03a9\U0001f600'", "\u03a9\U0001f600"), - ("NCHAR(5)", "N'\u03a9'", "\u03a9 "), - ("VARBINARY(10)", "0x00010200", b"\x00\x01\x02\x00"), - ("BINARY(4)", "0x00010203", b"\x00\x01\x02\x03"), - ("VARCHAR(1)", "''", ""), - ("NVARCHAR(1)", "N''", ""), -] - - -@pytest.mark.parametrize("size", [1, 10, 1000]) -def test_fetchmany_typed_nulls_and_values(metadata_cursor, size): - columns = [ - f"CASE WHEN id%3=0 THEN CAST(NULL AS {sqltype}) " - f"ELSE CAST({literal} AS {sqltype}) END AS c{i}" - for i, (sqltype, literal, _) in enumerate(_TYPES) - ] - cursor = metadata_cursor - cursor.execute(_query(columns)) - description = cursor.description - metadata = _describe(cursor) - assert len(metadata) == 24 - output = [] - while batch := cursor.fetchmany(size): - output.extend(batch) - values = tuple(value for _, _, value in _TYPES) - _assert_rows(output, [(None,) * 24 if i % 3 == 0 else values for i in range(1, 21)]) - assert cursor.description == description - assert _describe(cursor) == metadata - - -def test_reexecute_and_nextset_change_shape(metadata_cursor): - cursor = metadata_cursor - for _ in range(3): - cursor.execute("SELECT 1 AS first_name; SELECT N'new' AS second_name, 2 AS extra") - _assert_rows(cursor.fetchmany(1), [(1,)]) - assert cursor.nextset() - assert [col[0] for col in cursor.description] == ["second_name", "extra"] - _assert_rows(cursor.fetchmany(10), [("new", 2)]) - assert not cursor.nextset() - cursor.execute("SELECT CAST(3.5 AS DECIMAL(6,2)) AS replacement") - assert _describe(cursor)[0]["ColumnName"] == "replacement" - _assert_rows(cursor.fetchmany(), [(Decimal("3.50"),)]) - - -def test_converter_changes_on_execute_and_live_decoding(metadata_cursor): - cursor = metadata_cursor - connection = cursor.connection - cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 4)) - _assert_rows(cursor.fetchmany(1), [(1, "ascii")]) - calls = [] - - def convert(value): - calls.append(value) - return value + 100 - - connection.add_output_converter(mssql_python.SQL_INTEGER, convert) - connection.setdecoding(mssql_python.SQL_CHAR, "utf-8", mssql_python.SQL_CHAR) - cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 4)) - _assert_rows(cursor.fetchmany(1), [(101, "ascii")]) - assert calls == [1] - connection.setdecoding(mssql_python.SQL_CHAR, "latin1", mssql_python.SQL_CHAR) - _assert_rows(cursor.fetchmany(1), [(102, "ascii")]) - assert calls == [1, 2] - connection.remove_output_converter(mssql_python.SQL_INTEGER) - cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 1)) - _assert_rows(cursor.fetchmany(1), [(1, "ascii")]) - connection.setdecoding(mssql_python.SQL_CHAR) - cursor.execute(_query(["id", "CAST('ascii' AS VARCHAR(12)) AS txt"], 1)) - _assert_rows(cursor.fetchmany(10), [(1, "ascii")]) - - -@pytest.mark.parametrize("size", [1, 10]) -def test_fetchmany_lob_and_xml_typed_nulls(metadata_cursor, size): - cursor = metadata_cursor - columns = [ - "CASE WHEN id%2=0 THEN CAST(NULL AS NVARCHAR(MAX)) ELSE " - "REPLICATE(CAST(N'x' AS NVARCHAR(MAX)),9001) END AS txt", - "CASE WHEN id%2=0 THEN CAST(NULL AS VARBINARY(MAX)) ELSE " - "CAST(REPLICATE(CAST('a' AS VARCHAR(MAX)),10003) AS VARBINARY(MAX)) END AS bin", - "CASE WHEN id%2=0 THEN CAST(NULL AS XML) ELSE CAST('value' AS XML) END AS xml", - ] - cursor.execute(_query(columns, 4)) - metadata = _describe(cursor) - rows = [] - while batch := cursor.fetchmany(size): - rows.extend(batch) - values = ("x" * 9001, b"a" * 10003, "value") - _assert_rows(rows, [values, (None, None, None), values, (None, None, None)]) - assert _describe(cursor) == metadata - - -def _isolated(script, tmp_path): - environment = dict(os.environ) - root = str(Path(mssql_python.__file__).resolve().parent.parent) - environment["PYTHONPATH"] = os.pathsep.join([root, environment.get("PYTHONPATH", "")]) - result = subprocess.run( - [sys.executable, "-c", textwrap.dedent(script)], - cwd=tmp_path, - env=environment, - capture_output=True, - text=True, - timeout=45, - ) - assert result.returncode == 0, result.stdout + result.stderr - - -def test_interleaving_movement_and_variant_freshness(tmp_path): - _isolated( - """ - import os - import gc - import mssql_python as db - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - query = "SELECT id FROM (VALUES(1),(2),(3),(4),(5),(6),(7)) s(id) ORDER BY id" - for _ in range(4): - cursor.execute(query) - assert cursor.fetchmany(1)[0][0] == 1 - gc.collect() - assert cursor.fetchone()[0] == 2 - cursor.scroll(1) - assert cursor.fetchmany(1)[0][0] == 4 - cursor.skip(1) - assert [tuple(row) for row in cursor.fetchall()] == [(6,), (7,)] - cursor.execute("CREATE TABLE #metadata_variant(id INT, v SQL_VARIANT)") - cursor.execute( - "INSERT INTO #metadata_variant VALUES " - "(1,CAST('abc' AS VARCHAR(3)))," - "(2,CAST('abcdefgh' AS VARCHAR(8)))," - "(3,CAST(REPLICATE('x',30) AS VARCHAR(30)))" - ) - for method in ("fetchmany", "fetchall"): - cursor.execute("SELECT v FROM #metadata_variant ORDER BY id") - assert cursor.fetchone()[0] == "abc" - if method == "fetchmany": - assert cursor.fetchmany(1)[0][0] == "abcdefgh" - assert cursor.fetchmany(1)[0][0] == "x"*30 - else: - assert [r[0] for r in cursor.fetchall()] == ["abcdefgh", "x"*30] - """, - tmp_path, - ) - - -def test_closure_and_error_recovery(tmp_path): - _isolated( - """ - import os - import mssql_python as db - from mssql_python import Cursor, InterfaceError, ProgrammingError, DatabaseError - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute("SELECT 1 AS c") - assert cursor.fetchmany(0) == [] - assert cursor.fetchmany(-1) == [] - assert cursor.fetchmany(1)[0][0] == 1 - try: - cursor.execute("SELECT invalid_column FROM (VALUES(1)) t(c)") - except DatabaseError: - pass - else: - raise AssertionError("invalid query did not raise") - cursor.execute("SELECT 2 AS changed") - assert cursor.fetchmany(1)[0][0] == 2 - try: - cursor.fetchmany(1) - except ProgrammingError: - pass - else: - raise AssertionError("closed cursor did not raise") - connection = db.connect(os.environ["DB_CONNECTION_STRING"]) - cursor = Cursor(connection) - cursor.execute("SELECT 1") - connection.close() - try: - cursor.fetchmany(1) - except (InterfaceError, ProgrammingError): - pass - else: - raise AssertionError("closed connection did not raise") - cursor.close() - """, - tmp_path, - ) - - -def test_malformed_column_name_fails_before_fetch(tmp_path): - _isolated( - """ - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute( - "DECLARE @s NVARCHAR(200) = N'SELECT 1 AS [' + " - "CAST(0x00D8 AS NVARCHAR(1)) + N']'; EXEC(@s)" - ) - for operation in ( - lambda: native.DDBCSQLDescribeCol(cursor.hstmt, []), - lambda: cursor.fetchmany(1), - ): - try: - operation() - except UnicodeDecodeError: - pass - else: - raise AssertionError("malformed UTF-16 column name did not raise") - assert native.DDBCSQLFetch(cursor.hstmt) == 0 - assert native.DDBCSQLFetch(cursor.hstmt) == 100 - """, - tmp_path, - ) - - -@pytest.mark.skipif( - not hasattr(ddbc_bindings, "profiling"), reason="requires native profiling instrumentation" -) -def test_fetchmany_avoids_python_description_roundtrip(tmp_path): - _isolated( - """ - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute("SELECT id FROM (VALUES(1),(2)) s(id) ORDER BY id") - native.profiling.reset() - native.profiling.enable() - try: - metadata = [] - native.DDBCSQLDescribeCol(cursor.hstmt, metadata) - finally: - native.profiling.disable() - assert native.profiling.get_stats()["ddbc::SQLDescribeCol_wrap"]["calls"] == 1 - assert len(metadata) == 1 - native.profiling.reset() - native.profiling.enable() - try: - assert cursor.fetchmany(1)[0][0] == 1 - assert cursor.fetchmany(1)[0][0] == 2 - assert cursor.fetchmany(1) == [] - finally: - native.profiling.disable() - stats = native.profiling.get_stats() - assert stats["ddbc::FetchMany_wrap"]["calls"] == 3 - assert stats.get("ddbc::SQLDescribeCol_wrap", {}).get("calls", 0) == 0 - """, - tmp_path, - ) - - -@pytest.mark.parametrize("method", ["one", "many"]) -def test_result_metadata_prepared_reexecution(metadata_cursor, method): - cursor = metadata_cursor - statement = cursor.hstmt - query = "SELECT CAST(? AS INT) AS n, CAST(? AS NVARCHAR(30)) AS text_value" - for value in range(4): - cursor.execute(query, (value, f"value-{value}")) - assert cursor.hstmt is statement - assert cursor.is_stmt_prepared[0] - profiling = hasattr(ddbc_bindings, "profiling") - if profiling: - ddbc_bindings.profiling.reset() - ddbc_bindings.profiling.enable() - try: - rows = [cursor.fetchone()] if method == "one" else cursor.fetchmany() - finally: - if profiling: - ddbc_bindings.profiling.disable() - _assert_rows(rows, [(value, f"value-{value}")]) - if profiling: - assert ( - ddbc_bindings.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] - == 2 - ) - assert cursor.fetchone() is None - cursor.execute("SELECT CAST(? AS DECIMAL(8,2)) AS amount", (Decimal("3.25"),)) - _assert_rows(cursor.fetchall(), [(Decimal("3.25"),)]) - - -def test_result_metadata_catalog_replacement(metadata_cursor): - cursor = metadata_cursor - for _ in range(2): - cursor.execute("SELECT 42 AS previous_column") - _assert_rows(cursor.fetchmany(), [(42,)]) - cursor.getTypeInfo(mssql_python.SQL_INTEGER) - description = cursor.description - assert len(description) > 1 - row = cursor.fetchone() - assert row is not None and len(row) == len(description) - assert row[1] == mssql_python.SQL_INTEGER - cursor.fetchall() - cursor.execute("SELECT N'replaced' AS new_column, 5 AS extra") - _assert_rows(cursor.fetchmany(), [("replaced", 5)]) - - -def test_result_metadata_native_reset_and_replacement(tmp_path): - _isolated( - """ - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - stmt = connection._conn.alloc_statement_handle() - try: - for _ in range(3): - assert native.DDBCSQLExecDirect(stmt, "SELECT 1 AS a") in (0, 1) - rows = [] - assert native.DDBCSQLFetchMany(stmt, rows, 1) in (0, 1) - assert rows == [[1]] - assert native.DDBCSQLResetStmt(stmt) in (0, 1) - assert native.DDBCSQLExecDirect( - stmt, "SELECT CAST(2 AS BIGINT) AS b, N'new' AS c" - ) in (0, 1) - row = [] - assert native.DDBCSQLFetchOne(stmt, row) in (0, 1) - assert row == [2, "new"] - stmt._close_cursor() - finally: - stmt.free() - """, - tmp_path, - ) - - -@pytest.mark.parametrize("operation", ["commit", "rollback", "autocommit"]) -def test_result_metadata_transaction_recovery(metadata_cursor, operation): - cursor = metadata_cursor - connection = cursor.connection - cursor.execute(_query(["id"], 3)) - _assert_rows(cursor.fetchmany(), [(1,)]) - if operation == "autocommit": - connection.autocommit = True - else: - getattr(connection, operation)() - cursor.execute("SELECT CAST(5.75 AS DECIMAL(8,2)) AS changed, N'text' AS extra") - _assert_rows(cursor.fetchall(), [(Decimal("5.75"), "text")]) - - -@pytest.mark.parametrize("operation", ["commit", "rollback", "autocommit"]) -def test_result_metadata_transaction_preserved_cursor(metadata_cursor, operation): - cursor = metadata_cursor - connection = cursor.connection - info = ( - mssql_python.SQL_CURSOR_ROLLBACK_BEHAVIOR - if operation == "rollback" - else mssql_python.SQL_CURSOR_COMMIT_BEHAVIOR - ) - if connection.getinfo(info) != 2: # SQL_CB_PRESERVE - pytest.skip("Driver does not preserve cursors; cache/helper coverage is in tests/native") - cursor.execute(_query(["id"], 3)) - _assert_rows([cursor.fetchone()], [(1,)]) - if operation == "autocommit": - connection.autocommit = True - else: - getattr(connection, operation)() - profiling = hasattr(ddbc_bindings, "profiling") - if profiling: - ddbc_bindings.profiling.reset() - ddbc_bindings.profiling.enable() - try: - _assert_rows([cursor.fetchone()], [(2,)]) - _assert_rows(cursor.fetchmany(), [(3,)]) - finally: - if profiling: - ddbc_bindings.profiling.disable() - if profiling: - assert ( - ddbc_bindings.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 1 - ) - - -def test_result_metadata_arrow_interleave(tmp_path): - pytest.importorskip("pyarrow") - _isolated( - """ - import gc - import os - import mssql_python as db - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - query = ("SELECT id, CAST(id AS BIGINT) AS big FROM " - "(VALUES(1),(2),(3),(4),(5),(6)) s(id) ORDER BY id") - for _ in range(3): - cursor.execute(query) - assert tuple(cursor.fetchone()) == (1, 1) - assert [tuple(r) for r in cursor.fetchmany(2)] == [(2, 2), (3, 3)] - batch = cursor.arrow_batch(1) - assert [c.to_pylist() for c in batch.columns] == [[4], [4]] - gc.collect() - assert tuple(cursor.fetchone()) == (5, 5) - assert [tuple(r) for r in cursor.fetchall()] == [(6, 6)] - assert cursor.fetchmany() == [] - """, - tmp_path, - ) - - -@pytest.mark.parametrize("method", ["one", "many", "all"]) -def test_result_metadata_variant_type_and_size_changes(metadata_cursor, method): - cursor = metadata_cursor - cursor.execute( - "CREATE TABLE #metadata_mixed_variant (id INT, v SQL_VARIANT, txt NVARCHAR(MAX))" - ) - cursor.execute( - "INSERT INTO #metadata_mixed_variant VALUES " - "(1,CAST(NULL AS SQL_VARIANT),N'first')," - "(2,CAST(CAST('abc' AS VARCHAR(3)) AS SQL_VARIANT),NULL)," - "(3,CAST(CAST('abcdefgh' AS VARCHAR(8)) AS SQL_VARIANT),N'third')," - "(4,CAST(CAST(17 AS INT) AS SQL_VARIANT),NULL)," - "(5,CAST(CAST(3.25 AS DECIMAL(8,2)) AS SQL_VARIANT),N'fifth')," - "(6,CAST(NULL AS SQL_VARIANT),NULL)," - "(7,CAST(CAST(0x010200 AS VARBINARY(3)) AS SQL_VARIANT),N'last')" - ) - cursor.execute("SELECT v, txt FROM #metadata_mixed_variant ORDER BY id") - if method == "one": - rows = list(cursor) - elif method == "many": - rows = [] - while batch := cursor.fetchmany(): - rows.extend(batch) - else: - rows = cursor.fetchall() - _assert_rows( - rows, - [ - (None, "first"), - ("abc", None), - ("abcdefgh", "third"), - (17, None), - (Decimal("3.25"), "fifth"), - (None, None), - (b"\x01\x02\x00", "last"), - ], - ) - - -def test_result_metadata_one_then_malformed_name_many_does_not_advance(tmp_path): - _isolated( - """ - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute( - "DECLARE @s NVARCHAR(300) = N'SELECT id AS [' + " - "CAST(0x00D8 AS NVARCHAR(1)) + " - "N'] FROM (VALUES(1),(2),(3)) s(id) ORDER BY id'; EXEC(@s)" - ) - # The low-level row path never decoded a supported column's name. - row = [] - assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) - assert row == [1] - for _ in range(2): - try: - native.DDBCSQLFetchMany(cursor.hstmt, [], 1) - except UnicodeDecodeError: - pass - else: - raise AssertionError("many accepted the malformed column name") - row = [] - profiling = hasattr(native, "profiling") - if profiling: - native.profiling.reset() - native.profiling.enable() - try: - assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) - finally: - if profiling: - native.profiling.disable() - assert row == [2] - if profiling: - assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 1 - cursor.execute("SELECT 4 AS valid_name, N'recovered' AS text_value") - assert tuple(cursor.fetchone()) == (4, "recovered") - assert cursor.fetchmany() == [] - """, - tmp_path, - ) - - -@pytest.mark.skipif( - not hasattr(ddbc_bindings, "profiling"), reason="requires actual ODBC call instrumentation" -) -@pytest.mark.parametrize("method", ["one", "many"]) -def test_result_metadata_actual_description_counts(tmp_path, method): - _isolated( - f""" - import os - from decimal import Decimal - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - columns = ",".join(f"id AS c{{i}}" for i in range(24)) - query = ("WITH n AS (SELECT TOP(10000) ROW_NUMBER() OVER " - "(ORDER BY a.object_id,b.object_id) AS id " - "FROM sys.all_objects a CROSS JOIN sys.all_objects b) " - f"SELECT {{columns}} FROM n ORDER BY id") - cursor.execute(query) - native.profiling.reset() - native.profiling.enable() - try: - for value in range(1, 10001): - row = cursor.fetchone() if {method!r} == "one" else cursor.fetchmany()[0] - assert tuple(row) == (value,) * 24 - assert cursor.fetchone() is None - assert cursor.fetchmany() == [] - finally: - native.profiling.disable() - stats = native.profiling.get_stats() - assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == 24, stats - assert stats.get("ddbc::SQLDescribeCol_wrap", {{}}).get("calls", 0) == 0 - native.profiling.reset() - native.profiling.enable() - try: - metadata = [] - assert native.DDBCSQLDescribeCol(cursor.hstmt, metadata) in (0, 1) - finally: - native.profiling.disable() - assert len(metadata) == 24 - assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 24 - cursor.execute( - "SELECT CAST(3 AS INT) AS changed, CAST(N'x' AS NVARCHAR(1)) AS text_value; " - "SELECT CAST(7.25 AS DECIMAL(8,2)) AS amount, " - "CAST(N'next long value' AS NVARCHAR(40)) AS name" - ) - assert tuple(cursor.fetchone()) == (3, "x") - assert cursor.nextset() - native.profiling.reset() - native.profiling.enable() - try: - assert tuple(cursor.fetchmany()[0]) == (Decimal("7.25"), "next long value") - assert cursor.fetchmany() == [] - finally: - native.profiling.disable() - assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 2 - """, - tmp_path, - ) - - -@pytest.mark.skipif( - not hasattr(ddbc_bindings, "profiling"), reason="requires actual ODBC call instrumentation" -) -@pytest.mark.parametrize("method", ["one", "many", "all"]) -def test_result_metadata_variant_descriptions_remain_per_row(tmp_path, method): - _isolated( - f""" - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute( - "SELECT id, v FROM (VALUES " - "(1,CAST(NULL AS SQL_VARIANT))," - "(2,CAST('abc' AS SQL_VARIANT))," - "(3,CAST(17 AS SQL_VARIANT))) s(id,v) ORDER BY id" - ) - native.profiling.reset() - native.profiling.enable() - try: - if {method!r} == "one": - rows = list(cursor) - elif {method!r} == "all": - rows = cursor.fetchall() - else: - rows = [] - while batch := cursor.fetchmany(): - rows.extend(batch) - finally: - native.profiling.disable() - assert [tuple(row) for row in rows] == [(1,None),(2,"abc"),(3,17)] - stats = native.profiling.get_stats() - expected = 4 if {method!r} == "one" else 5 - assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == expected, stats - assert stats["ddbc::sql_variant::null_probe"]["calls"] == 3, stats - assert stats["ddbc::sql_variant::subtype"]["calls"] == 2, stats - """, - tmp_path, - ) - - -@pytest.mark.parametrize("method", ["one", "many", "all"]) -def test_result_metadata_all_null_rows(tmp_path, method): - _isolated( - f""" - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute("SELECT CAST(NULL AS INT) AS scalar_null") - scalar_null = [] - scalar_status = native.DDBCSQLFetchOne(cursor.hstmt, scalar_null) - diagnostics = native.DDBCSQLGetAllDiagRecords(cursor.hstmt) - assert scalar_status == 0, (scalar_status, diagnostics) - assert scalar_null == [None] - assert diagnostics == [] - values = ",".join(f"({{i}})" for i in range(1, 16)) - cursor.execute( - "SELECT CASE WHEN id%7=0 THEN NULL ELSE id END AS c0," - "CASE WHEN id%7=0 THEN CAST(NULL AS SQL_VARIANT) " - "WHEN id%3=0 THEN CAST(id AS SQL_VARIANT) " - "WHEN id%3=1 THEN CAST(N'row-'+CONVERT(NVARCHAR(12),id) AS SQL_VARIANT) " - "ELSE CAST(CONVERT(FLOAT,id)*0.25 AS SQL_VARIANT) END AS c1 " - f"FROM (VALUES{{values}}) s(id) ORDER BY id" - ) - profiling = hasattr(native, "profiling") - if profiling: - native.profiling.reset() - native.profiling.enable() - try: - if {method!r} == "one": - rows = list(cursor) - elif {method!r} == "all": - rows = cursor.fetchall() - else: - rows = [] - while batch := cursor.fetchmany(): - rows.extend(batch) - finally: - if profiling: - native.profiling.disable() - expected = [ - (None,None) if i%7==0 else (i,(i,f"row-{{i}}",i*0.25)[i%3]) - for i in range(1,16) - ] - assert [tuple(row) for row in rows] == expected - assert [[type(value) for value in row] for row in rows] == [ - [type(value) for value in row] for row in expected - ] - if profiling: - stats = native.profiling.get_stats() - expected_describes = 16 if {method!r} == "one" else 17 - assert stats["ddbc::SQLDescribeCol::driver_call"]["calls"] == expected_describes, stats - assert stats["ddbc::sql_variant::null_probe"]["calls"] == 15 - assert stats["ddbc::sql_variant::subtype"]["calls"] == 13 - """, - tmp_path, - ) - - -def test_result_metadata_odbc_error_invalidates(tmp_path): - _isolated( - """ - import os - import mssql_python as db - from mssql_python import ddbc_bindings as native - with db.connect(os.environ["DB_CONNECTION_STRING"]) as connection: - with connection.cursor() as cursor: - cursor.execute("SELECT id FROM (VALUES(1),(2),(3)) s(id) ORDER BY id") - row = [] - assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) - assert row == [1] - assert native.DDBCSQLGetData( - cursor.hstmt, 2, [], "utf-16le", "utf-16le", db.SQL_WCHAR - ) == -1 - diagnostics = native.DDBCSQLGetAllDiagRecords(cursor.hstmt) - assert any("07009" in state for state, _ in diagnostics), diagnostics - profiling = hasattr(native, "profiling") - if profiling: - native.profiling.reset() - native.profiling.enable() - try: - row = [] - assert native.DDBCSQLFetchOne(cursor.hstmt, row) in (0, 1) - assert row == [2] - finally: - if profiling: - native.profiling.disable() - if profiling: - assert native.profiling.get_stats()["ddbc::SQLDescribeCol::driver_call"]["calls"] == 1 - """, - tmp_path, - ) From 9eb586a63f251c98d1053673170b54b1e5c0015d Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 13:23:45 +0530 Subject: [PATCH 08/15] FIX: Address fetch buffer reuse review feedback Replace banned formatted C output without entering Python during cleanup. Add isolated profiler regression checks for binding reuse and size, encoding, and result-set transitions. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 9 +-- tests/test_025_profiler.py | 92 +++++++++++++++++++++++++++ 2 files changed, 97 insertions(+), 4 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 8d1aa43e9..42a5cdf56 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -1668,8 +1668,8 @@ SqlHandle::~SqlHandle() { // A failed free leaves a live handle. Detach if possible before // the last plan owner applies its native-only emergency policy. SQLRETURN detached = detachFetchBindings(); - std::fprintf(stderr, "mssql-python: native handle cleanup failed (%d), " - "fetch buffer detach returned %d\n", ret, detached); + std::cerr << "mssql-python: native handle cleanup failed (" << ret + << "), fetch buffer detach returned " << detached << '\n'; } } } catch (...) { @@ -2204,8 +2204,9 @@ static void AppendFetchBindingDiagnostics(py::handle messages, } catch (const std::exception& error) { if (!preserveFailure) throw; - std::fprintf(stderr, "mssql-python: failed to append fetch binding diagnostics: %s\n", - error.what()); + std::fputs("mssql-python: failed to append fetch binding diagnostics: ", stderr); + std::fputs(error.what(), stderr); + std::fputc('\n', stderr); } } diff --git a/tests/test_025_profiler.py b/tests/test_025_profiler.py index 1db0ebce1..18b00b2d7 100644 --- a/tests/test_025_profiler.py +++ b/tests/test_025_profiler.py @@ -394,6 +394,98 @@ def test_cpp_profiling_captures_query(): assert sample["calls"] >= 1 +@_needs_cpp +@_needs_db +@pytest.mark.parametrize("transition", ["size", "encoding", "result"]) +def test_fetchmany_reuses_bindings_until_transition(transition): + """Release + ENABLE_PROFILING=ON: count actual calls, isolated from other cursors.""" + script = textwrap.dedent(""" + import os + import sys + + sys.path.insert(0, sys.argv[2]) + import mssql_python as db + from mssql_python import ddbc_bindings as native + + assert hasattr(native, "profiling") + transition = sys.argv[1] + query = ( + "SELECT n, CAST(n AS VARCHAR(10)) AS txt " + "FROM (VALUES (1), (2), (3), (4), (5), (6)) AS v(n) ORDER BY n" + ) + + def counts(plans, binds, unbinds): + stats = native.profiling.get_stats() + for name, expected in ( + ("plan_allocation", plans), ("SQLBindCol", binds), ("SQL_UNBIND", unbinds) + ): + key = "ddbc::fetch_bindings::" + name + if expected: + assert key in stats, (key, stats) + assert stats[key]["calls"] == expected, (key, stats) + else: + assert key not in stats, (key, stats) + + def fetch(cursor, size, first): + expected = [(n, str(n)) for n in range(first, min(first + size, 7))] + assert [tuple(row) for row in cursor.fetchmany(size)] == expected + + try: + connection = db.connect(os.environ["DB_CONNECTION_STRING"], timeout=5) + except db.Error: + raise RuntimeError("SQL connection failed") from None + with connection, connection.cursor() as cursor: + connection.setdecoding(db.SQL_CHAR, encoding="utf-16le", ctype=db.SQL_WCHAR) + cursor.execute(query) + native.profiling.reset() + native.profiling.enable() + try: + fetch(cursor, 1, 1) + counts(1, 2, 0) + fetch(cursor, 1, 2) + counts(1, 2, 0) + size, first = 1, 3 + if transition == "size": + size = 2 + elif transition == "encoding": + connection.setdecoding( + db.SQL_CHAR, encoding="utf-16-le", ctype=db.SQL_WCHAR + ) + else: + cursor.execute(query) + first = 1 + fetch(cursor, size, first) + counts(2, 4, 1) + first += size + fetch(cursor, size, first) + counts(2, 4, 1) + first += size + while first <= 6: + fetch(cursor, size, first) + counts(2, 4, 1) + first += size + assert cursor.fetchmany(size) == [] + counts(2, 4, 2) + finally: + native.profiling.disable() + native.profiling.reset() + """) + result = subprocess.run( + [ + sys.executable, + "-E", + "-c", + script, + transition, + os.path.dirname(os.path.dirname(perf_timer.__file__)), + ], + capture_output=True, + text=True, + timeout=45, + ) + assert result.returncode == 0, result.stdout + result.stderr + + @_needs_cpp @_needs_db def test_cpp_timeline_captures_events(): From be219eb370d860dc003f6f72b59eb1760f2a8f95 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 13:40:43 +0530 Subject: [PATCH 09/15] FIX: Use supported codecs in fetch binding reuse regression Keep SQL_CHAR fixed while changing ASCII to Latin-1 so the encoding-only cache miss reaches the native fetch path. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- tests/test_025_profiler.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/test_025_profiler.py b/tests/test_025_profiler.py index 18b00b2d7..5ad1b7a58 100644 --- a/tests/test_025_profiler.py +++ b/tests/test_025_profiler.py @@ -435,7 +435,7 @@ def fetch(cursor, size, first): except db.Error: raise RuntimeError("SQL connection failed") from None with connection, connection.cursor() as cursor: - connection.setdecoding(db.SQL_CHAR, encoding="utf-16le", ctype=db.SQL_WCHAR) + connection.setdecoding(db.SQL_CHAR, encoding="ascii", ctype=db.SQL_CHAR) cursor.execute(query) native.profiling.reset() native.profiling.enable() @@ -449,7 +449,7 @@ def fetch(cursor, size, first): size = 2 elif transition == "encoding": connection.setdecoding( - db.SQL_CHAR, encoding="utf-16-le", ctype=db.SQL_WCHAR + db.SQL_CHAR, encoding="latin-1", ctype=db.SQL_CHAR ) else: cursor.execute(query) From 1c05d11e1a0f0e519874513bff74c011d5d287d0 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 14:22:35 +0530 Subject: [PATCH 10/15] FIX: Cover native fetch cleanup failure ownership Exercise real statement-owned plans with injected unbind and rows-fetched-pointer cleanup failures, runtime lifetime assertions, and successful release recovery. Keep the native target opt-in and run it in the existing Ubuntu PR jobs. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .github/prompts/run-tests.prompt.md | 17 ++ eng/pipelines/pr-validation-pipeline.yml | 3 + mssql_python/pybind/CMakeLists.txt | 32 +++ tests/native/test_fetch_binding_cleanup.cpp | 256 ++++++++++++++++++++ 4 files changed, 308 insertions(+) create mode 100644 tests/native/test_fetch_binding_cleanup.cpp diff --git a/.github/prompts/run-tests.prompt.md b/.github/prompts/run-tests.prompt.md index da8bcfa88..ce0bef757 100644 --- a/.github/prompts/run-tests.prompt.md +++ b/.github/prompts/run-tests.prompt.md @@ -116,6 +116,23 @@ python main.py ## TASK +### Native fetch cleanup ownership regressions (no database) + +After configuring the native extension with the normal platform/compiler settings, +enable `-DBUILD_FETCH_BINDING_TESTS=ON` in the same CMake build directory. Build +the `fetch_binding_cleanup_test` target in Release, then run +`ctest --test-dir -C Release --output-on-failure -R "^fetch_binding_cleanup_"`. + +These opt-in tests compile the actual native sources with the extension's settings +and embed the selected Python interpreter. They require Python development headers +and its embedding library, but no SQL Server, connection string, or profiling build. +Only the existing ODBC function pointers are replaced, in the test executable. +Runtime checks remain enabled under `NDEBUG` and verify retained buffer ownership, +non-reuse, original-error preservation, and destruction after successful cleanup +following injected unbind and rows-fetched-pointer reset failures. Ubuntu PR CI runs +them before uninstalling the ODBC development headers. Normal wheel builds leave +the option OFF. + Help the developer run tests to validate their changes. Follow this process based on what they need. --- diff --git a/eng/pipelines/pr-validation-pipeline.yml b/eng/pipelines/pr-validation-pipeline.yml index 99507c8e6..75e254c79 100644 --- a/eng/pipelines/pr-validation-pipeline.yml +++ b/eng/pipelines/pr-validation-pipeline.yml @@ -885,6 +885,9 @@ jobs: source /opt/venv/bin/activate cd /workspace python -m eng.profiler_benchmarks.controller --check-build on + cmake -S mssql_python/pybind -B mssql_python/pybind/build -DBUILD_FETCH_BINDING_TESTS=ON + cmake --build mssql_python/pybind/build --config Release --target fetch_binding_cleanup_test + ctest --test-dir mssql_python/pybind/build -C Release --output-on-failure -R '^fetch_binding_cleanup_' " fi displayName: 'Build pybind bindings (.so) in $(distroName) container' diff --git a/mssql_python/pybind/CMakeLists.txt b/mssql_python/pybind/CMakeLists.txt index 77d599bd5..e062683e1 100644 --- a/mssql_python/pybind/CMakeLists.txt +++ b/mssql_python/pybind/CMakeLists.txt @@ -393,3 +393,35 @@ if(APPLE) target_compile_definitions(ddbc_bindings PRIVATE MACOS_STRING_FIX) target_compile_options(ddbc_bindings PRIVATE -DAPPLE_SILICON) endif() + +option(BUILD_FETCH_BINDING_TESTS "Build native fetch cleanup ownership tests (no database)" OFF) +if(BUILD_FETCH_BINDING_TESTS) + enable_testing() + execute_process( + COMMAND python -c "import sys; print(sys.executable)" + OUTPUT_VARIABLE Python3_EXECUTABLE + OUTPUT_STRIP_TRAILING_WHITESPACE + RESULT_VARIABLE test_python_status + ) + if(NOT test_python_status EQUAL 0) + message(FATAL_ERROR "Cannot locate the Python interpreter for native cleanup tests") + endif() + find_package(Python3 REQUIRED COMPONENTS Interpreter Development) + add_executable(fetch_binding_cleanup_test + ../../tests/native/test_fetch_binding_cleanup.cpp + $ + ) + target_include_directories(fetch_binding_cleanup_test PRIVATE + $) + target_compile_definitions(fetch_binding_cleanup_test PRIVATE + $) + target_compile_options(fetch_binding_cleanup_test PRIVATE + $) + target_link_libraries(fetch_binding_cleanup_test PRIVATE + $ Python3::Python ${CMAKE_DL_LIBS}) + foreach(failure IN ITEMS unbind rows_pointer) + add_test(NAME fetch_binding_cleanup_${failure} + COMMAND fetch_binding_cleanup_test ${failure}) + set_tests_properties(fetch_binding_cleanup_${failure} PROPERTIES TIMEOUT 30) + endforeach() +endif() diff --git a/tests/native/test_fetch_binding_cleanup.cpp b/tests/native/test_fetch_binding_cleanup.cpp new file mode 100644 index 000000000..ad02f16d3 --- /dev/null +++ b/tests/native/test_fetch_binding_cleanup.cpp @@ -0,0 +1,256 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#include "ddbc_bindings.h" +#include +#include +#include +#include +#include +#include + +SQLRETURN FetchMany_wrap(SqlHandlePtr, py::list&, int, const std::string&, + const std::string&, int, py::handle); + +namespace { + +enum class Call { bind, unbind, rowsPointer, rowArraySize, getArraySize, freeHandle }; +enum class Failure { none, unbind, rowsPointer }; +constexpr SQLLEN integerBytes = sizeof(SQLINTEGER); + +void require(bool condition, const char* message) { + if (!condition) { + throw std::runtime_error(message); + } +} + +struct Driver { + Failure failure; + std::array calls{}; + size_t callCount = 0; + SQLULEN arraySize = 1; + SQLULEN* rowsFetched = nullptr; + SQLINTEGER* values = nullptr; + SQLLEN* indicators = nullptr; + int diagnostic = 0; + std::weak_ptr owner; + bool checkOwner = false; + bool prematureRelease = false; + + SQLRETURN record(Call call) noexcept { + diagnostic = 0; // Another ODBC call would overwrite the original failure. + if (checkOwner && call != Call::freeHandle && owner.expired()) { + prematureRelease = true; + } + if (callCount == calls.size()) { + return SQL_ERROR; + } + calls[callCount++] = call; + return SQL_SUCCESS; + } + + SQLRETURN fail() noexcept { + diagnostic = 42; + return SQL_ERROR; + } + + void expect(std::initializer_list expected) { + require(callCount == expected.size(), "Unexpected ODBC call after failed cleanup"); + require(std::equal(expected.begin(), expected.end(), calls.begin()), + "Incorrect cleanup call order"); + callCount = 0; + } +}; + +SQLRETURN SQL_API setAttribute(SQLHSTMT handle, SQLINTEGER attribute, SQLPOINTER value, + SQLINTEGER) { + auto& driver = *static_cast(handle); + if (attribute == SQL_ATTR_ROWS_FETCHED_PTR) { + if (!SQL_SUCCEEDED(driver.record(Call::rowsPointer))) { + return SQL_ERROR; + } + if (!value && driver.failure == Failure::rowsPointer) { + return driver.fail(); + } + driver.rowsFetched = static_cast(value); + return SQL_SUCCESS; + } + if (attribute == SQL_ATTR_ROW_ARRAY_SIZE) { + if (!SQL_SUCCEEDED(driver.record(Call::rowArraySize))) { + return SQL_ERROR; + } + driver.arraySize = static_cast(reinterpret_cast(value)); + return SQL_SUCCESS; + } + return SQL_ERROR; +} + +SQLRETURN SQL_API getAttribute(SQLHSTMT handle, SQLINTEGER attribute, SQLPOINTER value, + SQLINTEGER, SQLINTEGER*) { + auto& driver = *static_cast(handle); + if (attribute != SQL_ATTR_ROW_ARRAY_SIZE || + !SQL_SUCCEEDED(driver.record(Call::getArraySize))) { + return SQL_ERROR; + } + *static_cast(value) = driver.arraySize; + return SQL_SUCCESS; +} + +SQLRETURN SQL_API bindColumn(SQLHSTMT handle, SQLUSMALLINT column, SQLSMALLINT type, + SQLPOINTER values, SQLLEN length, SQLLEN* indicators) { + auto& driver = *static_cast(handle); + if (!SQL_SUCCEEDED(driver.record(Call::bind)) || column != 1 || type != SQL_C_LONG || + length != integerBytes) { + return SQL_ERROR; + } + driver.values = static_cast(values); + driver.indicators = indicators; + return SQL_SUCCESS; +} + +SQLRETURN SQL_API freeStatement(SQLHSTMT handle, SQLUSMALLINT option) { + auto& driver = *static_cast(handle); + if (option != SQL_UNBIND || !SQL_SUCCEEDED(driver.record(Call::unbind))) { + return SQL_ERROR; + } + if (driver.failure == Failure::unbind) { + return driver.fail(); + } + driver.values = nullptr; + driver.indicators = nullptr; + return SQL_SUCCESS; +} + +SQLRETURN SQL_API freeHandle(SQLSMALLINT, SQLHANDLE handle) { + auto& driver = *static_cast(handle); + driver.record(Call::freeHandle); + driver.values = nullptr; + driver.indicators = nullptr; + driver.rowsFetched = nullptr; + return SQL_SUCCESS; +} + +void checkRetained(const SqlHandlePtr& statement, Driver& driver, + FetchBindingPlan* address, SQLINTEGER* values, SQLLEN* indicators, + SQLULEN* rowsFetched) { + require(driver.owner.use_count() == 1, "The statement must be the only plan owner"); + auto retained = statement->fetchBindings.snapshot(); + require(retained && retained.get() == address, "Failed detach discarded the owned plan"); + require(!retained->matches({retained->generation, retained->metadata}, 2, + "ascii", "utf-16le", SQL_C_CHAR), + "Failed detach left the plan eligible for reuse"); + require(retained->buffers.intBuffers[0].data() == values && + retained->buffers.indicators[0].data() == indicators && + &retained->rowsFetched == rowsFetched, + "Failed detach replaced driver-referenced storage"); + require(driver.rowsFetched == rowsFetched && driver.arraySize == 2, + "Failed cleanup continued resetting statement attributes"); + *driver.rowsFetched = 1; + require(retained->rowsFetched == 1, "Driver rows-fetched pointer lost its owner"); + if (driver.failure == Failure::unbind) { + require(driver.values == values && driver.indicators == indicators, + "Failed unbind lost the column addresses"); + driver.values[0] = 73; + driver.indicators[0] = integerBytes; + require(retained->buffers.intBuffers[0][0] == 73 && + retained->buffers.indicators[0][0] == integerBytes, + "Driver column pointers lost their owners"); + } + bool replacementRejected = false; + try { + statement->fetchBindings.install(retained); + } catch (const std::logic_error&) { + replacementRejected = true; + } + require(replacementRejected, "A still-bound plan could be replaced"); + require(!driver.prematureRelease, "An ODBC callback observed expired ownership"); +} + +void testCleanup(Failure failure) { + Driver driver{Failure::none}; + auto statement = std::make_shared( + SQL_HANDLE_STMT, &driver, std::make_shared()); + auto metadata = std::make_shared(); + metadata->columns.push_back({u"value", SQL_INTEGER, 10}); + metadata->namesValidated = true; + statement->resultMetadata.publish(0, metadata); + auto plan = std::shared_ptr( + new FetchBindingPlan(statement->resultMetadata.snapshot(), 2, + "ascii", "utf-16le", SQL_C_CHAR), + FetchBindingPlan::Deleter{}); + plan->buffers.intBuffers[0].resize(2); + auto* values = plan->buffers.intBuffers[0].data(); + auto* indicators = plan->buffers.indicators[0].data(); + auto* rowsFetched = &plan->rowsFetched; + auto* address = plan.get(); + plan->bindings.push_back({1, SQL_C_LONG, values, integerBytes, indicators}); + statement->fetchBindings.install(plan); + require(SQL_SUCCEEDED(plan->attach(statement->get())), "Initial binding failed"); + require(plan->matches(statement->resultMetadata.snapshot(), 2, + "ascii", "utf-16le", SQL_C_CHAR), + "Initial plan was not reusable"); + driver.expect({Call::rowArraySize, Call::getArraySize, Call::rowsPointer, Call::bind}); + driver.owner = plan; + std::weak_ptr metadataLifetime = metadata; + metadata.reset(); + plan.reset(); + driver.checkOwner = true; + driver.failure = failure; + + for (int attempt = 0; attempt < 3; ++attempt) { + py::list rows; + // Both formerly-compatible and resized fetches must retry cleanup, not rebind/fetch. + SQLRETURN result = attempt == 0 + ? statement->detachFetchBindings() + : FetchMany_wrap(statement, rows, attempt == 1 ? 2 : 3, + "ascii", "utf-16le", SQL_C_CHAR, py::none()); + require(result == SQL_ERROR && rows.empty(), "Cleanup failure was not propagated"); + require(driver.diagnostic == 42, "Cleanup overwrote the first error diagnostic"); + if (failure == Failure::unbind) { + driver.expect({Call::unbind}); + } else { + driver.expect({Call::unbind, Call::rowsPointer}); + } + require(!statement->resultMetadata.snapshot().metadata, + "Failed cleanup must invalidate the result metadata"); + checkRetained(statement, driver, address, values, indicators, rowsFetched); + require(metadataLifetime.use_count() == 1, "Plan no longer owns its native metadata"); + } + + driver.failure = Failure::none; + require(SQL_SUCCEEDED(statement->detachFetchBindings()), "Cleanup retry failed"); + driver.expect({Call::unbind, Call::rowsPointer, Call::rowArraySize}); + require(!driver.values && !driver.indicators && !driver.rowsFetched && driver.arraySize == 1, + "Successful cleanup left driver pointers installed"); + require(!statement->fetchBindings.hasPlan() && driver.owner.expired(), + "Successful cleanup retained the plan"); + // Metadata is the first plan member and is destroyed after all column buffers. + // Expiry also rejects the emergency deleter's deliberate terminal-retention path. + require(metadataLifetime.expired(), "Successful cleanup leaked the plan allocation"); + require(!driver.prematureRelease, "Storage was released before the final ODBC use"); + require(SQL_SUCCEEDED(statement->freeHandle()), "Statement free failed"); + driver.expect({Call::freeHandle}); +} + +} // namespace + +int main(int argc, char** argv) { + py::scoped_interpreter interpreter{}; + try { + require(argc == 2, "Expected unbind or rows_pointer"); + std::string mode = argv[1]; + require(mode == "unbind" || mode == "rows_pointer", "Unknown failure mode"); + SQLSetStmtAttr_ptr = setAttribute; + SQLGetStmtAttr_ptr = getAttribute; + SQLBindCol_ptr = bindColumn; + SQLFreeStmt_ptr = freeStatement; + SQLFreeHandle_ptr = freeHandle; + testCleanup(mode == "unbind" ? Failure::unbind : Failure::rowsPointer); + std::cout << "PASS " << mode + << ": retained ownership, blocked reuse, successful retry and destruction\n"; + return 0; + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + return 1; + } +} From ad0a8550f3523b237db29bcde9c152e168aad136 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 14:49:48 +0530 Subject: [PATCH 11/15] CHORE: Revert cleanup-failure test infrastructure Revert 1c05d11e1a0f0e519874513bff74c011d5d287d0 after the user rejected the expanded scope. Restore the exact be219eb3 source tree, preserving the earlier review fixes and profiler regressions. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .github/prompts/run-tests.prompt.md | 17 -- eng/pipelines/pr-validation-pipeline.yml | 3 - mssql_python/pybind/CMakeLists.txt | 32 --- tests/native/test_fetch_binding_cleanup.cpp | 256 -------------------- 4 files changed, 308 deletions(-) delete mode 100644 tests/native/test_fetch_binding_cleanup.cpp diff --git a/.github/prompts/run-tests.prompt.md b/.github/prompts/run-tests.prompt.md index ce0bef757..da8bcfa88 100644 --- a/.github/prompts/run-tests.prompt.md +++ b/.github/prompts/run-tests.prompt.md @@ -116,23 +116,6 @@ python main.py ## TASK -### Native fetch cleanup ownership regressions (no database) - -After configuring the native extension with the normal platform/compiler settings, -enable `-DBUILD_FETCH_BINDING_TESTS=ON` in the same CMake build directory. Build -the `fetch_binding_cleanup_test` target in Release, then run -`ctest --test-dir -C Release --output-on-failure -R "^fetch_binding_cleanup_"`. - -These opt-in tests compile the actual native sources with the extension's settings -and embed the selected Python interpreter. They require Python development headers -and its embedding library, but no SQL Server, connection string, or profiling build. -Only the existing ODBC function pointers are replaced, in the test executable. -Runtime checks remain enabled under `NDEBUG` and verify retained buffer ownership, -non-reuse, original-error preservation, and destruction after successful cleanup -following injected unbind and rows-fetched-pointer reset failures. Ubuntu PR CI runs -them before uninstalling the ODBC development headers. Normal wheel builds leave -the option OFF. - Help the developer run tests to validate their changes. Follow this process based on what they need. --- diff --git a/eng/pipelines/pr-validation-pipeline.yml b/eng/pipelines/pr-validation-pipeline.yml index 75e254c79..99507c8e6 100644 --- a/eng/pipelines/pr-validation-pipeline.yml +++ b/eng/pipelines/pr-validation-pipeline.yml @@ -885,9 +885,6 @@ jobs: source /opt/venv/bin/activate cd /workspace python -m eng.profiler_benchmarks.controller --check-build on - cmake -S mssql_python/pybind -B mssql_python/pybind/build -DBUILD_FETCH_BINDING_TESTS=ON - cmake --build mssql_python/pybind/build --config Release --target fetch_binding_cleanup_test - ctest --test-dir mssql_python/pybind/build -C Release --output-on-failure -R '^fetch_binding_cleanup_' " fi displayName: 'Build pybind bindings (.so) in $(distroName) container' diff --git a/mssql_python/pybind/CMakeLists.txt b/mssql_python/pybind/CMakeLists.txt index e062683e1..77d599bd5 100644 --- a/mssql_python/pybind/CMakeLists.txt +++ b/mssql_python/pybind/CMakeLists.txt @@ -393,35 +393,3 @@ if(APPLE) target_compile_definitions(ddbc_bindings PRIVATE MACOS_STRING_FIX) target_compile_options(ddbc_bindings PRIVATE -DAPPLE_SILICON) endif() - -option(BUILD_FETCH_BINDING_TESTS "Build native fetch cleanup ownership tests (no database)" OFF) -if(BUILD_FETCH_BINDING_TESTS) - enable_testing() - execute_process( - COMMAND python -c "import sys; print(sys.executable)" - OUTPUT_VARIABLE Python3_EXECUTABLE - OUTPUT_STRIP_TRAILING_WHITESPACE - RESULT_VARIABLE test_python_status - ) - if(NOT test_python_status EQUAL 0) - message(FATAL_ERROR "Cannot locate the Python interpreter for native cleanup tests") - endif() - find_package(Python3 REQUIRED COMPONENTS Interpreter Development) - add_executable(fetch_binding_cleanup_test - ../../tests/native/test_fetch_binding_cleanup.cpp - $ - ) - target_include_directories(fetch_binding_cleanup_test PRIVATE - $) - target_compile_definitions(fetch_binding_cleanup_test PRIVATE - $) - target_compile_options(fetch_binding_cleanup_test PRIVATE - $) - target_link_libraries(fetch_binding_cleanup_test PRIVATE - $ Python3::Python ${CMAKE_DL_LIBS}) - foreach(failure IN ITEMS unbind rows_pointer) - add_test(NAME fetch_binding_cleanup_${failure} - COMMAND fetch_binding_cleanup_test ${failure}) - set_tests_properties(fetch_binding_cleanup_${failure} PROPERTIES TIMEOUT 30) - endforeach() -endif() diff --git a/tests/native/test_fetch_binding_cleanup.cpp b/tests/native/test_fetch_binding_cleanup.cpp deleted file mode 100644 index ad02f16d3..000000000 --- a/tests/native/test_fetch_binding_cleanup.cpp +++ /dev/null @@ -1,256 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -#include "ddbc_bindings.h" -#include -#include -#include -#include -#include -#include - -SQLRETURN FetchMany_wrap(SqlHandlePtr, py::list&, int, const std::string&, - const std::string&, int, py::handle); - -namespace { - -enum class Call { bind, unbind, rowsPointer, rowArraySize, getArraySize, freeHandle }; -enum class Failure { none, unbind, rowsPointer }; -constexpr SQLLEN integerBytes = sizeof(SQLINTEGER); - -void require(bool condition, const char* message) { - if (!condition) { - throw std::runtime_error(message); - } -} - -struct Driver { - Failure failure; - std::array calls{}; - size_t callCount = 0; - SQLULEN arraySize = 1; - SQLULEN* rowsFetched = nullptr; - SQLINTEGER* values = nullptr; - SQLLEN* indicators = nullptr; - int diagnostic = 0; - std::weak_ptr owner; - bool checkOwner = false; - bool prematureRelease = false; - - SQLRETURN record(Call call) noexcept { - diagnostic = 0; // Another ODBC call would overwrite the original failure. - if (checkOwner && call != Call::freeHandle && owner.expired()) { - prematureRelease = true; - } - if (callCount == calls.size()) { - return SQL_ERROR; - } - calls[callCount++] = call; - return SQL_SUCCESS; - } - - SQLRETURN fail() noexcept { - diagnostic = 42; - return SQL_ERROR; - } - - void expect(std::initializer_list expected) { - require(callCount == expected.size(), "Unexpected ODBC call after failed cleanup"); - require(std::equal(expected.begin(), expected.end(), calls.begin()), - "Incorrect cleanup call order"); - callCount = 0; - } -}; - -SQLRETURN SQL_API setAttribute(SQLHSTMT handle, SQLINTEGER attribute, SQLPOINTER value, - SQLINTEGER) { - auto& driver = *static_cast(handle); - if (attribute == SQL_ATTR_ROWS_FETCHED_PTR) { - if (!SQL_SUCCEEDED(driver.record(Call::rowsPointer))) { - return SQL_ERROR; - } - if (!value && driver.failure == Failure::rowsPointer) { - return driver.fail(); - } - driver.rowsFetched = static_cast(value); - return SQL_SUCCESS; - } - if (attribute == SQL_ATTR_ROW_ARRAY_SIZE) { - if (!SQL_SUCCEEDED(driver.record(Call::rowArraySize))) { - return SQL_ERROR; - } - driver.arraySize = static_cast(reinterpret_cast(value)); - return SQL_SUCCESS; - } - return SQL_ERROR; -} - -SQLRETURN SQL_API getAttribute(SQLHSTMT handle, SQLINTEGER attribute, SQLPOINTER value, - SQLINTEGER, SQLINTEGER*) { - auto& driver = *static_cast(handle); - if (attribute != SQL_ATTR_ROW_ARRAY_SIZE || - !SQL_SUCCEEDED(driver.record(Call::getArraySize))) { - return SQL_ERROR; - } - *static_cast(value) = driver.arraySize; - return SQL_SUCCESS; -} - -SQLRETURN SQL_API bindColumn(SQLHSTMT handle, SQLUSMALLINT column, SQLSMALLINT type, - SQLPOINTER values, SQLLEN length, SQLLEN* indicators) { - auto& driver = *static_cast(handle); - if (!SQL_SUCCEEDED(driver.record(Call::bind)) || column != 1 || type != SQL_C_LONG || - length != integerBytes) { - return SQL_ERROR; - } - driver.values = static_cast(values); - driver.indicators = indicators; - return SQL_SUCCESS; -} - -SQLRETURN SQL_API freeStatement(SQLHSTMT handle, SQLUSMALLINT option) { - auto& driver = *static_cast(handle); - if (option != SQL_UNBIND || !SQL_SUCCEEDED(driver.record(Call::unbind))) { - return SQL_ERROR; - } - if (driver.failure == Failure::unbind) { - return driver.fail(); - } - driver.values = nullptr; - driver.indicators = nullptr; - return SQL_SUCCESS; -} - -SQLRETURN SQL_API freeHandle(SQLSMALLINT, SQLHANDLE handle) { - auto& driver = *static_cast(handle); - driver.record(Call::freeHandle); - driver.values = nullptr; - driver.indicators = nullptr; - driver.rowsFetched = nullptr; - return SQL_SUCCESS; -} - -void checkRetained(const SqlHandlePtr& statement, Driver& driver, - FetchBindingPlan* address, SQLINTEGER* values, SQLLEN* indicators, - SQLULEN* rowsFetched) { - require(driver.owner.use_count() == 1, "The statement must be the only plan owner"); - auto retained = statement->fetchBindings.snapshot(); - require(retained && retained.get() == address, "Failed detach discarded the owned plan"); - require(!retained->matches({retained->generation, retained->metadata}, 2, - "ascii", "utf-16le", SQL_C_CHAR), - "Failed detach left the plan eligible for reuse"); - require(retained->buffers.intBuffers[0].data() == values && - retained->buffers.indicators[0].data() == indicators && - &retained->rowsFetched == rowsFetched, - "Failed detach replaced driver-referenced storage"); - require(driver.rowsFetched == rowsFetched && driver.arraySize == 2, - "Failed cleanup continued resetting statement attributes"); - *driver.rowsFetched = 1; - require(retained->rowsFetched == 1, "Driver rows-fetched pointer lost its owner"); - if (driver.failure == Failure::unbind) { - require(driver.values == values && driver.indicators == indicators, - "Failed unbind lost the column addresses"); - driver.values[0] = 73; - driver.indicators[0] = integerBytes; - require(retained->buffers.intBuffers[0][0] == 73 && - retained->buffers.indicators[0][0] == integerBytes, - "Driver column pointers lost their owners"); - } - bool replacementRejected = false; - try { - statement->fetchBindings.install(retained); - } catch (const std::logic_error&) { - replacementRejected = true; - } - require(replacementRejected, "A still-bound plan could be replaced"); - require(!driver.prematureRelease, "An ODBC callback observed expired ownership"); -} - -void testCleanup(Failure failure) { - Driver driver{Failure::none}; - auto statement = std::make_shared( - SQL_HANDLE_STMT, &driver, std::make_shared()); - auto metadata = std::make_shared(); - metadata->columns.push_back({u"value", SQL_INTEGER, 10}); - metadata->namesValidated = true; - statement->resultMetadata.publish(0, metadata); - auto plan = std::shared_ptr( - new FetchBindingPlan(statement->resultMetadata.snapshot(), 2, - "ascii", "utf-16le", SQL_C_CHAR), - FetchBindingPlan::Deleter{}); - plan->buffers.intBuffers[0].resize(2); - auto* values = plan->buffers.intBuffers[0].data(); - auto* indicators = plan->buffers.indicators[0].data(); - auto* rowsFetched = &plan->rowsFetched; - auto* address = plan.get(); - plan->bindings.push_back({1, SQL_C_LONG, values, integerBytes, indicators}); - statement->fetchBindings.install(plan); - require(SQL_SUCCEEDED(plan->attach(statement->get())), "Initial binding failed"); - require(plan->matches(statement->resultMetadata.snapshot(), 2, - "ascii", "utf-16le", SQL_C_CHAR), - "Initial plan was not reusable"); - driver.expect({Call::rowArraySize, Call::getArraySize, Call::rowsPointer, Call::bind}); - driver.owner = plan; - std::weak_ptr metadataLifetime = metadata; - metadata.reset(); - plan.reset(); - driver.checkOwner = true; - driver.failure = failure; - - for (int attempt = 0; attempt < 3; ++attempt) { - py::list rows; - // Both formerly-compatible and resized fetches must retry cleanup, not rebind/fetch. - SQLRETURN result = attempt == 0 - ? statement->detachFetchBindings() - : FetchMany_wrap(statement, rows, attempt == 1 ? 2 : 3, - "ascii", "utf-16le", SQL_C_CHAR, py::none()); - require(result == SQL_ERROR && rows.empty(), "Cleanup failure was not propagated"); - require(driver.diagnostic == 42, "Cleanup overwrote the first error diagnostic"); - if (failure == Failure::unbind) { - driver.expect({Call::unbind}); - } else { - driver.expect({Call::unbind, Call::rowsPointer}); - } - require(!statement->resultMetadata.snapshot().metadata, - "Failed cleanup must invalidate the result metadata"); - checkRetained(statement, driver, address, values, indicators, rowsFetched); - require(metadataLifetime.use_count() == 1, "Plan no longer owns its native metadata"); - } - - driver.failure = Failure::none; - require(SQL_SUCCEEDED(statement->detachFetchBindings()), "Cleanup retry failed"); - driver.expect({Call::unbind, Call::rowsPointer, Call::rowArraySize}); - require(!driver.values && !driver.indicators && !driver.rowsFetched && driver.arraySize == 1, - "Successful cleanup left driver pointers installed"); - require(!statement->fetchBindings.hasPlan() && driver.owner.expired(), - "Successful cleanup retained the plan"); - // Metadata is the first plan member and is destroyed after all column buffers. - // Expiry also rejects the emergency deleter's deliberate terminal-retention path. - require(metadataLifetime.expired(), "Successful cleanup leaked the plan allocation"); - require(!driver.prematureRelease, "Storage was released before the final ODBC use"); - require(SQL_SUCCEEDED(statement->freeHandle()), "Statement free failed"); - driver.expect({Call::freeHandle}); -} - -} // namespace - -int main(int argc, char** argv) { - py::scoped_interpreter interpreter{}; - try { - require(argc == 2, "Expected unbind or rows_pointer"); - std::string mode = argv[1]; - require(mode == "unbind" || mode == "rows_pointer", "Unknown failure mode"); - SQLSetStmtAttr_ptr = setAttribute; - SQLGetStmtAttr_ptr = getAttribute; - SQLBindCol_ptr = bindColumn; - SQLFreeStmt_ptr = freeStatement; - SQLFreeHandle_ptr = freeHandle; - testCleanup(mode == "unbind" ? Failure::unbind : Failure::rowsPointer); - std::cout << "PASS " << mode - << ": retained ownership, blocked reuse, successful retry and destruction\n"; - return 0; - } catch (const std::exception& error) { - std::cerr << error.what() << '\n'; - return 1; - } -} From 99c5a56913b18ff8080ef64702eb8074139a70f1 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 15:09:27 +0530 Subject: [PATCH 12/15] FIX: Add focused failed-unbind regression Exercise an injected SQL_UNBIND failure in the existing isolated profiler test suite. Assert that failed cleanup blocks reuse and rebinding, then verify recovery and EOF without changing production or build configuration. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- tests/test_025_profiler.py | 92 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 92 insertions(+) diff --git a/tests/test_025_profiler.py b/tests/test_025_profiler.py index 5ad1b7a58..073945813 100644 --- a/tests/test_025_profiler.py +++ b/tests/test_025_profiler.py @@ -486,6 +486,98 @@ def fetch(cursor, size, first): assert result.returncode == 0, result.stdout + result.stderr +@_needs_cpp +@_needs_db +@pytest.mark.skipif( + sys.platform == "win32", reason="Windows does not export the ODBC function-pointer globals" +) +def test_fetchmany_failed_unbind_blocks_reuse_until_cleanup_succeeds(): + """A failed detach must not permit fetching or replacing the retained binding plan.""" + script = textwrap.dedent(""" + import ctypes + import os + import sys + + sys.path.insert(0, sys.argv[1]) + import mssql_python as db + from mssql_python import ddbc_bindings as native + + assert hasattr(native, "profiling") + assert os.path.realpath(native.module.__file__) == sys.argv[2] + library = ctypes.CDLL(sys.argv[2]) + free_stmt = ctypes.c_void_p.in_dll(library, "SQLFreeStmt_ptr") + callback_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p, ctypes.c_ushort) + + def counts(expected): + stats = native.profiling.get_stats() + actual = tuple( + stats.get("ddbc::fetch_bindings::" + name, {}).get("calls", 0) + for name in ("plan_allocation", "SQLBindCol", "SQL_UNBIND") + ) + assert actual == expected, (actual, expected, stats) + + try: + connection = db.connect(os.environ["DB_CONNECTION_STRING"], timeout=5) + except db.Error: + raise RuntimeError("SQL connection failed") from None + with connection, connection.cursor() as cursor: + cursor.execute("SELECT n FROM (VALUES (1), (2), (3), (4)) AS v(n) ORDER BY n") + native.profiling.reset() + native.profiling.enable() + try: + assert [tuple(row) for row in cursor.fetchmany(2)] == [(1,), (2,)] + counts((1, 1, 0)) + original = free_stmt.value + assert original + original_call = callback_type(original) + failures = [] + + @callback_type + def fail_unbind(handle, option): + if option == 2: # SQL_UNBIND + failures.append(handle) + return -1 # SQL_ERROR, without releasing the driver's bindings + return original_call(handle, option) + + try: + free_stmt.value = ctypes.cast(fail_unbind, ctypes.c_void_p).value + for attempt, size in enumerate((3, 2), 1): + rows = [] + ret = native.DDBCSQLFetchMany( + cursor.hstmt, rows, size, cursor._cached_char_encoding, + cursor._cached_wchar_encoding, cursor._cached_char_ctype, + ) + assert ret == -1 and rows == [], (ret, rows) + assert len(failures) == attempt, failures + counts((1, 1, attempt)) + finally: + free_stmt.value = original + + assert [tuple(row) for row in cursor.fetchmany(2)] == [(3,), (4,)] + counts((2, 2, 3)) + assert cursor.fetchmany(2) == [] + counts((2, 2, 4)) + finally: + native.profiling.disable() + native.profiling.reset() + """) + result = subprocess.run( + [ + sys.executable, + "-E", + "-c", + script, + os.path.dirname(os.path.dirname(perf_timer.__file__)), + os.path.realpath(ddbc.module.__file__), + ], + capture_output=True, + text=True, + timeout=45, + ) + assert result.returncode == 0, result.stdout + result.stderr + assert "retaining fetch buffers" not in result.stderr, result.stderr + + @_needs_cpp @_needs_db def test_cpp_timeline_captures_events(): From db19e6375c138cfa08ac3fe51ed6d5c9837c546e Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 15:10:45 +0530 Subject: [PATCH 13/15] FIX: Assert failed cleanup cannot advance native fetches Check the existing SQLFetchScroll call counter alongside plan allocation and binding counts in the single cleanup regression. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- tests/test_025_profiler.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/tests/test_025_profiler.py b/tests/test_025_profiler.py index 073945813..4da895762 100644 --- a/tests/test_025_profiler.py +++ b/tests/test_025_profiler.py @@ -508,13 +508,14 @@ def test_fetchmany_failed_unbind_blocks_reuse_until_cleanup_succeeds(): free_stmt = ctypes.c_void_p.in_dll(library, "SQLFreeStmt_ptr") callback_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p, ctypes.c_ushort) - def counts(expected): + def counts(expected, fetches): stats = native.profiling.get_stats() actual = tuple( stats.get("ddbc::fetch_bindings::" + name, {}).get("calls", 0) for name in ("plan_allocation", "SQLBindCol", "SQL_UNBIND") ) assert actual == expected, (actual, expected, stats) + assert stats["ddbc::FetchBatchData::SQLFetchScroll_call"]["calls"] == fetches, stats try: connection = db.connect(os.environ["DB_CONNECTION_STRING"], timeout=5) @@ -526,7 +527,7 @@ def counts(expected): native.profiling.enable() try: assert [tuple(row) for row in cursor.fetchmany(2)] == [(1,), (2,)] - counts((1, 1, 0)) + counts((1, 1, 0), 1) original = free_stmt.value assert original original_call = callback_type(original) @@ -549,14 +550,14 @@ def fail_unbind(handle, option): ) assert ret == -1 and rows == [], (ret, rows) assert len(failures) == attempt, failures - counts((1, 1, attempt)) + counts((1, 1, attempt), 1) finally: free_stmt.value = original assert [tuple(row) for row in cursor.fetchmany(2)] == [(3,), (4,)] - counts((2, 2, 3)) + counts((2, 2, 3), 2) assert cursor.fetchmany(2) == [] - counts((2, 2, 4)) + counts((2, 2, 4), 3) finally: native.profiling.disable() native.profiling.reset() From 39e220980e4eac745bb023255e8ae76d12a4a9b4 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 15:59:56 +0530 Subject: [PATCH 14/15] FIX: Cover failed rows-fetched pointer cleanup Parameterize the existing isolated cleanup regression to fail either SQL_UNBIND or clearing SQL_ATTR_ROWS_FETCHED_PTR. Preserve no-allocation, no-bind, no-fetch, retry, and recovery assertions without changing production or build configuration. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- tests/test_025_profiler.py | 36 ++++++++++++++++++++++++++---------- 1 file changed, 26 insertions(+), 10 deletions(-) diff --git a/tests/test_025_profiler.py b/tests/test_025_profiler.py index 4da895762..8d094d118 100644 --- a/tests/test_025_profiler.py +++ b/tests/test_025_profiler.py @@ -491,7 +491,8 @@ def fetch(cursor, size, first): @pytest.mark.skipif( sys.platform == "win32", reason="Windows does not export the ODBC function-pointer globals" ) -def test_fetchmany_failed_unbind_blocks_reuse_until_cleanup_succeeds(): +@pytest.mark.parametrize("failure_point", ("unbind", "rows_fetched_ptr")) +def test_fetchmany_failed_cleanup_blocks_reuse_until_cleanup_succeeds(failure_point): """A failed detach must not permit fetching or replacing the retained binding plan.""" script = textwrap.dedent(""" import ctypes @@ -505,8 +506,17 @@ def test_fetchmany_failed_unbind_blocks_reuse_until_cleanup_succeeds(): assert hasattr(native, "profiling") assert os.path.realpath(native.module.__file__) == sys.argv[2] library = ctypes.CDLL(sys.argv[2]) - free_stmt = ctypes.c_void_p.in_dll(library, "SQLFreeStmt_ptr") - callback_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p, ctypes.c_ushort) + failure_point = sys.argv[3] + if failure_point == "unbind": + pointer_name = "SQLFreeStmt_ptr" + callback_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p, ctypes.c_ushort) + else: + pointer_name = "SQLSetStmtAttr_ptr" + callback_type = ctypes.CFUNCTYPE( + ctypes.c_short, ctypes.c_void_p, ctypes.c_int32, + ctypes.c_void_p, ctypes.c_int32, + ) + pointer = ctypes.c_void_p.in_dll(library, pointer_name) def counts(expected, fetches): stats = native.profiling.get_stats() @@ -528,20 +538,25 @@ def counts(expected, fetches): try: assert [tuple(row) for row in cursor.fetchmany(2)] == [(1,), (2,)] counts((1, 1, 0), 1) - original = free_stmt.value + original = pointer.value assert original original_call = callback_type(original) failures = [] @callback_type - def fail_unbind(handle, option): - if option == 2: # SQL_UNBIND + def fail_cleanup(handle, operation, *args): + # Fail only SQL_UNBIND or clearing SQL_ATTR_ROWS_FETCHED_PTR. + should_fail = ( + operation == 2 if failure_point == "unbind" + else operation == 26 and args[0] is None + ) + if should_fail: failures.append(handle) - return -1 # SQL_ERROR, without releasing the driver's bindings - return original_call(handle, option) + return -1 # SQL_ERROR, leaving the driver's pointers unchanged + return original_call(handle, operation, *args) try: - free_stmt.value = ctypes.cast(fail_unbind, ctypes.c_void_p).value + pointer.value = ctypes.cast(fail_cleanup, ctypes.c_void_p).value for attempt, size in enumerate((3, 2), 1): rows = [] ret = native.DDBCSQLFetchMany( @@ -552,7 +567,7 @@ def fail_unbind(handle, option): assert len(failures) == attempt, failures counts((1, 1, attempt), 1) finally: - free_stmt.value = original + pointer.value = original assert [tuple(row) for row in cursor.fetchmany(2)] == [(3,), (4,)] counts((2, 2, 3), 2) @@ -570,6 +585,7 @@ def fail_unbind(handle, option): script, os.path.dirname(os.path.dirname(perf_timer.__file__)), os.path.realpath(ddbc.module.__file__), + failure_point, ], capture_output=True, text=True, From 4d12d4f4a90dee64acd43f3ab81e714116e58be3 Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Mon, 28 Sep 2026 15:47:32 +0530 Subject: [PATCH 15/15] FIX: Avoid iostreams in native handle shutdown diagnostics Use fixed native stderr messages for failed handle cleanup and detach status without relying on C++ iostream teardown ordering. Cleanup and retained-buffer ownership remain unchanged. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index e4798b3cf..1470967df 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -1709,8 +1709,11 @@ SqlHandle::~SqlHandle() { // A failed free leaves a live handle. Detach if possible before // the last plan owner applies its native-only emergency policy. SQLRETURN detached = detachFetchBindings(); - std::cerr << "mssql-python: native handle cleanup failed (" << ret - << "), fetch buffer detach returned " << detached << '\n'; + std::fputs( + SQL_SUCCEEDED(detached) + ? "mssql-python: native handle cleanup failed; fetch buffer detach succeeded\n" + : "mssql-python: native handle cleanup failed; fetch buffer detach failed\n", + stderr); } } } catch (...) {