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 f410631d4..1470967df 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -1587,14 +1587,137 @@ void DriverLoader::loadDriver() { } } +static void CaptureFetchBindingDiagnostics(SQLHSTMT stmt, SQLRETURN ret, + FetchBindingDiagnostics* diagnostics); +static void AppendFetchBindingDiagnostics(py::handle messages, + const FetchBindingDiagnostics& diagnostics, + bool preserveFailure = false); + +SQLRETURN FetchBindingPlan::attach(SQLHSTMT stmt, FetchBindingDiagnostics* diagnostics) { + reusable = false; + needsReset = true; + SQLRETURN ret; + { + PERF_TIMER("fetch_bindings::SQLSetStmtAttr::ROW_ARRAY_SIZE"); + ret = SQLSetStmtAttr_ptr( + stmt, SQL_ATTR_ROW_ARRAY_SIZE, + reinterpret_cast(static_cast(fetchSize)), 0); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + SQLULEN activeSize = 0; + { + PERF_TIMER("fetch_bindings::SQLGetStmtAttr"); + ret = SQLGetStmtAttr_ptr(stmt, SQL_ATTR_ROW_ARRAY_SIZE, &activeSize, 0, nullptr); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + 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; + { + PERF_TIMER("fetch_bindings::SQLSetStmtAttr::ROWS_FETCHED_PTR"); + ret = SQLSetStmtAttr_ptr(stmt, SQL_ATTR_ROWS_FETCHED_PTR, &rowsFetched, 0); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + for (const auto& column : bindings) { + { + PERF_TIMER("fetch_bindings::SQLBindCol"); + ret = SQLBindCol_ptr(stmt, column.column, column.cType, column.data, + column.bufferLength, column.indicators); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } + reusable = true; + return ret; +} + +SQLRETURN FetchBindingPlan::detach(SQLHSTMT stmt, FetchBindingDiagnostics* diagnostics) { + reusable = false; + if (!needsReset) { + return SQL_SUCCESS; + } + SQLRETURN ret; + { + PERF_TIMER("fetch_bindings::SQL_UNBIND"); + ret = SQLFreeStmt_ptr(stmt, SQL_UNBIND); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + { + PERF_TIMER("fetch_bindings::SQLSetStmtAttr::ROWS_FETCHED_PTR"); + ret = SQLSetStmtAttr_ptr(stmt, SQL_ATTR_ROWS_FETCHED_PTR, nullptr, 0); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + driverMayReference = false; + { + PERF_TIMER("fetch_bindings::SQLSetStmtAttr::ROW_ARRAY_SIZE"); + ret = SQLSetStmtAttr_ptr(stmt, SQL_ATTR_ROW_ARRAY_SIZE, reinterpret_cast(1), 0); + } + CaptureFetchBindingDiagnostics(stmt, ret, diagnostics); + if (SQL_SUCCEEDED(ret)) { + needsReset = false; + } + return ret; +} + +namespace { + +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::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 (...) { + std::fputs("mssql-python: unexpected failure during native handle cleanup\n", stderr); } } @@ -1615,7 +1738,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) { @@ -1627,6 +1750,7 @@ void SqlHandle::markImplicitlyFreed() { return; // Refuse to mark - let normal free() handle it } resultMetadata.clear(); + fetchBindings.nativeReleased(); _implicitly_freed = true; } @@ -1638,7 +1762,72 @@ 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(FetchBindingDiagnostics* diagnostics) { + 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, diagnostics); + if (SQL_SUCCEEDED(ret)) { + fetchBindings.remove(plan); + } + return ret; +} + +SQLRETURN SqlHandle::detachPresentFetchBindings(py::handle messages) { + auto plan = fetchBindings.snapshot(); + if (!plan) { + return SQL_SUCCESS; + } + if (is_python_finalizing()) { + return SQL_ERROR; + } + FetchBindingDiagnostics diagnostics; + auto* sink = messages && !messages.is_none() ? &diagnostics : nullptr; + auto detachNative = [this, sink]() { + auto cleanupLock = lockForCleanup(); + return detachFetchBindingsNative(sink); + }; + SQLRETURN ret; + try { + if (PyGILState_Check()) { + py::gil_scoped_release release; + ret = detachNative(); + } else { + ret = detachNative(); + } + } catch (...) { + resultMetadata.clear(); + AppendFetchBindingDiagnostics(messages, diagnostics, true); + throw; + } + if (!SQL_SUCCEEDED(ret)) { + resultMetadata.clear(); + } + // No Python objects or callbacks while the native cleanup gate is held. + AppendFetchBindingDiagnostics(messages, diagnostics, !SQL_SUCCEEDED(ret)); + return ret; +} + +void SqlHandle::requireDetachedFetchBindings() { + SQLRETURN ret = detachFetchBindings(); + if (!SQL_SUCCEEDED(ret)) { + ThrowFetchCleanupError(_type, _handle, ret, "Detaching retained fetch buffers"); + } } SQLRETURN SqlHandle::freeHandle() { @@ -1666,11 +1855,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; }; @@ -1697,6 +1888,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; @@ -1707,7 +1904,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"); } } @@ -1747,7 +1944,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(); } @@ -1769,7 +1968,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"); } @@ -1782,7 +1983,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"); } @@ -1807,7 +2010,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"); } @@ -1840,7 +2045,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"); } @@ -1863,7 +2070,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"); } @@ -1886,7 +2095,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"); } @@ -1952,8 +2163,16 @@ ErrorInfo SQLReadError(SQLSMALLINT handleType, SQLHANDLE rawHandle, SQLRETURN re return errorInfo; } +static void AppendDiagRecord(py::handle records, const std::string& state, + const std::string& message) { + py::tuple record = py::make_tuple(py::str(state), py::str(message)); + if (PyList_Append(records.ptr(), record.ptr()) < 0) + throw py::error_already_set(); +} + static void AppendDiagRecords(SQLHANDLE rawHandle, SQLSMALLINT handleType, py::handle records, - bool internalTruncation = false) { + bool internalTruncation = false, + FetchBindingDiagnostics* nativeRecords = nullptr) { // Iterate through all available diagnostic records for (SQLSMALLINT recNumber = 1;; recNumber++) { SQLWCHAR sqlState[6] = {0}; @@ -2003,10 +2222,35 @@ static void AppendDiagRecords(SQLHANDLE rawHandle, SQLSMALLINT handleType, py::h // Format the state string std::string stateWithError = "[" + stateStr + "] (" + std::to_string(nativeError) + ")"; - // Create the tuple with converted strings - py::tuple record = py::make_tuple(py::str(stateWithError), py::str(msgStr)); - if (PyList_Append(records.ptr(), record.ptr()) < 0) - throw py::error_already_set(); + if (nativeRecords) + nativeRecords->emplace_back(std::move(stateWithError), std::move(msgStr)); + else + AppendDiagRecord(records, stateWithError, msgStr); + } +} + +static void CaptureFetchBindingDiagnostics(SQLHSTMT stmt, SQLRETURN ret, + FetchBindingDiagnostics* diagnostics) { + if (diagnostics && (ret == SQL_SUCCESS_WITH_INFO || ret == SQL_NO_DATA)) + AppendDiagRecords(stmt, SQL_HANDLE_STMT, {}, false, diagnostics); +} + +static void AppendFetchBindingDiagnostics(py::handle messages, + const FetchBindingDiagnostics& diagnostics, + bool preserveFailure) { + try { + for (const auto& record : diagnostics) + AppendDiagRecord(messages, record.first, record.second); + } catch (py::error_already_set& error) { + if (!preserveFailure) + throw; + error.discard_as_unraisable("fetch binding diagnostics"); + } catch (const std::exception& error) { + if (!preserveFailure) + throw; + std::fputs("mssql-python: failed to append fetch binding diagnostics: ", stderr); + std::fputs(error.what(), stderr); + std::fputc('\n', stderr); } } @@ -2041,7 +2285,9 @@ static void CheckFetchError(const SqlHandlePtr& handle, SQLRETURN ret) { // 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()); @@ -2078,7 +2324,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(); @@ -2125,7 +2373,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). @@ -2929,7 +3179,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); @@ -3214,7 +3466,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"); } @@ -3236,6 +3490,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"); @@ -3440,6 +3697,10 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p const std::string& wcharEncoding = "utf-16le", int charCtype = SQL_C_WCHAR, py::handle messages = {}) { PERF_TIMER("SQLGetData_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(nullptr, messages); + !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 @@ -4249,9 +4510,16 @@ 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) { + PERF_TIMER("fetch_bindings::SQL_UNBIND"); + SQLFreeStmt_ptr(StatementHandle->get(), SQL_UNBIND); + } // Perform scroll operation SQLRETURN ret = SQL_ERROR; @@ -4279,10 +4547,20 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri template SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames, SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR, - py::handle messages = {}) { + py::handle messages = {}, + 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 (bindings) { + bindings->push_back({column, cType, data, length, indicators}); + return SQL_SUCCESS; + } + PERF_TIMER("fetch_bindings::SQLBindCol"); + return SQLBindCol_ptr(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); @@ -4299,7 +4577,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // 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( + ret = bindColumn( hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), fetchBufferSize * sizeof(SQLWCHAR), buffers.indicators[col - 1].data()); } else { @@ -4310,7 +4588,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; #endif buffers.charBuffers[col - 1].resize(fetchSize * fetchBufferSize); - ret = SQLBindCol_ptr( + ret = bindColumn( hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), fetchBufferSize * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); } @@ -4324,41 +4602,41 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& 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(), + 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(), + 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, + 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(), + 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(), + 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(), + 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(), + ret = bindColumn(hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), MAX_DIGITS_IN_NUMERIC * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; @@ -4366,38 +4644,38 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& case SQL_FLOAT: buffers.doubleBuffers[col - 1].resize(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( + 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); 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); 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); 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(), + ret = bindColumn(hStmt, col, SQL_C_GUID, buffers.guidBuffers[col - 1].data(), sizeof(SQLGUID), buffers.indicators[col - 1].data()); break; case SQL_SS_UDT: @@ -4408,12 +4686,12 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // suffice HandleZeroColumnSizeAtFetch(columnSize); buffers.charBuffers[col - 1].resize(fetchSize * columnSize); - ret = SQLBindCol_ptr(hStmt, col, SQL_C_BINARY, buffers.charBuffers[col - 1].data(), + 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, + ret = bindColumn(hStmt, col, SQL_C_SS_TIMESTAMPOFFSET, buffers.datetimeoffsetBuffers[col - 1].data(), sizeof(DateTimeOffset) * fetchSize, buffers.indicators[col - 1].data()); @@ -4443,12 +4721,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, - py::handle messages = {}) { + py::handle messages = {}, SQLULEN rowCapacity = 0) { PERF_TIMER("FetchBatchData"); LOG("FetchBatchData: Fetching data in batches"); SQLRETURN ret; @@ -4469,6 +4747,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 @@ -4976,13 +5259,24 @@ 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(nullptr, messages); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + plan.reset(); + } SQLHSTMT hStmt = StatementHandle->get(); // Retrieve column count SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle, messages); // Retrieve column metadata - auto snapshot = StatementHandle->resultMetadata.snapshot(); - auto metadata = std::move(snapshot.metadata); + auto metadata = metadataSnapshot.metadata; const bool matches = metadata && numCols >= 0 && metadata->columns.size() == static_cast(numCols); if (!matches || !metadata->namesValidated) { @@ -5008,7 +5302,7 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch } } pending->namesValidated = true; - StatementHandle->resultMetadata.publish(snapshot.generation, pending); + StatementHandle->resultMetadata.publish(metadataSnapshot.generation, pending); metadata = std::move(pending); } ret = SQL_SUCCESS; @@ -5030,6 +5324,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(nullptr, messages); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } LOG("FetchMany_wrap: LOB columns detected (%zu columns), using per-row " "SQLGetData path", lobColumns.size()); @@ -5055,6 +5355,60 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch return SQL_SUCCESS; } + if (StatementHandle->fetchBindings.eligible()) { + const ResultMetadataCache::Snapshot snapshot{metadataSnapshot.generation, metadata}; + if (plan && plan->metadata != metadata) { + ret = StatementHandle->detachFetchBindings(nullptr, messages); + 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, + messages, &plan->bindings); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + StatementHandle->fetchBindings.install(plan); + FetchBindingDiagnostics diagnostics; + try { + ret = plan->attach(hStmt, messages && !messages.is_none() ? &diagnostics : nullptr); + } catch (...) { + AppendFetchBindingDiagnostics(messages, diagnostics, true); + throw; + } + AppendFetchBindingDiagnostics(messages, diagnostics, !SQL_SUCCEEDED(ret)); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } + plan->resetValues(); + ret = FetchBatchData(hStmt, plan->buffers, columnNames, rows, numCols, plan->rowsFetched, + lobColumns, charEncoding, charCtype, messages, fetchSize); + if (ret == SQL_NO_DATA || + (SQL_SUCCEEDED(ret) && + StatementHandle->resultMetadata.snapshot().generation != metadataSnapshot.generation)) { + SQLRETURN detached = StatementHandle->detachFetchBindings(nullptr, messages); + if (!SQL_SUCCEEDED(detached)) { + ret = detached; + } + } + return ret; + } + + if (plan) { + ret = StatementHandle->detachFetchBindings(nullptr, messages); + if (!SQL_SUCCEEDED(ret)) { + return ret; + } + } // Initialize column buffers ColumnBuffers buffers(numCols, fetchSize); FetchStateGuard fetchStateGuard(StatementHandle, messages); @@ -5187,6 +5541,10 @@ int32_t days_from_civil(int y, int m, int d) { SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, int arrowBatchSize, int charCtype, py::handle messages = {}) { PERF_TIMER("FetchArrowBatch_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(nullptr, messages); + !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"); @@ -6119,6 +6477,10 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, const std::string& wcharEncoding = "utf-16le", int charCtype = SQL_C_WCHAR, py::handle messages = {}) { PERF_TIMER("FetchAll_wrap"); + if (SQLRETURN ret = StatementHandle->detachFetchBindings(nullptr, messages); + !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); @@ -6290,13 +6652,20 @@ 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. - ret = SQLFreeStmt_ptr(hStmt, SQL_UNBIND); - CaptureFetchDiagnostics(hStmt, ret, messages); - if (!SQL_SUCCEEDED(ret)) + bool hadFetchPlan; + ret = StatementHandle->detachFetchBindings(&hadFetchPlan, messages); + if (!SQL_SUCCEEDED(ret)) { return ret; + } + if (!hadFetchPlan) { + { + PERF_TIMER("fetch_bindings::SQL_UNBIND"); + ret = SQLFreeStmt_ptr(hStmt, SQL_UNBIND); + } + CaptureFetchDiagnostics(hStmt, ret, messages); + if (!SQL_SUCCEEDED(ret)) + return ret; + } // Assume hStmt is already allocated and a query has been executed { @@ -6323,7 +6692,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 " @@ -6548,8 +6919,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 32d9f8067..ce90f5652 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 @@ -299,6 +300,14 @@ class SqlHandle { SQLSMALLINT type() const; void free(); SQLRETURN freeHandle(); + SQLRETURN detachFetchBindings(bool* hadPlan = nullptr, py::handle messages = {}) { + const bool present = fetchBindings.hasPlan(); + if (hadPlan) { + *hadPlan = present; + } + return present ? detachPresentFetchBindings(messages) : 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* @@ -309,18 +318,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. @@ -331,10 +337,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(py::handle messages); + SQLRETURN detachFetchBindingsNative(FetchBindingDiagnostics* diagnostics = nullptr); SQLSMALLINT _type; SQLHANDLE _handle; bool _implicitly_freed = false; // Tracks if handle was freed by parent @@ -401,51 +410,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..d1068b4f5 --- /dev/null +++ b/mssql_python/pybind/fetch_bindings.hpp @@ -0,0 +1,193 @@ +// 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; +}; + +using FetchBindingDiagnostics = std::vector>; + +// 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; + } + + SQLRETURN attach(SQLHSTMT stmt, FetchBindingDiagnostics* diagnostics = nullptr); + SQLRETURN detach(SQLHSTMT stmt, FetchBindingDiagnostics* diagnostics = nullptr); + + 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/test_025_profiler.py b/tests/test_025_profiler.py index 1db0ebce1..8d094d118 100644 --- a/tests/test_025_profiler.py +++ b/tests/test_025_profiler.py @@ -394,6 +394,207 @@ 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="ascii", ctype=db.SQL_CHAR) + 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="latin-1", ctype=db.SQL_CHAR + ) + 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 +@pytest.mark.skipif( + sys.platform == "win32", reason="Windows does not export the ODBC function-pointer globals" +) +@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 + 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]) + 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() + 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) + 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), 1) + original = pointer.value + assert original + original_call = callback_type(original) + failures = [] + + @callback_type + 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, leaving the driver's pointers unchanged + return original_call(handle, operation, *args) + + try: + pointer.value = ctypes.cast(fail_cleanup, 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), 1) + finally: + pointer.value = original + + assert [tuple(row) for row in cursor.fetchmany(2)] == [(3,), (4,)] + counts((2, 2, 3), 2) + assert cursor.fetchmany(2) == [] + counts((2, 2, 4), 3) + 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__), + failure_point, + ], + 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():