diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index f410631d..1ad82c63 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -6306,8 +6306,12 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row, } CaptureFetchDiagnostics(hStmt, ret, messages); if (SQL_SUCCEEDED(ret)) { - // Retrieve column count - SQLSMALLINT colCount = SQLNumResultCols_wrap(StatementHandle, messages); + const auto snapshot = StatementHandle->resultMetadata.snapshot(); + SQLSMALLINT colCount = snapshot.fullColumnCount; + if (colCount < 0) { + colCount = SQLNumResultCols_wrap(StatementHandle, messages); + StatementHandle->resultMetadata.publishFullColumnCount(snapshot.generation, colCount); + } ret = SQLGetData_wrap(StatementHandle, colCount, row, charEncoding, wcharEncoding, charCtype, messages); if (!SQL_SUCCEEDED(ret)) { diff --git a/mssql_python/pybind/result_metadata.hpp b/mssql_python/pybind/result_metadata.hpp index 3a991286..72ec0010 100644 --- a/mssql_python/pybind/result_metadata.hpp +++ b/mssql_python/pybind/result_metadata.hpp @@ -30,11 +30,12 @@ class ResultMetadataCache { struct Snapshot { uint64_t generation; std::shared_ptr metadata; + SQLSMALLINT fullColumnCount; }; Snapshot snapshot() const { std::lock_guard lock(mutex_); - return {generation_, metadata_}; + return {generation_, metadata_, fullColumnCount_}; } void publish(uint64_t generation, std::shared_ptr metadata) { @@ -44,16 +45,26 @@ class ResultMetadataCache { } } + void publishFullColumnCount(uint64_t generation, SQLSMALLINT columnCount) { + std::lock_guard lock(mutex_); + if (generation == generation_ && columnCount >= 0) { + fullColumnCount_ = columnCount; + } + } + void clear() { std::lock_guard lock(mutex_); ++generation_; metadata_.reset(); + fullColumnCount_ = -1; } private: mutable std::mutex mutex_; uint64_t generation_ = 0; std::shared_ptr metadata_; + // Prefix GetData metadata does not establish the full result cardinality. + SQLSMALLINT fullColumnCount_ = -1; }; class ResultMetadataFailureGuard { diff --git a/tests/test_fetch_settings_cache.py b/tests/test_fetch_settings_cache.py index 78ad5229..b9e63e7c 100644 --- a/tests/test_fetch_settings_cache.py +++ b/tests/test_fetch_settings_cache.py @@ -1111,3 +1111,253 @@ def test_fetch_error_is_raised_before_wrapping_rows(connection, method, bridge_n fetch_rows(cursor, method) assert cursor._next_row_index == position assert tuple(fetch_rows(cursor, method)[0]) == (1,) + + +@pytest.mark.skipif( + sys.platform == "win32", + reason="Windows does not export the native driver function-pointer globals", +) +@pytest.mark.parametrize( + "mode", + ( + "warm", + "prefix", + "shapes", + "nextset", + "fetch-error", + "count-error", + "decode-error", + "generation", + "mixed", + ), +) +def test_native_fetchone_full_column_count_cache(conn_str, mode): + if not conn_str: + pytest.skip("DB_CONNECTION_STRING is required") + code = ( + "import runpy, sys; " + "runpy.run_path(sys.argv[1])['_check_native_fetchone_full_column_count_cache']" + "(sys.argv[2], sys.argv[3])" + ) + result = subprocess.run( + [ + sys.executable, + "-c", + code, + str(Path(__file__).resolve()), + mode, + str(Path(mssql_python.ddbc_bindings.module.__file__).resolve()), + ], + capture_output=True, + text=True, + timeout=60, + ) + assert result.returncode == 0, (result.returncode, result.stdout, result.stderr) + + +def _check_native_fetchone_full_column_count_cache(mode, expected_native): + import ctypes + import os + + native = Path(mssql_python.ddbc_bindings.module.__file__).resolve() + assert native == Path(expected_native) + library = ctypes.CDLL(str(native)) + count_pointer = ctypes.c_void_p.in_dll(library, "SQLNumResultCols_ptr") + fetch_pointer = ctypes.c_void_p.in_dll(library, "SQLFetch_ptr") + count_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p, ctypes.POINTER(ctypes.c_short)) + fetch_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p) + success = ConstantsDDBC.SQL_SUCCESS.value + error = ConstantsDDBC.SQL_ERROR.value + ddbc = mssql_python.ddbc_bindings + count_calls, callback_errors = [], [] + fail_fetch = fail_count = invalidate_count = False + + @count_type + def count_columns(handle, count): + nonlocal fail_count, invalidate_count + try: + count_calls.append(handle) + if fail_count: + fail_count = False + return error + result = original_count(handle, count) + if invalidate_count: + invalidate_count = False + status = ddbc.DDBCSQLSetStmtAttr( + cursor.hstmt, ConstantsDDBC.SQL_ATTR_QUERY_TIMEOUT.value, 0 + ) + if status != success: + callback_errors.append("Failed to invalidate the metadata generation") + return error + return result + except BaseException as failure: + callback_errors.append(type(failure).__name__) + return error + + @fetch_type + def fetch_row(handle): + nonlocal fail_fetch + try: + if fail_fetch: + fail_fetch = False + return error + return original_fetch(handle) + except BaseException as failure: + callback_errors.append(type(failure).__name__) + return error + + try: + connection = mssql_python.connect(os.environ["DB_CONNECTION_STRING"], timeout=5) + except mssql_python.Error as failure: + raise AssertionError( + f"Connection failed: {type(failure).__name__}; connection details withheld" + ) from None + with connection, connection.cursor() as cursor: + query = ( + "SELECT n AS number, CAST(N'text' AS NVARCHAR(10)) AS txt " + "FROM (VALUES (1), (2), (3), (4), (5), (6)) AS v(n) ORDER BY n" + ) + cursor.execute(query) + saved_count, saved_fetch = count_pointer.value, fetch_pointer.value + assert saved_count and saved_fetch + original_count, original_fetch = count_type(saved_count), fetch_type(saved_fetch) + try: + count_pointer.value = ctypes.cast(count_columns, ctypes.c_void_p).value + fetch_pointer.value = ctypes.cast(fetch_row, ctypes.c_void_p).value + if mode == "prefix": + assert ddbc.DDBCSQLFetch(cursor.hstmt) == success + prefix = [] + assert ( + ddbc.DDBCSQLGetData( + cursor.hstmt, + 1, + prefix, + "utf-16le", + "utf-16le", + ConstantsDDBC.SQL_C_WCHAR.value, + ) + == success + ) + assert prefix == [1] + assert count_calls == [] + assert tuple(cursor.fetchone()) == (2, "text") + assert len(count_calls) == 1 + assert tuple(cursor.fetchone()) == (3, "text") + assert len(count_calls) == 1 + elif mode == "shapes": + for sql, expected in ( + ("SELECT 7 AS number", (7,)), + ("SELECT CAST(N'new' AS NVARCHAR(10)) AS txt", ("new",)), + ("SELECT CAST(NULL AS INT) AS empty_value, 9 AS number", (None, 9)), + ): + cursor.execute(sql + " FROM (VALUES (1), (2)) AS v(n)") + count_calls.clear() + for _ in range(2): + row = tuple(cursor.fetchone()) + assert row == expected + assert tuple(map(type, row)) == tuple(map(type, expected)) + assert len(count_calls) == 1 + assert cursor.fetchone() is None + assert len(count_calls) == 1 + elif mode == "nextset": + cursor.execute( + "SELECT 1 AS number FROM (VALUES (1), (2)) AS v(n); " + "SELECT CAST(N'changed' AS NVARCHAR(10)) AS txt " + "FROM (VALUES (1), (2)) AS v(n)" + ) + count_calls.clear() + assert [cursor.fetchone()[0] for _ in range(2)] == [1, 1] + assert len(count_calls) == 1 + assert cursor.nextset() + count_calls.clear() + assert [cursor.fetchone()[0] for _ in range(2)] == ["changed", "changed"] + assert len(count_calls) == 1 + elif mode == "fetch-error": + assert tuple(cursor.fetchone()) == (1, "text") + assert len(count_calls) == 1 + fail_fetch = True + row = [] + assert ddbc.DDBCSQLFetchOne(cursor.hstmt, row) == error + assert row == [] + assert len(count_calls) == 1 + assert tuple(cursor.fetchone()) == (2, "text") + assert len(count_calls) == 2 + assert tuple(cursor.fetchone()) == (3, "text") + assert len(count_calls) == 2 + elif mode == "count-error": + fail_count = True + with pytest.raises(mssql_python.DatabaseError): + ddbc.DDBCSQLFetchOne(cursor.hstmt, []) + assert len(count_calls) == 1 + assert tuple(cursor.fetchone()) == (2, "text") + assert len(count_calls) == 2 + assert tuple(cursor.fetchone()) == (3, "text") + assert len(count_calls) == 2 + elif mode == "decode-error": + cursor.execute( + "SELECT CASE WHEN n = 2 THEN CAST(0x00D8 AS NVARCHAR(10)) " + "ELSE CAST(N'ok' AS NVARCHAR(10)) END AS txt " + "FROM (VALUES (1), (2), (3), (4)) AS v(n) ORDER BY n" + ) + count_calls.clear() + assert cursor.fetchone()[0] == "ok" + with pytest.raises(UnicodeDecodeError): + cursor.fetchone() + assert len(count_calls) == 1 + assert [cursor.fetchone()[0] for _ in range(2)] == ["ok", "ok"] + assert len(count_calls) == 2 + elif mode == "generation": + invalidate_count = True + assert tuple(cursor.fetchone()) == (1, "text") + assert len(count_calls) == 1 + assert tuple(cursor.fetchone()) == (2, "text") + assert len(count_calls) == 2 + assert tuple(cursor.fetchone()) == (3, "text") + assert len(count_calls) == 2 + elif mode == "mixed": + assert [tuple(row) for row in cursor.fetchmany(2)] == [(1, "text"), (2, "text")] + count_calls.clear() + assert tuple(cursor.fetchone()) == (3, "text") + assert len(count_calls) == 1 + assert tuple(cursor.fetchmany(1)[0]) == (4, "text") + count_calls.clear() + assert tuple(cursor.fetchone()) == (5, "text") + assert count_calls == [] + assert [tuple(row) for row in cursor.fetchall()] == [(6, "text")] + count_calls.clear() + assert cursor.fetchone() is None + assert count_calls == [] + cursor.execute(query) + try: + other_connection = mssql_python.connect( + os.environ["DB_CONNECTION_STRING"], timeout=5 + ) + except mssql_python.Error as failure: + raise AssertionError( + f"Connection failed: {type(failure).__name__}; connection details withheld" + ) from None + with other_connection, other_connection.cursor() as other: + other.execute("SELECT CAST(N'other' AS NVARCHAR(10))") + count_calls.clear() + assert cursor.fetchone()[0] == 1 + assert other.fetchone()[0] == "other" + assert cursor.fetchone()[0] == 2 + assert len(count_calls) == 2 + else: + assert mode == "warm" + assert [tuple(cursor.fetchone()) for _ in range(6)] == [ + (n, "text") for n in range(1, 7) + ] + assert len(count_calls) == 1 + assert cursor.fetchone() is None + assert len(count_calls) == 1 + # Direct callers still perform a real query on every invocation. + assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2 + assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2 + assert len(count_calls) == 3 + assert not callback_errors, callback_errors + assert not cursor.messages + finally: + count_pointer.value, fetch_pointer.value = saved_count, saved_fetch + cursor.execute("SELECT 42") + assert cursor.fetchone()[0] == 42