diff --git a/CHANGELOG.md b/CHANGELOG.md index 97fb442b1..2cac01d4d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -57,6 +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 +- 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 3cb568b4b..5068f3c19 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). @@ -2829,6 +2843,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); @@ -3008,9 +3023,43 @@ SQLSMALLINT SQLNumResultCols_wrap(SqlHandlePtr statementHandle) { return columnCount; } -// Wrap SQLDescribeCol -SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMetadata) { - PERF_TIMER("SQLDescribeCol_wrap"); +namespace { + +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 py::cast(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) { @@ -3033,19 +3082,18 @@ SQLRETURN SQLDescribeCol_wrap(SqlHandlePtr StatementHandle, py::list& ColumnMeta 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)) { - // 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)); + auto name = dupeSqlWCharAsUtf16Le( + ColumnName, std::min(static_cast(NameLength), + (sizeof(ColumnName) / sizeof(SQLWCHAR)) - 1)); + appendColumn(std::move(name), DataType, ColumnSize, DecimalDigits, Nullable); } else { return retcode; } @@ -3053,11 +3101,29 @@ 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"); + 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 = name, "DataType"_a = type, "ColumnSize"_a = size, + "DecimalDigits"_a = digits, "Nullable"_a = nullable)); + }); + return ret; +} + SQLRETURN SQLSpecialColumns_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT identifierType, const py::object& catalogObj, const py::object& schemaObj, const std::u16string& table, SQLSMALLINT scope, SQLSMALLINT nullable) { PERF_TIMER("SQLSpecialColumns_wrap"); + StatementHandle->resultMetadata.clear(); if (!SQLSpecialColumns_ptr) { ThrowStdException("SQLSpecialColumns function not loaded"); } @@ -3086,8 +3152,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 @@ -3290,24 +3361,57 @@ 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}); + } } // Preprocess sql_variant: detect underlying type to route to correct conversion logic @@ -3320,10 +3424,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; } @@ -3333,10 +3441,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; } @@ -3983,6 +4095,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; } @@ -4003,7 +4123,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; @@ -4024,16 +4145,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: @@ -4165,7 +4287,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; @@ -4174,7 +4296,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; @@ -4188,7 +4310,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", @@ -4238,9 +4361,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; @@ -4520,8 +4643,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; @@ -4660,26 +4783,54 @@ 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 - py::list columnNames; - ret = SQLDescribeCol_wrap(StatementHandle, columnNames); - 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"); + 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(); - - if (IsLobOrVariantColumn(dataType, columnSize)) { + const auto& column = columnNames.at(i); + if (IsLobOrVariantColumn(column.dataType, column.columnSize)) { lobColumns.push_back(i + 1); // 1-based } } @@ -4873,7 +5024,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); @@ -5781,12 +5933,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)) { @@ -5812,6 +5966,23 @@ 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 = columnNames[i].cast(); + SQLSMALLINT type = column["DataType"].cast(); + metadata->columns.push_back({ + column["ColumnName"].cast(), type, + type == SQL_SS_VARIANT ? 0 : column["ColumnSize"].cast()}); + } + StatementHandle->resultMetadata.publish(metadataSnapshot.generation, std::move(metadata)); while (true) { { // Release GIL during the blocking fetch @@ -5928,7 +6099,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) @@ -5960,6 +6132,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 " @@ -6177,6 +6350,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..3a9912860 --- /dev/null +++ b/mssql_python/pybind/result_metadata.hpp @@ -0,0 +1,75 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +struct FetchColumnMetadata { + std::u16string name; + SQLSMALLINT dataType; + SQLULEN columnSize; +}; + +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_004_cursor.py b/tests/test_004_cursor.py index 40ea4d09d..f30a24202 100644 --- a/tests/test_004_cursor.py +++ b/tests/test_004_cursor.py @@ -2961,13 +2961,26 @@ def test_executemany_DecimalMix_List(cursor, db_connection): def test_nextset(cursor): - """Test nextset""" - cursor.execute("SELECT * FROM #pytest_all_data_types WHERE id = 1;") - assert cursor.nextset() is False, "Nextset should return False" + """Test metadata invalidation when re-executing and skipping unread results.""" + cursor.execute( + "SELECT id AS first_id FROM #pytest_all_data_types WHERE id IN (1, 2) ORDER BY id;" + ) + first_row = cursor.fetchmany(1)[0] + assert first_row.first_id == 1 + # Equal column counts must not hide changed types or names. cursor.execute( - "SELECT * FROM #pytest_all_data_types WHERE id = 2; SELECT * FROM #pytest_all_data_types WHERE id = 3;" + "SELECT CAST(id AS NVARCHAR(10)) AS second_id FROM #pytest_all_data_types " + "WHERE id IN (2, 3) ORDER BY id; " + "SELECT id AS third_id FROM #pytest_all_data_types WHERE id = 3;" ) + second_row = cursor.fetchmany(1)[0] + assert second_row.second_id == "2" assert cursor.nextset() is True, "Nextset should return True" + assert cursor.fetchmany(1)[0].third_id == 3 + assert cursor.description[0][0] == "third_id" + assert first_row.first_id == 1 + assert second_row.second_id == "2" + assert cursor.nextset() is False, "Nextset should return False" def test_delete_table(cursor, db_connection): diff --git a/tests/test_016_connection_invalidation_segfault.py b/tests/test_016_connection_invalidation_segfault.py index 4ae07306a..e3c7e760d 100644 --- a/tests/test_016_connection_invalidation_segfault.py +++ b/tests/test_016_connection_invalidation_segfault.py @@ -49,10 +49,11 @@ def test_connection_invalidation_with_multiple_cursors(conn_str): # Create multiple cursors with statement handles cursors = [] + rows = [] for i in range(5): cursor = conn.cursor() cursor.execute("SELECT 1 AS id, 'test' AS name") - cursor.fetchall() # Fetch results to complete the query + rows.append(cursor.fetchmany(1)[0]) cursors.append(cursor) # Close connection without explicitly closing cursors first @@ -62,10 +63,10 @@ def test_connection_invalidation_with_multiple_cursors(conn_str): # Force garbage collection to trigger cursor cleanup # This is where the segfault would occur without the fix cursors = None + del cursor gc.collect() - # If we reach here without crashing, the fix is working - assert True + assert [(row.id, row.name) for row in rows] == [(1, "test")] * 5 def test_connection_invalidation_without_cursor_close(conn_str):