From 6be230ac8dbf505bd89c978a646801c40da693f7 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Mon, 21 Sep 2026 12:17:09 +0530 Subject: [PATCH 01/21] FIX: Validate driver-provided fetch sizes --- mssql_python/pybind/ddbc_bindings.cpp | 371 +++++++++++++++++++------- mssql_python/pybind/ddbc_bindings.h | 6 + 2 files changed, 287 insertions(+), 90 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 1f235f829..14dc809f5 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -21,6 +21,7 @@ #include // For std::memcpy #include #include +#include #include // std::forward #include // CPython datetime API (PyDateTime_IMPORT, PyDateTime_GET_*, etc.) @@ -328,6 +329,91 @@ ParamType* AllocateParamBufferArray(std::vector>& paramBuf return raw; } +size_t CheckedFetchAdd(size_t left, size_t right, const char* errorMessage) { + if (left > std::numeric_limits::max() - right) { + ThrowStdException(errorMessage); + } + return left + right; +} + +size_t CheckedFetchMultiply(size_t left, size_t right, const char* errorMessage) { + if (left != 0 && right > std::numeric_limits::max() / left) { + ThrowStdException(errorMessage); + } + return left * right; +} + +size_t CheckedFetchColumnSize(SQLULEN columnSize) { + if (columnSize > std::numeric_limits::max()) { + ThrowStdException("Column size is too large"); + } + return static_cast(columnSize); +} + +template +void ResizeFetchBuffer(std::vector& buffer, size_t rowCount, size_t stride) { + const size_t count = + CheckedFetchMultiply(rowCount, stride, "Column fetch buffer is too large"); + if (count > buffer.max_size()) { + ThrowStdException("Column fetch buffer is too large"); + } + buffer.resize(count); +} + +SQLLEN CheckedFetchBufferLength(size_t elementCount, size_t elementSize) { + const size_t byteCount = + CheckedFetchMultiply(elementCount, elementSize, "Column fetch stride is too large"); + if (byteCount > static_cast(std::numeric_limits::max())) { + ThrowStdException("Column fetch stride is too large"); + } + return static_cast(byteCount); +} + +void EnsureArrowBufferSize(std::vector& buffer, size_t start, size_t required) { + const size_t targetSize = + CheckedFetchAdd(start, required, "Arrow variable-length buffer is too large"); + if (targetSize > buffer.max_size()) { + ThrowStdException("Arrow variable-length buffer is too large"); + } + size_t newSize = buffer.size(); + while (newSize < targetSize) { + if (newSize == 0) { + newSize = 1; + } else if (newSize > buffer.max_size() / 2) { + newSize = targetSize; + } else { + newSize *= 2; + } + } + if (newSize != buffer.size()) { + buffer.resize(newSize); + } +} + +template +size_t CheckedArrowSourceOffset(const std::vector& buffer, size_t rowIndex, + size_t stride, size_t dataBytes) { + const size_t offset = + CheckedFetchMultiply(rowIndex, stride, "Arrow source offset is too large"); + if (offset > buffer.size() || stride > buffer.size() - offset) { + ThrowStdException("Driver data length exceeds the allocated fetch buffer"); + } + const size_t availableBytes = + CheckedFetchMultiply(stride, sizeof(ElementType), "Arrow source size is too large"); + if (dataBytes > availableBytes) { + ThrowStdException("Driver data length exceeds the allocated fetch buffer"); + } + return offset; +} + +template +std::unique_ptr AllocateFetchArray(size_t count) { + if (count > std::numeric_limits::max() / sizeof(ElementType)) { + ThrowStdException("Arrow value buffer is too large"); + } + return std::make_unique(count); +} + std::string DescribeChar(unsigned char ch) { if (ch >= 32 && ch <= 126) { return std::string("'") + static_cast(ch) + "'"; @@ -3987,6 +4073,10 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& columnNames, SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR) { PERF_TIMER("SQLBindColums"); + if (fetchSize < 0) { + ThrowStdException("Fetch size must be non-negative"); + } + const size_t rowCount = static_cast(fetchSize); SQLRETURN ret = SQL_SUCCESS; const bool useWideChar = (charCtype == SQL_C_WCHAR); // Bind columns based on their data types @@ -4000,25 +4090,33 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column case SQL_VARCHAR: case SQL_LONGVARCHAR: { HandleZeroColumnSizeAtFetch(columnSize); + const size_t baseColumnSize = CheckedFetchColumnSize(columnSize); if (useWideChar) { // Bind VARCHAR columns as SQL_C_WCHAR so the ODBC driver // returns UTF-16 data, avoiding code-page decode issues. - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; - buffers.wcharBuffers[col - 1].resize(fetchSize * fetchBufferSize); + const size_t fetchBufferSize = CheckedFetchAdd( + baseColumnSize, 1, "Column fetch stride is too large"); + ResizeFetchBuffer(buffers.wcharBuffers[col - 1], rowCount, fetchBufferSize); ret = SQLBindCol_ptr( hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), - fetchBufferSize * sizeof(SQLWCHAR), buffers.indicators[col - 1].data()); + CheckedFetchBufferLength(fetchBufferSize, sizeof(SQLWCHAR)), + buffers.indicators[col - 1].data()); } else { // Original narrow-char path #if defined(__APPLE__) || defined(__linux__) - uint64_t fetchBufferSize = columnSize * 4 + 1 /*null-terminator*/; + const size_t fetchBufferSize = CheckedFetchAdd( + CheckedFetchMultiply(baseColumnSize, 4, + "Column fetch stride is too large"), + 1, "Column fetch stride is too large"); #else - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; + const size_t fetchBufferSize = CheckedFetchAdd( + baseColumnSize, 1, "Column fetch stride is too large"); #endif - buffers.charBuffers[col - 1].resize(fetchSize * fetchBufferSize); + ResizeFetchBuffer(buffers.charBuffers[col - 1], rowCount, fetchBufferSize); ret = SQLBindCol_ptr( hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), - fetchBufferSize * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); + CheckedFetchBufferLength(fetchBufferSize, sizeof(SQLCHAR)), + buffers.indicators[col - 1].data()); } break; } @@ -4028,10 +4126,11 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column // TODO: handle variable length data correctly. This logic wont // suffice HandleZeroColumnSizeAtFetch(columnSize); - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; - buffers.wcharBuffers[col - 1].resize(fetchSize * fetchBufferSize); + const size_t fetchBufferSize = CheckedFetchAdd( + CheckedFetchColumnSize(columnSize), 1, "Column fetch stride is too large"); + ResizeFetchBuffer(buffers.wcharBuffers[col - 1], rowCount, fetchBufferSize); ret = SQLBindCol_ptr(hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), - fetchBufferSize * sizeof(SQLWCHAR), + CheckedFetchBufferLength(fetchBufferSize, sizeof(SQLWCHAR)), buffers.indicators[col - 1].data()); break; } @@ -4063,7 +4162,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column break; case SQL_DECIMAL: case SQL_NUMERIC: - buffers.charBuffers[col - 1].resize(fetchSize * MAX_DIGITS_IN_NUMERIC); + ResizeFetchBuffer(buffers.charBuffers[col - 1], rowCount, MAX_DIGITS_IN_NUMERIC); ret = SQLBindCol_ptr(hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), MAX_DIGITS_IN_NUMERIC * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); @@ -4113,15 +4212,17 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column // TODO: handle variable length data correctly. This logic wont // suffice HandleZeroColumnSizeAtFetch(columnSize); - buffers.charBuffers[col - 1].resize(fetchSize * columnSize); + ResizeFetchBuffer(buffers.charBuffers[col - 1], rowCount, + CheckedFetchColumnSize(columnSize)); ret = SQLBindCol_ptr(hStmt, col, SQL_C_BINARY, buffers.charBuffers[col - 1].data(), - columnSize, buffers.indicators[col - 1].data()); + CheckedFetchBufferLength(CheckedFetchColumnSize(columnSize), 1), + 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, buffers.datetimeoffsetBuffers[col - 1].data(), - sizeof(DateTimeOffset) * fetchSize, + sizeof(DateTimeOffset), buffers.indicators[col - 1].data()); break; default: @@ -4166,6 +4267,11 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum LOG("FetchBatchData: No data to fetch"); return ret; } + for (SQLUSMALLINT col = 0; col < numCols; ++col) { + if (numRowsFetched > buffers.indicators[col].size()) { + ThrowStdException("Driver returned more rows than the allocated fetch buffers"); + } + } if (!SQL_SUCCEEDED(ret)) { LOG("FetchBatchData: Error while fetching rows in batches - " "SQLRETURN=%d", @@ -4348,6 +4454,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum PyList_SET_ITEM(row, col - 1, Py_None); continue; } + if (dataLen < 0) { + ThrowStdException("Unexpected negative data length"); + } // Performance: Use function pointer dispatch for simple types (fast // path) This eliminates the switch statement from hot loop - @@ -4373,13 +4482,6 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum Py_INCREF(Py_None); PyList_SET_ITEM(row, col - 1, Py_None); continue; - } else if (dataLen < 0) { - // Negative value is unexpected, log column index, SQL type & - // raise exception - LOG("FetchBatchData: Unexpected negative data length - " - "column=%d, SQL_type=%d, dataLen=%ld", - col, dataType, (long)dataLen); - ThrowStdException("Unexpected negative data length, check logs for details"); } assert(dataLen > 0 && "Data length must be > 0"); @@ -4389,6 +4491,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum case SQL_NUMERIC: { try { SQLLEN decimalDataLen = buffers.indicators[col - 1][i]; + if (decimalDataLen > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } const char* rawData = reinterpret_cast( &buffers.charBuffers[col - 1][i * MAX_DIGITS_IN_NUMERIC]); @@ -4742,35 +4847,54 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, ThrowStdException("GetDataVar only supports SQL_C_CHAR, SQL_C_WCHAR, and SQL_C_BINARY"); } - // Ensure initial buffer has space for at least the null terminator - if (dataVec.size() < sizeNullTerminator) { - dataVec.resize(sizeNullTerminator); + // Binary data has no terminator, but still needs storage to make progress. + const size_t initialSize = std::max(sizeNullTerminator, 1); + if (dataVec.size() < initialSize) { + dataVec.resize(initialSize); } while (true) { + if (start > dataVec.size()) { + ThrowStdException("Invalid variable-length fetch buffer offset"); + } + const size_t availableBytes = CheckedFetchMultiply( + dataVec.size() - start, sizeof(T), "Variable-length fetch buffer is too large"); + if (availableBytes > static_cast(std::numeric_limits::max())) { + ThrowStdException("Variable-length fetch buffer is too large"); + } SQLLEN localInd = 0; SQLRETURN ret = SQLGetData_ptr( hStmt, colNumber, cType, reinterpret_cast(dataVec.data() + start), - sizeof(T) * (dataVec.size() - start), // Available buffer size from start position + static_cast(availableBytes), &localInd); + // Indicator contents are undefined when the ODBC call fails. + if (ret == SQL_ERROR || ret == SQL_INVALID_HANDLE) { + return ret; + } + // Handle NULL data if (localInd == SQL_NULL_DATA) { *indicator = SQL_NULL_DATA; return SQL_SUCCESS; } - // Check for errors (excluding SQL_SUCCESS_WITH_INFO which means more data available) - if (ret == SQL_ERROR || ret == SQL_INVALID_HANDLE) { - return ret; - } - // SQL_SUCCESS or SQL_NO_DATA means we got all the data if (ret == SQL_SUCCESS || ret == SQL_NO_DATA) { if (localInd >= 0) { - *indicator = static_cast(start) * sizeof(T) + localInd; + const size_t prefixBytes = CheckedFetchMultiply( + start, sizeof(T), "Variable-length fetch result is too large"); + if (prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch result is too large"); + } + *indicator = static_cast(prefixBytes) + localInd; } else { - *indicator = localInd; // Preserve SQL_NO_TOTAL or other negative values + if (localInd != SQL_NO_TOTAL) { + ThrowStdException("Unexpected negative variable-length data indicator"); + } + *indicator = localInd; } break; } @@ -4778,19 +4902,40 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // SQL_SUCCESS_WITH_INFO means buffer was too small, need to continue fetching if (ret == SQL_SUCCESS_WITH_INFO) { // Determine how much more space we need - if (localInd < 0) { + if (localInd == SQL_NO_TOTAL) { // SQL_NO_TOTAL: driver doesn't know total size, double the buffer - end = dataVec.size() * 2; + end = CheckedFetchMultiply(dataVec.size(), 2, + "Variable-length fetch buffer is too large"); + } else if (localInd < 0) { + ThrowStdException("Unexpected negative variable-length data indicator"); } else { // Driver returned total size: allocate exactly what we need - assert(localInd % sizeof(T) == 0); - end = start + static_cast(localInd) / sizeof(T) + sizeNullTerminator; + if (localInd % sizeof(T) != 0) { + ThrowStdException("Variable-length data has an invalid byte length"); + } + end = CheckedFetchAdd( + CheckedFetchAdd(start, static_cast(localInd) / sizeof(T), + "Variable-length fetch buffer is too large"), + sizeNullTerminator, "Variable-length fetch buffer is too large"); } // The next read starts where the null terminator would have been placed + if (end <= dataVec.size()) { + const size_t prefixBytes = CheckedFetchMultiply( + start, sizeof(T), "Variable-length fetch result is too large"); + if (localInd < 0 || + prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch made no progress"); + } + *indicator = static_cast(prefixBytes) + localInd; + return SQL_SUCCESS; + } + if (end > dataVec.max_size()) { + ThrowStdException("Variable-length fetch buffer is too large"); + } start = dataVec.size() - sizeNullTerminator; - - // Resize buffer for next iteration dataVec.resize(end); } else { // Unexpected return code @@ -4836,6 +4981,16 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, int arrowBatchSize, int charCtype) { PERF_TIMER("FetchArrowBatch_wrap"); + if (arrowBatchSize < 0) { + ThrowStdException("Arrow batch size must be non-negative"); + } + const size_t batchSize = static_cast(arrowBatchSize); + const size_t offsetCount = + CheckedFetchAdd(batchSize, 1, "Arrow offset buffer is too large"); + const size_t initialVarDataSize = + CheckedFetchMultiply(batchSize, 42, "Arrow variable-length buffer is too large"); + const size_t bitmapSize = + CheckedFetchAdd(batchSize, 7, "Arrow bitmap is too large") / 8; // Fetch narrow char data as SQL_C_CHAR if on Linux/macOS and configured by the user charCtype = EffectiveCharCtypeForFetch(charCtype, "utf-8"); @@ -4879,7 +5034,6 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, SQLSMALLINT nullable = colMeta["Nullable"].cast(); dataTypes[i] = dataType; - columnSizes[i] = columnSize; columnNullable[i] = (nullable != SQL_NO_NULLS); if ((dataType == SQL_WVARCHAR || dataType == SQL_WLONGVARCHAR || dataType == SQL_VARCHAR || @@ -4892,6 +5046,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } } + columnSizes[i] = columnSize; + std::string columnName = colMeta["ColumnName"].cast(); size_t nameLen = columnName.length() + 1; arrowSchemaPrivateData[i]->name = std::make_unique(nameLen); @@ -4908,8 +5064,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_WLONGVARCHAR: case SQL_GUID: format = "U"; - arrowColumnProducer->varVal = std::make_unique(arrowBatchSize + 1); - arrowColumnProducer->varData.resize(arrowBatchSize * 42); + arrowColumnProducer->varVal = AllocateFetchArray(offsetCount); + arrowColumnProducer->varData.resize(initialVarDataSize); columnVarLen[i] = true; // start at offset 0 arrowColumnProducer->varVal[0] = 0; @@ -4920,8 +5076,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_VARBINARY: case SQL_LONGVARBINARY: format = "Z"; - arrowColumnProducer->varVal = std::make_unique(arrowBatchSize + 1); - arrowColumnProducer->varData.resize(arrowBatchSize * 42); + arrowColumnProducer->varVal = AllocateFetchArray(offsetCount); + arrowColumnProducer->varData.resize(initialVarDataSize); columnVarLen[i] = true; // start at offset 0 arrowColumnProducer->varVal[0] = 0; @@ -4929,33 +5085,33 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, break; case SQL_TINYINT: format = "C"; - arrowColumnProducer->uint8Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->uint8Val = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->uint8Val.get(); break; case SQL_SMALLINT: format = "s"; - arrowColumnProducer->int16Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->int16Val = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->int16Val.get(); break; case SQL_INTEGER: format = "i"; - arrowColumnProducer->int32Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->int32Val = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->int32Val.get(); break; case SQL_BIGINT: format = "l"; - arrowColumnProducer->int64Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->int64Val = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->int64Val.get(); break; case SQL_REAL: format = "f"; - arrowColumnProducer->float32Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->float32Val = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->float32Val.get(); break; case SQL_FLOAT: case SQL_DOUBLE: format = "g"; - arrowColumnProducer->float64Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->float64Val = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->float64Val.get(); break; case SQL_DECIMAL: @@ -4968,7 +5124,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, arrowSchemaPrivateData[i]->format = std::make_unique(formatLen); std::memcpy(arrowSchemaPrivateData[i]->format.get(), formatStr.c_str(), formatLen); format = arrowSchemaPrivateData[i]->format.get(); - arrowColumnProducer->decimalVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->decimalVal = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->decimalVal.get(); break; } @@ -4976,28 +5132,28 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_TYPE_TIMESTAMP: case SQL_DATETIME: format = "tsu:"; - arrowColumnProducer->tsMicroVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->tsMicroVal = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->tsMicroVal.get(); break; case SQL_SS_TIMESTAMPOFFSET: format = "tsu:+00:00"; - arrowColumnProducer->tsMicroVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->tsMicroVal = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->tsMicroVal.get(); break; case SQL_TYPE_DATE: format = "tdD"; - arrowColumnProducer->dateVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->dateVal = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->dateVal.get(); break; case SQL_SS_TIME2: format = "ttn"; - arrowColumnProducer->timeNanoVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->timeNanoVal = AllocateFetchArray(batchSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->timeNanoVal.get(); break; case SQL_BIT: format = "b"; - arrowColumnProducer->bitVal = std::make_unique((arrowBatchSize + 7) / 8); - std::memset(arrowColumnProducer->bitVal.get(), 0, (arrowBatchSize + 7) / 8); + arrowColumnProducer->bitVal = AllocateFetchArray(bitmapSize); + std::memset(arrowColumnProducer->bitVal.get(), 0, bitmapSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->bitVal.get(); break; default: @@ -5018,9 +5174,9 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, std::memcpy(arrowSchemaPrivateData[i]->format.get(), format.c_str(), formatLen); } - arrowColumnProducer->valid = std::make_unique((arrowBatchSize + 7) / 8); + arrowColumnProducer->valid = AllocateFetchArray(bitmapSize); // Initialize validity bitmap to all valid - std::memset(arrowColumnProducer->valid.get(), 0xFF, (arrowBatchSize + 7) / 8); + std::memset(arrowColumnProducer->valid.get(), 0xFF, bitmapSize); } // Initialize column buffers @@ -5041,9 +5197,11 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, while (idxRowArrow < arrowBatchSize) { int spaceLeftInArrowBatch = arrowBatchSize - idxRowArrow; + int currentFetchSize = fetchSize; if (fetchSize > spaceLeftInArrowBatch) { // Adjust fetch size for final batch to avoid overfetching - fetchStateGuard.setRowArraySize(spaceLeftInArrowBatch); + currentFetchSize = spaceLeftInArrowBatch; + fetchStateGuard.setRowArraySize(currentFetchSize); } { // Release GIL during the blocking ODBC fetch @@ -5060,7 +5218,10 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } // numRowsFetched is the SQL_ATTR_ROWS_FETCHED_PTR attribute. // It'll be populated by SQLFetch - assert(numRowsFetched + idxRowArrow <= static_cast(arrowBatchSize)); + if (numRowsFetched > static_cast(currentFetchSize) || + numRowsFetched > static_cast(spaceLeftInArrowBatch)) { + ThrowStdException("Driver returned more rows than the allocated Arrow buffers"); + } for (SQLULEN idxRowSql = 0; idxRowSql < numRowsFetched; idxRowSql++) { for (SQLUSMALLINT idxCol = 0; idxCol < numCols; idxCol++) { auto& arrowColumnProducer = arrowArrayPrivateData[idxCol]; @@ -5326,16 +5487,20 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_BINARY: case SQL_VARBINARY: case SQL_LONGVARBINARY: { - uint64_t fetchBufferSize = columnSize /* bytes are not null terminated */; + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchColumnSize( + processedColumnSize); auto target_vec = &arrowColumnProducer->varData; - auto start = arrowColumnProducer->varVal[idxRowArrow]; - while (target_vec->size() < start + dataLen) { - target_vec->resize(target_vec->size() * 2); - } + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); + EnsureArrowBufferSize(*target_vec, start, dataLen); std::memcpy(&(*target_vec)[start], - &buffers.charBuffers[idxCol][idxRowSql * fetchBufferSize], - dataLen); + &buffers.charBuffers[idxCol][sourceOffset], dataLen); arrowColumnProducer->varVal[idxRowArrow + 1] = start + dataLen; break; } @@ -5343,20 +5508,36 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_VARCHAR: case SQL_LONGVARCHAR: { if (charCtype == SQL_C_CHAR) { + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); #if defined(__APPLE__) || defined(__linux__) - uint64_t fetchBufferSize = columnSize * 4 + 1 /*null-terminator*/; + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchAdd( + CheckedFetchMultiply( + CheckedFetchColumnSize( + processedColumnSize), + 4, + "Column fetch stride is too large"), + 1, + "Column fetch stride is too large"); #else - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchAdd( + CheckedFetchColumnSize( + processedColumnSize), + 1, + "Column fetch stride is too large"); #endif auto target_vec = &arrowColumnProducer->varData; - auto start = arrowColumnProducer->varVal[idxRowArrow]; - while (target_vec->size() < start + dataLen) { - target_vec->resize(target_vec->size() * 2); - } + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); + EnsureArrowBufferSize(*target_vec, start, dataLen); std::memcpy(&(*target_vec)[start], - &buffers.charBuffers[idxCol][idxRowSql * fetchBufferSize], - dataLen); + &buffers.charBuffers[idxCol][sourceOffset], dataLen); arrowColumnProducer->varVal[idxRowArrow + 1] = start + dataLen; break; } @@ -5367,20 +5548,30 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_WVARCHAR: case SQL_WLONGVARCHAR: { // We have previously fetched these as WCHARs, even for SQL_CHAR types. - assert(dataLen % sizeof(SQLWCHAR) == 0); - auto dataLenW = dataLen / sizeof(SQLWCHAR); - auto wcharSource = - &buffers.wcharBuffers[idxCol][idxRowSql * (columnSize + 1)]; - auto start = arrowColumnProducer->varVal[idxRowArrow]; + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } + const size_t dataLenW = dataLen / sizeof(SQLWCHAR); + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); + const size_t fetchBufferSize = hasLobColumns + ? buffers.wcharBuffers[idxCol].size() + : CheckedFetchAdd( + CheckedFetchColumnSize( + processedColumnSize), + 1, + "Column fetch stride is too large"); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.wcharBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); + auto wcharSource = &buffers.wcharBuffers[idxCol][sourceOffset]; + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; auto target_vec = &arrowColumnProducer->varData; static_assert(sizeof(SQLWCHAR) == sizeof(char16_t)); static_assert(alignof(SQLWCHAR) == alignof(char16_t)); const auto* utf16Source = reinterpret_cast(wcharSource); - size_t maxUtf8Size = dataLenW * 3; - - while (target_vec->size() < start + maxUtf8Size) { - target_vec->resize(target_vec->size() * 2); - } + const size_t maxUtf8Size = CheckedFetchMultiply( + dataLenW, 3, "Arrow UTF-8 buffer is too large"); + EnsureArrowBufferSize(*target_vec, start, maxUtf8Size); size_t bytesWritten = simdutf::convert_utf16le_to_utf8_with_replacement( utf16Source, dataLenW, @@ -5394,12 +5585,10 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, // "550e8400-e29b-41d4-a716-446655440000") Each GUID is exactly 36 bytes in // UTF-8 auto target_vec = &arrowColumnProducer->varData; - auto start = arrowColumnProducer->varVal[idxRowArrow]; + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; // Ensure buffer has space for the GUID string + null terminator - while (target_vec->size() < start + 37) { - target_vec->resize(target_vec->size() * 2); - } + EnsureArrowBufferSize(*target_vec, start, 37); // Get the GUID from the buffer const SQLGUID& guidValue = buffers.guidBuffers[idxCol][idxRowSql]; @@ -5444,7 +5633,9 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_DECIMAL: case SQL_NUMERIC: { // Relies on overloaded operators defined in Int128_t struct - assert(dataLen <= MAX_DIGITS_IN_NUMERIC); + if (dataLen > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } Int128_t decimalValue(0, 0); auto start = idxRowSql * MAX_DIGITS_IN_NUMERIC; int sign = 1; diff --git a/mssql_python/pybind/ddbc_bindings.h b/mssql_python/pybind/ddbc_bindings.h index 11c33d8d2..f77583ceb 100644 --- a/mssql_python/pybind/ddbc_bindings.h +++ b/mssql_python/pybind/ddbc_bindings.h @@ -600,6 +600,9 @@ inline void ProcessChar(PyObject* row, ColumnBuffers& buffers, const void* colIn } if (colInfo->useWideChar) { + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } // Wide-char path: data was bound as SQL_C_WCHAR, lives in wcharBuffers uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR); if (!colInfo->isLob && numCharsInData < colInfo->fetchBufferSize) { @@ -712,6 +715,9 @@ inline void ProcessWChar(PyObject* row, ColumnBuffers& buffers, const void* colI } uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR); + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } // Fast path: Data fits in buffer (not LOB or truncated) // fetchBufferSize includes null-terminator, numCharsInData doesn't. Hence // '<' From d085633591c5de8d3595e52579e0b9e3787772a5 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Mon, 21 Sep 2026 12:22:41 +0530 Subject: [PATCH 02/21] FIX: Unify checked size helpers --- mssql_python/pybind/ddbc_bindings.cpp | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 42e01c49a..61d78bde3 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -2356,9 +2356,9 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& LOG("BindParameterArray: Binding SQL_C_WCHAR array - " "param_index=%d, count=%zu, column_size=%zu", paramIndex, paramSetSize, info.columnSize); - const size_t elementWidth = CheckedAddSize( + const size_t elementWidth = CheckedFetchAdd( info.columnSize, 1, "Wide-character parameter size is too large"); - const size_t bufferBytes = CheckedMultiplySize( + const size_t bufferBytes = CheckedFetchMultiply( elementWidth, sizeof(SQLWCHAR), "Wide-character parameter length is too large"); if (bufferBytes > static_cast(std::numeric_limits::max())) { @@ -2366,7 +2366,7 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& } SQLWCHAR* wcharArray = AllocateParamBufferArray( tempBuffers, - CheckedMultiplySize(paramSetSize, elementWidth, + CheckedFetchMultiply(paramSetSize, elementWidth, "Wide-character parameter buffer is too large")); strLenOrIndArray = AllocateParamBufferArray(tempBuffers, paramSetSize); for (size_t i = 0; i < paramSetSize; ++i) { @@ -2463,14 +2463,14 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& LOG("BindParameterArray: Binding SQL_C_CHAR/BINARY array - " "param_index=%d, count=%zu, column_size=%zu, encoding='%s'", paramIndex, paramSetSize, info.columnSize, charEncoding.c_str()); - const size_t elementWidth = CheckedAddSize( + const size_t elementWidth = CheckedFetchAdd( info.columnSize, 1, "Character parameter size is too large"); if (elementWidth > static_cast(std::numeric_limits::max())) { ThrowStdException("Character parameter length is too large"); } char* charArray = AllocateParamBufferArray( tempBuffers, - CheckedMultiplySize(paramSetSize, elementWidth, + CheckedFetchMultiply(paramSetSize, elementWidth, "Character parameter buffer is too large")); strLenOrIndArray = AllocateParamBufferArray(tempBuffers, paramSetSize); for (size_t i = 0; i < paramSetSize; ++i) { From b2c89b89a4531ab9521dce3e0b774a2c54053ba3 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Mon, 21 Sep 2026 12:17:09 +0530 Subject: [PATCH 03/21] FIX: Validate driver-provided fetch sizes --- mssql_python/pybind/ddbc_bindings.cpp | 165 +++++++++++++++++++------- 1 file changed, 125 insertions(+), 40 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index fd7f5ca41..a5399ef51 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -4118,6 +4118,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column case SQL_VARCHAR: case SQL_LONGVARCHAR: { HandleZeroColumnSizeAtFetch(columnSize); + const size_t baseColumnSize = CheckedFetchColumnSize(columnSize); if (useWideChar) { // Bind VARCHAR columns as SQL_C_WCHAR so the ODBC driver // returns UTF-16 data, avoiding code-page decode issues. @@ -4128,13 +4129,18 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column reservedBytes); ret = SQLBindCol_ptr( hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), - fetchBufferSize * sizeof(SQLWCHAR), buffers.indicators[col - 1].data()); + CheckedFetchBufferLength(fetchBufferSize, sizeof(SQLWCHAR)), + buffers.indicators[col - 1].data()); } else { // Original narrow-char path #if defined(__APPLE__) || defined(__linux__) - uint64_t fetchBufferSize = columnSize * 4 + 1 /*null-terminator*/; + const size_t fetchBufferSize = CheckedFetchAdd( + CheckedFetchMultiply(baseColumnSize, 4, + "Column fetch stride is too large"), + 1, "Column fetch stride is too large"); #else - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; + const size_t fetchBufferSize = CheckedFetchAdd( + baseColumnSize, 1, "Column fetch stride is too large"); #endif ResizeNativeFetchBuffer(buffers.charBuffers[col - 1], CheckedMultiplySize(fetchSize, fetchBufferSize, @@ -4142,7 +4148,8 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column reservedBytes); ret = SQLBindCol_ptr( hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), - fetchBufferSize * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); + CheckedFetchBufferLength(fetchBufferSize, sizeof(SQLCHAR)), + buffers.indicators[col - 1].data()); } break; } @@ -4158,7 +4165,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column "Native fetch buffer is too large"), reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_WCHAR, buffers.wcharBuffers[col - 1].data(), - fetchBufferSize * sizeof(SQLWCHAR), + CheckedFetchBufferLength(fetchBufferSize, sizeof(SQLWCHAR)), buffers.indicators[col - 1].data()); break; } @@ -4250,14 +4257,15 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column "Native fetch buffer is too large"), reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_BINARY, buffers.charBuffers[col - 1].data(), - columnSize, buffers.indicators[col - 1].data()); + CheckedFetchBufferLength(CheckedFetchColumnSize(columnSize), 1), + buffers.indicators[col - 1].data()); break; case SQL_SS_TIMESTAMPOFFSET: ResizeNativeFetchBuffer(buffers.datetimeoffsetBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_SS_TIMESTAMPOFFSET, buffers.datetimeoffsetBuffers[col - 1].data(), - sizeof(DateTimeOffset) * fetchSize, + sizeof(DateTimeOffset), buffers.indicators[col - 1].data()); break; default: @@ -4303,6 +4311,11 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum LOG("FetchBatchData: No data to fetch"); return ret; } + for (SQLUSMALLINT col = 0; col < numCols; ++col) { + if (numRowsFetched > buffers.indicators[col].size()) { + ThrowStdException("Driver returned more rows than the allocated fetch buffers"); + } + } if (!SQL_SUCCEEDED(ret)) { LOG("FetchBatchData: Error while fetching rows in batches - " "SQLRETURN=%d", @@ -4485,6 +4498,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum PyList_SET_ITEM(row, col - 1, Py_None); continue; } + if (dataLen < 0) { + ThrowStdException("Unexpected negative data length"); + } // Performance: Use function pointer dispatch for simple types (fast // path) This eliminates the switch statement from hot loop - @@ -4510,13 +4526,6 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum Py_INCREF(Py_None); PyList_SET_ITEM(row, col - 1, Py_None); continue; - } else if (dataLen < 0) { - // Negative value is unexpected, log column index, SQL type & - // raise exception - LOG("FetchBatchData: Unexpected negative data length - " - "column=%d, SQL_type=%d, dataLen=%ld", - col, dataType, (long)dataLen); - ThrowStdException("Unexpected negative data length, check logs for details"); } assert(dataLen > 0 && "Data length must be > 0"); @@ -4526,6 +4535,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum case SQL_NUMERIC: { try { SQLLEN decimalDataLen = buffers.indicators[col - 1][i]; + if (decimalDataLen > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } const char* rawData = reinterpret_cast( &buffers.charBuffers[col - 1][i * MAX_DIGITS_IN_NUMERIC]); @@ -4892,29 +4904,47 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, } while (true) { + if (start > dataVec.size()) { + ThrowStdException("Invalid variable-length fetch buffer offset"); + } + const size_t availableBytes = CheckedFetchMultiply( + dataVec.size() - start, sizeof(T), "Variable-length fetch buffer is too large"); + if (availableBytes > static_cast(std::numeric_limits::max())) { + ThrowStdException("Variable-length fetch buffer is too large"); + } SQLLEN localInd = 0; SQLRETURN ret = SQLGetData_ptr( hStmt, colNumber, cType, reinterpret_cast(dataVec.data() + start), - sizeof(T) * (dataVec.size() - start), // Available buffer size from start position + static_cast(availableBytes), &localInd); + // Indicator contents are undefined when the ODBC call fails. + if (ret == SQL_ERROR || ret == SQL_INVALID_HANDLE) { + return ret; + } + // Handle NULL data if (localInd == SQL_NULL_DATA) { *indicator = SQL_NULL_DATA; return SQL_SUCCESS; } - // Check for errors (excluding SQL_SUCCESS_WITH_INFO which means more data available) - if (ret == SQL_ERROR || ret == SQL_INVALID_HANDLE) { - return ret; - } - // SQL_SUCCESS or SQL_NO_DATA means we got all the data if (ret == SQL_SUCCESS || ret == SQL_NO_DATA) { if (localInd >= 0) { - *indicator = static_cast(start) * sizeof(T) + localInd; + const size_t prefixBytes = CheckedFetchMultiply( + start, sizeof(T), "Variable-length fetch result is too large"); + if (prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch result is too large"); + } + *indicator = static_cast(prefixBytes) + localInd; } else { - *indicator = localInd; // Preserve SQL_NO_TOTAL or other negative values + if (localInd != SQL_NO_TOTAL) { + ThrowStdException("Unexpected negative variable-length data indicator"); + } + *indicator = localInd; } break; } @@ -4922,7 +4952,7 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // SQL_SUCCESS_WITH_INFO means buffer was too small, need to continue fetching if (ret == SQL_SUCCESS_WITH_INFO) { // Determine how much more space we need - if (localInd < 0) { + if (localInd == SQL_NO_TOTAL) { // SQL_NO_TOTAL: driver doesn't know total size, double the buffer end = CheckedMultiplySize(dataVec.size(), 2, "Native fetch buffer size is too large"); @@ -4936,6 +4966,21 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, } // The next read starts where the null terminator would have been placed + if (end <= dataVec.size()) { + const size_t prefixBytes = CheckedFetchMultiply( + start, sizeof(T), "Variable-length fetch result is too large"); + if (localInd < 0 || + prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch made no progress"); + } + *indicator = static_cast(prefixBytes) + localInd; + return SQL_SUCCESS; + } + if (end > dataVec.max_size()) { + ThrowStdException("Variable-length fetch buffer is too large"); + } start = dataVec.size() - sizeNullTerminator; // Resize buffer for next iteration @@ -5036,7 +5081,6 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, SQLSMALLINT nullable = colMeta["Nullable"].cast(); dataTypes[i] = dataType; - columnSizes[i] = columnSize; columnNullable[i] = (nullable != SQL_NO_NULLS); if ((dataType == SQL_WVARCHAR || dataType == SQL_WLONGVARCHAR || dataType == SQL_VARCHAR || @@ -5049,6 +5093,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } } + columnSizes[i] = columnSize; + std::string columnName = colMeta["ColumnName"].cast(); size_t nameLen = columnName.length() + 1; arrowSchemaPrivateData[i]->name = std::make_unique(nameLen); @@ -5235,9 +5281,11 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, while (idxRowArrow < arrowBatchSize) { int spaceLeftInArrowBatch = arrowBatchSize - idxRowArrow; + int currentFetchSize = fetchSize; if (fetchSize > spaceLeftInArrowBatch) { // Adjust fetch size for final batch to avoid overfetching - fetchStateGuard.setRowArraySize(spaceLeftInArrowBatch); + currentFetchSize = spaceLeftInArrowBatch; + fetchStateGuard.setRowArraySize(currentFetchSize); } { // Release GIL during the blocking ODBC fetch @@ -5254,7 +5302,10 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } // numRowsFetched is the SQL_ATTR_ROWS_FETCHED_PTR attribute. // It'll be populated by SQLFetch - assert(numRowsFetched + idxRowArrow <= static_cast(arrowBatchSize)); + if (numRowsFetched > static_cast(currentFetchSize) || + numRowsFetched > static_cast(spaceLeftInArrowBatch)) { + ThrowStdException("Driver returned more rows than the allocated Arrow buffers"); + } for (SQLULEN idxRowSql = 0; idxRowSql < numRowsFetched; idxRowSql++) { for (SQLUSMALLINT idxCol = 0; idxCol < numCols; idxCol++) { auto& arrowColumnProducer = arrowArrayPrivateData[idxCol]; @@ -5520,7 +5571,12 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_BINARY: case SQL_VARBINARY: case SQL_LONGVARBINARY: { - uint64_t fetchBufferSize = columnSize /* bytes are not null terminated */; + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchColumnSize( + processedColumnSize); auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; EnsureNativeFetchBufferSize( @@ -5529,8 +5585,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, reservedBytes); std::memcpy(&(*target_vec)[start], - &buffers.charBuffers[idxCol][idxRowSql * fetchBufferSize], - dataLen); + &buffers.charBuffers[idxCol][sourceOffset], dataLen); arrowColumnProducer->varVal[idxRowArrow + 1] = start + dataLen; break; } @@ -5538,10 +5593,27 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_VARCHAR: case SQL_LONGVARCHAR: { if (charCtype == SQL_C_CHAR) { + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); #if defined(__APPLE__) || defined(__linux__) - uint64_t fetchBufferSize = columnSize * 4 + 1 /*null-terminator*/; + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchAdd( + CheckedFetchMultiply( + CheckedFetchColumnSize( + processedColumnSize), + 4, + "Column fetch stride is too large"), + 1, + "Column fetch stride is too large"); #else - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchAdd( + CheckedFetchColumnSize( + processedColumnSize), + 1, + "Column fetch stride is too large"); #endif auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; @@ -5551,8 +5623,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, reservedBytes); std::memcpy(&(*target_vec)[start], - &buffers.charBuffers[idxCol][idxRowSql * fetchBufferSize], - dataLen); + &buffers.charBuffers[idxCol][sourceOffset], dataLen); arrowColumnProducer->varVal[idxRowArrow + 1] = start + dataLen; break; } @@ -5563,11 +5634,23 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_WVARCHAR: case SQL_WLONGVARCHAR: { // We have previously fetched these as WCHARs, even for SQL_CHAR types. - assert(dataLen % sizeof(SQLWCHAR) == 0); - auto dataLenW = dataLen / sizeof(SQLWCHAR); - auto wcharSource = - &buffers.wcharBuffers[idxCol][idxRowSql * (columnSize + 1)]; - auto start = arrowColumnProducer->varVal[idxRowArrow]; + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } + const size_t dataLenW = dataLen / sizeof(SQLWCHAR); + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); + const size_t fetchBufferSize = hasLobColumns + ? buffers.wcharBuffers[idxCol].size() + : CheckedFetchAdd( + CheckedFetchColumnSize( + processedColumnSize), + 1, + "Column fetch stride is too large"); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.wcharBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); + auto wcharSource = &buffers.wcharBuffers[idxCol][sourceOffset]; + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; auto target_vec = &arrowColumnProducer->varData; static_assert(sizeof(SQLWCHAR) == sizeof(char16_t)); static_assert(alignof(SQLWCHAR) == alignof(char16_t)); @@ -5592,7 +5675,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, // "550e8400-e29b-41d4-a716-446655440000") Each GUID is exactly 36 bytes in // UTF-8 auto target_vec = &arrowColumnProducer->varData; - auto start = arrowColumnProducer->varVal[idxRowArrow]; + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; // Ensure buffer has space for the GUID string + null terminator EnsureNativeFetchBufferSize( @@ -5643,7 +5726,9 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_DECIMAL: case SQL_NUMERIC: { // Relies on overloaded operators defined in Int128_t struct - assert(dataLen <= MAX_DIGITS_IN_NUMERIC); + if (dataLen > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } Int128_t decimalValue(0, 0); auto start = idxRowSql * MAX_DIGITS_IN_NUMERIC; int sign = 1; From d40dc895643453fd0c198fa263b0626a817c132b Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Mon, 21 Sep 2026 15:27:24 +0530 Subject: [PATCH 04/21] FIX: Preserve stacked fetch validation --- mssql_python/pybind/ddbc_bindings.cpp | 136 ++++++++++++++++---------- 1 file changed, 86 insertions(+), 50 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index a5399ef51..76017fd00 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -346,6 +346,46 @@ size_t CheckedMultiplySize(size_t left, size_t right, const char* errorMessage) return left * right; } +size_t CheckedFetchAdd(size_t left, size_t right, const char* errorMessage) { + return CheckedAddSize(left, right, errorMessage); +} + +size_t CheckedFetchMultiply(size_t left, size_t right, const char* errorMessage) { + return CheckedMultiplySize(left, right, errorMessage); +} + +size_t CheckedFetchColumnSize(SQLULEN columnSize) { + if (columnSize > std::numeric_limits::max()) { + ThrowStdException("Column size is too large"); + } + return static_cast(columnSize); +} + +SQLLEN CheckedFetchBufferLength(size_t elementCount, size_t elementSize) { + const size_t byteCount = + CheckedFetchMultiply(elementCount, elementSize, "Column fetch stride is too large"); + if (byteCount > static_cast(std::numeric_limits::max())) { + ThrowStdException("Column fetch stride is too large"); + } + return static_cast(byteCount); +} + +template +size_t CheckedArrowSourceOffset(const std::vector& buffer, size_t rowIndex, + size_t stride, size_t dataBytes) { + const size_t offset = + CheckedFetchMultiply(rowIndex, stride, "Arrow source offset is too large"); + if (offset > buffer.size() || stride > buffer.size() - offset) { + ThrowStdException("Driver data length exceeds the allocated fetch buffer"); + } + const size_t availableBytes = + CheckedFetchMultiply(stride, sizeof(ElementType), "Arrow source size is too large"); + if (dataBytes > availableBytes) { + ThrowStdException("Driver data length exceeds the allocated fetch buffer"); + } + return offset; +} + constexpr int MAX_NATIVE_ROW_COUNT = 1000000; constexpr size_t MAX_NATIVE_FETCH_BYTES = 256ULL * 1024 * 1024; @@ -4122,7 +4162,8 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column if (useWideChar) { // Bind VARCHAR columns as SQL_C_WCHAR so the ODBC driver // returns UTF-16 data, avoiding code-page decode issues. - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; + const size_t fetchBufferSize = CheckedFetchAdd( + baseColumnSize, 1, "Column fetch stride is too large"); ResizeNativeFetchBuffer(buffers.wcharBuffers[col - 1], CheckedMultiplySize(fetchSize, fetchBufferSize, "Native fetch buffer is too large"), @@ -4159,7 +4200,8 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& column // TODO: handle variable length data correctly. This logic wont // suffice HandleZeroColumnSizeAtFetch(columnSize); - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; + const size_t fetchBufferSize = CheckedFetchAdd( + CheckedFetchColumnSize(columnSize), 1, "Column fetch stride is too large"); ResizeNativeFetchBuffer(buffers.wcharBuffers[col - 1], CheckedMultiplySize(fetchSize, fetchBufferSize, "Native fetch buffer is too large"), @@ -4311,17 +4353,17 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, py::list& colum LOG("FetchBatchData: No data to fetch"); return ret; } - for (SQLUSMALLINT col = 0; col < numCols; ++col) { - if (numRowsFetched > buffers.indicators[col].size()) { - ThrowStdException("Driver returned more rows than the allocated fetch buffers"); - } - } if (!SQL_SUCCEEDED(ret)) { LOG("FetchBatchData: Error while fetching rows in batches - " "SQLRETURN=%d", ret); return ret; } + for (SQLUSMALLINT col = 0; col < numCols; ++col) { + if (numRowsFetched > buffers.indicators[col].size()) { + ThrowStdException("Driver returned more rows than the allocated fetch buffers"); + } + } // 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 @@ -4753,6 +4795,26 @@ size_t calculateRowSize(py::list& columnNames, SQLUSMALLINT numCols) { return rowSize; } +struct FetchStateGuard { + SQLHSTMT hStmt; + + FetchStateGuard(SQLHSTMT stmtHandle, SQLULEN* numRowsFetched, SQLULEN rowArraySize) + : hStmt(stmtHandle) { + SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)(intptr_t)rowArraySize, 0); + SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, numRowsFetched, 0); + } + + ~FetchStateGuard() { + SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)1, 0); + SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, NULL, 0); + SQLFreeStmt_ptr(hStmt, SQL_UNBIND); + } + + void setRowArraySize(SQLULEN rowArraySize) const { + SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)(intptr_t)rowArraySize, 0); + } +}; + // FetchMany_wrap - Fetches multiple rows of data from the result set. // // @param StatementHandle: Handle to the statement from which data is to be @@ -4846,8 +4908,7 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch return ret; } - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)(intptr_t)fetchSize, 0); - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, &numRowsFetched, 0); + FetchStateGuard fetchStateGuard(hStmt, &numRowsFetched, fetchSize); ret = FetchBatchData(hStmt, buffers, columnNames, rows, numCols, numRowsFetched, lobColumns, charEncoding, charCtype); @@ -4856,13 +4917,6 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch return ret; } - // Reset attributes before returning to avoid using stack pointers later - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)1, 0); - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, NULL, 0); - - // Unbind columns to allow subsequent fetchone() calls to use SQLGetData - SQLFreeStmt_ptr(hStmt, SQL_UNBIND); - return ret; } @@ -4898,9 +4952,10 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, ThrowStdException("GetDataVar only supports SQL_C_CHAR, SQL_C_WCHAR, and SQL_C_BINARY"); } - // Ensure initial buffer has space for at least the null terminator - if (dataVec.size() < sizeNullTerminator) { - ResizeNativeFetchBuffer(dataVec, sizeNullTerminator, reservedBytes); + // Binary data has no terminator, but SQL_NO_TOTAL still needs room to make progress. + const size_t initialSize = std::max(sizeNullTerminator, 1); + if (dataVec.size() < initialSize) { + ResizeNativeFetchBuffer(dataVec, initialSize, reservedBytes); } while (true) { @@ -4957,8 +5012,13 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, end = CheckedMultiplySize(dataVec.size(), 2, "Native fetch buffer size is too large"); } else { + if (localInd < 0) { + ThrowStdException("Unexpected negative variable-length data indicator"); + } + if (localInd % sizeof(T) != 0) { + ThrowStdException("Variable-length data has an invalid byte length"); + } // Driver returned total size: allocate exactly what we need - assert(localInd % sizeof(T) == 0); end = CheckedAddSize( CheckedAddSize(start, static_cast(localInd) / sizeof(T), "Native fetch buffer size is too large"), @@ -4995,26 +5055,6 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, return SQL_SUCCESS; } -struct FetchStateGuard { - SQLHSTMT hStmt; - - FetchStateGuard(SQLHSTMT stmtHandle, SQLULEN* numRowsFetched, SQLULEN rowArraySize) - : hStmt(stmtHandle) { - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)(intptr_t)rowArraySize, 0); - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, numRowsFetched, 0); - } - - ~FetchStateGuard() { - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)1, 0); - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, NULL, 0); - SQLFreeStmt_ptr(hStmt, SQL_UNBIND); - } - - void setRowArraySize(SQLULEN rowArraySize) const { - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)(intptr_t)rowArraySize, 0); - } -}; - int32_t days_from_civil(int y, int m, int d) { // Implements the "days_from_civil" algorithm by Howard Hinnant // Returns number of days since Unix epoch (1970-01-01) @@ -5583,6 +5623,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, *target_vec, CheckedAddSize(start, dataLen, "Arrow value buffer is too large"), reservedBytes); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -5621,6 +5663,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, *target_vec, CheckedAddSize(start, dataLen, "Arrow value buffer is too large"), reservedBytes); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -6140,9 +6184,8 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, return ret; } - SQLULEN numRowsFetched; - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)(intptr_t)fetchSize, 0); - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, &numRowsFetched, 0); + SQLULEN numRowsFetched = 0; + FetchStateGuard fetchStateGuard(hStmt, &numRowsFetched, fetchSize); while (ret != SQL_NO_DATA) { ret = FetchBatchData(hStmt, buffers, columnNames, rows, numCols, numRowsFetched, lobColumns, @@ -6153,13 +6196,6 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, } } - // Reset attributes before returning to avoid using stack pointers later - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROW_ARRAY_SIZE, (SQLPOINTER)1, 0); - SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_ROWS_FETCHED_PTR, NULL, 0); - - // Unbind columns to allow subsequent fetchone() calls to use SQLGetData - SQLFreeStmt_ptr(hStmt, SQL_UNBIND); - return ret; } From d9ebebb97ec5769232e6d68cd2905d16f0e9bba9 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 11:41:01 +0530 Subject: [PATCH 05/21] FIX: Validate driver-provided fetch sizes Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 264 ++++++++++++++++++++------ mssql_python/pybind/ddbc_bindings.h | 6 + tests/test_010_pybind_functions.py | 37 ++++ 3 files changed, 246 insertions(+), 61 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index fc4aa8658..ab3972874 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -366,6 +366,22 @@ SQLLEN CheckedFetchBufferLength(size_t elementCount, size_t elementSize) { return static_cast(byteCount); } +template +size_t CheckedArrowSourceOffset(const std::vector& buffer, size_t rowIndex, + size_t stride, size_t dataBytes) { + const size_t offset = + CheckedMultiplySize(rowIndex, stride, "Arrow source offset is too large"); + if (offset > buffer.size() || stride > buffer.size() - offset) { + ThrowStdException("Driver data length exceeds the allocated fetch buffer"); + } + const size_t availableBytes = + CheckedMultiplySize(stride, sizeof(ElementType), "Arrow source size is too large"); + if (dataBytes > availableBytes) { + ThrowStdException("Driver data length exceeds the allocated fetch buffer"); + } + return offset; +} + constexpr int MAX_NATIVE_ROW_COUNT = 1000000; constexpr size_t MAX_NATIVE_FETCH_BYTES = 256ULL * 1024 * 1024; constexpr size_t MAX_NATIVE_PARAMETER_BYTES = 256ULL * 1024 * 1024; @@ -438,20 +454,6 @@ size_t ParameterArrayElementSize(const ParamInfo& info) { return 0; } -template -size_t CheckedArrowSourceOffset(const std::vector& buffer, size_t rowIndex, - size_t rowStride, size_t dataBytes) { - const size_t offset = - CheckedMultiplySize(rowIndex, rowStride, "Arrow source offset is too large"); - const size_t rowCapacity = CheckedMultiplySize( - rowStride, sizeof(ElementType), "Arrow source capacity is too large"); - if (offset > buffer.size() || rowStride > buffer.size() - offset || - dataBytes > rowCapacity) { - ThrowStdException("Driver data length exceeds the allocated fetch buffer"); - } - return offset; -} - template std::unique_ptr AllocateUniqueArray(size_t count, const char* errorMessage) { if (count > std::numeric_limits::max() / sizeof(ElementType)) { @@ -4617,7 +4619,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_SS_TIMESTAMPOFFSET, buffers.datetimeoffsetBuffers[col - 1].data(), - sizeof(DateTimeOffset) * fetchSize, + sizeof(DateTimeOffset), buffers.indicators[col - 1].data()); break; default: @@ -4644,6 +4646,15 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& return ret; } +void ValidateFetchedRowCount(const ColumnBuffers& buffers, SQLUSMALLINT numCols, + SQLULEN numRowsFetched) { + for (SQLUSMALLINT col = 0; col < numCols; ++col) { + if (numRowsFetched > buffers.indicators[col].size()) { + ThrowStdException("Driver returned more rows than the allocated fetch buffers"); + } + } +} + // Fetch rows in batches // TODO: Move to anonymous namespace, since it is not used outside this file template @@ -4672,6 +4683,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& ret); return ret; } + ValidateFetchedRowCount(buffers, numCols, numRowsFetched); // 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 @@ -4849,6 +4861,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& PyList_SET_ITEM(row, col - 1, Py_None); continue; } + if (dataLen < 0) { + ThrowStdException("Unexpected negative data length"); + } // Performance: Use function pointer dispatch for simple types (fast // path) This eliminates the switch statement from hot loop - @@ -4874,13 +4889,6 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& Py_INCREF(Py_None); PyList_SET_ITEM(row, col - 1, Py_None); continue; - } else if (dataLen < 0) { - // Negative value is unexpected, log column index, SQL type & - // raise exception - LOG("FetchBatchData: Unexpected negative data length - " - "column=%d, SQL_type=%d, dataLen=%ld", - col, dataType, (long)dataLen); - ThrowStdException("Unexpected negative data length, check logs for details"); } assert(dataLen > 0 && "Data length must be > 0"); @@ -4890,6 +4898,9 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& case SQL_NUMERIC: { try { SQLLEN decimalDataLen = buffers.indicators[col - 1][i]; + if (decimalDataLen > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } const char* rawData = reinterpret_cast( &buffers.charBuffers[col - 1][i * MAX_DIGITS_IN_NUMERIC]); @@ -5308,7 +5319,6 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch } fetchStateGuard.close(); - return ret; } @@ -5327,7 +5337,7 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch template SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, std::vector& dataVec, SQLLEN* indicator, size_t& reservedBytes, - py::handle messages) { + py::handle messages, bool captureDiagnostics = true) { size_t start = 0; size_t end = 0; @@ -5345,18 +5355,34 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, ThrowStdException("GetDataVar only supports SQL_C_CHAR, SQL_C_WCHAR, and SQL_C_BINARY"); } - // Ensure initial buffer has space for at least the null terminator - if (dataVec.size() < sizeNullTerminator) { - ResizeNativeFetchBuffer(dataVec, sizeNullTerminator, reservedBytes); + // Binary data has no terminator, but SQL_NO_TOTAL still needs room to make progress. + const size_t initialSize = std::max(sizeNullTerminator, 1); + if (dataVec.size() < initialSize) { + ResizeNativeFetchBuffer(dataVec, initialSize, reservedBytes); } while (true) { + if (start > dataVec.size()) { + ThrowStdException("Invalid variable-length fetch buffer offset"); + } + const size_t availableBytes = CheckedMultiplySize( + dataVec.size() - start, sizeof(T), "Variable-length fetch buffer is too large"); + if (availableBytes > static_cast(std::numeric_limits::max())) { + ThrowStdException("Variable-length fetch buffer is too large"); + } SQLLEN localInd = 0; SQLRETURN ret = SQLGetData_ptr( hStmt, colNumber, cType, reinterpret_cast(dataVec.data() + start), - sizeof(T) * (dataVec.size() - start), // Available buffer size from start position + static_cast(availableBytes), &localInd); - CaptureFetchDiagnostics(hStmt, ret, messages, true); + if (captureDiagnostics) { + CaptureFetchDiagnostics(hStmt, ret, messages, true); + } + + // Indicator contents are undefined when the ODBC call fails. + if (ret == SQL_ERROR || ret == SQL_INVALID_HANDLE) { + return ret; + } // Handle NULL data if (localInd == SQL_NULL_DATA) { @@ -5364,17 +5390,22 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, return SQL_SUCCESS; } - // Check for errors (excluding SQL_SUCCESS_WITH_INFO which means more data available) - if (ret == SQL_ERROR || ret == SQL_INVALID_HANDLE) { - return ret; - } - // SQL_SUCCESS or SQL_NO_DATA means we got all the data if (ret == SQL_SUCCESS || ret == SQL_NO_DATA) { if (localInd >= 0) { - *indicator = static_cast(start) * sizeof(T) + localInd; + const size_t prefixBytes = CheckedMultiplySize( + start, sizeof(T), "Variable-length fetch result is too large"); + if (prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch result is too large"); + } + *indicator = static_cast(prefixBytes) + localInd; } else { - *indicator = localInd; // Preserve SQL_NO_TOTAL or other negative values + if (localInd != SQL_NO_TOTAL) { + ThrowStdException("Unexpected negative variable-length data indicator"); + } + *indicator = localInd; } break; } @@ -5382,13 +5413,18 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // SQL_SUCCESS_WITH_INFO means buffer was too small, need to continue fetching if (ret == SQL_SUCCESS_WITH_INFO) { // Determine how much more space we need - if (localInd < 0) { + if (localInd == SQL_NO_TOTAL) { // SQL_NO_TOTAL: driver doesn't know total size, double the buffer end = CheckedMultiplySize(dataVec.size(), 2, "Native fetch buffer size is too large"); } else { + if (localInd < 0) { + ThrowStdException("Unexpected negative variable-length data indicator"); + } + if (localInd % sizeof(T) != 0) { + ThrowStdException("Variable-length data has an invalid byte length"); + } // Driver returned total size: allocate exactly what we need - assert(localInd % sizeof(T) == 0); end = CheckedAddSize( CheckedAddSize(start, static_cast(localInd) / sizeof(T), "Native fetch buffer size is too large"), @@ -5396,6 +5432,21 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, } // The next read starts where the null terminator would have been placed + if (end <= dataVec.size()) { + const size_t prefixBytes = CheckedMultiplySize( + start, sizeof(T), "Variable-length fetch result is too large"); + if (localInd < 0 || + prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch made no progress"); + } + *indicator = static_cast(prefixBytes) + localInd; + return SQL_SUCCESS; + } + if (end > dataVec.max_size()) { + ThrowStdException("Variable-length fetch buffer is too large"); + } start = dataVec.size() - sizeNullTerminator; // Resize buffer for next iteration @@ -5410,6 +5461,63 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, return SQL_SUCCESS; } +thread_local std::vector> testGetDataResults; +thread_local size_t testGetDataResultIndex = 0; + +SQLRETURN SQL_API TestSQLGetData(SQLHANDLE, SQLUSMALLINT, SQLSMALLINT, SQLPOINTER target, + SQLLEN targetLength, SQLLEN* indicator) { + if (testGetDataResultIndex >= testGetDataResults.size()) { + return SQL_ERROR; + } + const auto [ret, value] = testGetDataResults[testGetDataResultIndex++]; + if (target != nullptr && targetLength > 0) { + std::memset(target, 'x', static_cast(targetLength)); + } + *indicator = value; + return ret; +} + +py::object RunFetchValidationTest(const std::string& scenario) { + if (scenario == "oversized_rows") { + ColumnBuffers buffers(1, 1); + ValidateFetchedRowCount(buffers, 1, 2); + } else if (scenario == "odd_wchar") { + ColumnBuffers buffers(1, 1); + buffers.wcharBuffers[0].resize(2); + buffers.indicators[0][0] = 3; + ColumnInfoExt columnInfo{}; + columnInfo.useWideChar = true; + columnInfo.fetchBufferSize = 2; + py::list row; + row.append(py::none()); + ColumnProcessors::ProcessWChar(row.ptr(), buffers, &columnInfo, 1, 0, nullptr); + } else if (scenario == "oversized_indicator") { + std::vector buffer(4); + CheckedArrowSourceOffset(buffer, 0, buffer.size(), buffer.size() + 1); + } else if (scenario == "sql_no_total_progress") { + struct RestoreSQLGetData { + SQLGetDataFunc original = SQLGetData_ptr; + ~RestoreSQLGetData() { SQLGetData_ptr = original; } + } restore; + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, + {static_cast(SQL_SUCCESS), static_cast(1)}, + }; + testGetDataResultIndex = 0; + SQLGetData_ptr = TestSQLGetData; + + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, + reservedBytes, py::none(), false); + return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); + } else { + throw py::value_error("Unknown fetch validation test scenario"); + } + return py::none(); +} + int32_t days_from_civil(int y, int m, int d) { // Implements the "days_from_civil" algorithm by Howard Hinnant // Returns number of days since Unix epoch (1970-01-01) @@ -5476,7 +5584,6 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, SQLSMALLINT nullable = colMeta["Nullable"].cast(); dataTypes[i] = dataType; - columnSizes[i] = columnSize; columnNullable[i] = (nullable != SQL_NO_NULLS); if ((dataType == SQL_WVARCHAR || dataType == SQL_WLONGVARCHAR || dataType == SQL_VARCHAR || @@ -5489,6 +5596,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } } + columnSizes[i] = columnSize; + std::string columnName = colMeta["ColumnName"].cast(); size_t nameLen = columnName.length() + 1; arrowSchemaPrivateData[i]->name = std::make_unique(nameLen); @@ -5690,9 +5799,11 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, while (idxRowArrow < arrowBatchSize) { int spaceLeftInArrowBatch = arrowBatchSize - idxRowArrow; + int currentFetchSize = fetchSize; if (fetchSize > spaceLeftInArrowBatch) { // Adjust fetch size for final batch to avoid overfetching - fetchStateGuard.setRowArraySize(spaceLeftInArrowBatch); + currentFetchSize = spaceLeftInArrowBatch; + fetchStateGuard.setRowArraySize(currentFetchSize); } { // Release GIL during the blocking ODBC fetch @@ -5710,11 +5821,15 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } // numRowsFetched is the SQL_ATTR_ROWS_FETCHED_PTR attribute. // It'll be populated by SQLFetch - assert(numRowsFetched + idxRowArrow <= static_cast(arrowBatchSize)); + if (numRowsFetched > static_cast(currentFetchSize) || + numRowsFetched > static_cast(spaceLeftInArrowBatch)) { + ThrowStdException("Driver returned more rows than the allocated Arrow buffers"); + } for (SQLULEN idxRowSql = 0; idxRowSql < numRowsFetched; idxRowSql++) { for (SQLUSMALLINT idxCol = 0; idxCol < numCols; idxCol++) { auto& arrowColumnProducer = arrowArrayPrivateData[idxCol]; auto dataType = dataTypes[idxCol]; + auto columnSize = columnSizes[idxCol]; if (hasLobColumns) { assert(idxRowSql == 0 && "GetData only works one row at a time"); @@ -5991,18 +6106,20 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_BINARY: case SQL_VARBINARY: case SQL_LONGVARBINARY: { + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedFetchColumnSize( + processedColumnSize); auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; EnsureNativeFetchBufferSize( *target_vec, CheckedAddSize(start, dataLen, "Arrow value buffer is too large"), reservedBytes); - const size_t sourceStride = hasLobColumns - ? buffers.charBuffers[idxCol].size() - : buffers.charBuffers[idxCol].size() / - static_cast(fetchSize); const size_t sourceOffset = CheckedArrowSourceOffset( - buffers.charBuffers[idxCol], idxRowSql, sourceStride, dataLen); + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -6013,18 +6130,36 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_VARCHAR: case SQL_LONGVARCHAR: { if (charCtype == SQL_C_CHAR) { + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); +#if defined(__APPLE__) || defined(__linux__) + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedAddSize( + CheckedMultiplySize( + CheckedFetchColumnSize( + processedColumnSize), + 4, + "Column fetch stride is too large"), + 1, + "Column fetch stride is too large"); +#else + const size_t fetchBufferSize = hasLobColumns + ? buffers.charBuffers[idxCol].size() + : CheckedAddSize( + CheckedFetchColumnSize( + processedColumnSize), + 1, + "Column fetch stride is too large"); +#endif auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; EnsureNativeFetchBufferSize( *target_vec, CheckedAddSize(start, dataLen, "Arrow value buffer is too large"), reservedBytes); - const size_t sourceStride = hasLobColumns - ? buffers.charBuffers[idxCol].size() - : buffers.charBuffers[idxCol].size() / - static_cast(fetchSize); const size_t sourceOffset = CheckedArrowSourceOffset( - buffers.charBuffers[idxCol], idxRowSql, sourceStride, dataLen); + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -6041,15 +6176,20 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, if (dataLen % sizeof(SQLWCHAR) != 0) { ThrowStdException("Wide-character data has an invalid byte length"); } - auto dataLenW = dataLen / sizeof(SQLWCHAR); - const size_t sourceStride = hasLobColumns - ? buffers.wcharBuffers[idxCol].size() - : buffers.wcharBuffers[idxCol].size() / - static_cast(fetchSize); + const size_t dataLenW = dataLen / sizeof(SQLWCHAR); + SQLULEN processedColumnSize = columnSize; + HandleZeroColumnSizeAtFetch(processedColumnSize); + const size_t fetchBufferSize = hasLobColumns + ? buffers.wcharBuffers[idxCol].size() + : CheckedAddSize( + CheckedFetchColumnSize( + processedColumnSize), + 1, + "Column fetch stride is too large"); const size_t sourceOffset = CheckedArrowSourceOffset( - buffers.wcharBuffers[idxCol], idxRowSql, sourceStride, dataLen); + buffers.wcharBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); auto wcharSource = &buffers.wcharBuffers[idxCol][sourceOffset]; - auto start = arrowColumnProducer->varVal[idxRowArrow]; + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; auto target_vec = &arrowColumnProducer->varData; static_assert(sizeof(SQLWCHAR) == sizeof(char16_t)); static_assert(alignof(SQLWCHAR) == alignof(char16_t)); @@ -6074,7 +6214,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, // "550e8400-e29b-41d4-a716-446655440000") Each GUID is exactly 36 bytes in // UTF-8 auto target_vec = &arrowColumnProducer->varData; - auto start = arrowColumnProducer->varVal[idxRowArrow]; + const size_t start = arrowColumnProducer->varVal[idxRowArrow]; // Ensure buffer has space for the GUID string + null terminator EnsureNativeFetchBufferSize( @@ -6125,7 +6265,9 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_DECIMAL: case SQL_NUMERIC: { // Relies on overloaded operators defined in Int128_t struct - assert(dataLen <= MAX_DIGITS_IN_NUMERIC); + if (dataLen > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } Int128_t decimalValue(0, 0); auto start = idxRowSql * MAX_DIGITS_IN_NUMERIC; int sign = 1; @@ -6576,7 +6718,6 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, } fetchStateGuard.close(); - return ret; } @@ -6733,6 +6874,7 @@ PYBIND11_MODULE(ddbc_bindings, m) { m.attr("ARCHITECTURE") = ARCHITECTURE; m.attr("SQL_NO_TOTAL") = static_cast(SQL_NO_TOTAL); + m.def("_test_fetch_validation", &RunFetchValidationTest); // Expose the C++ functions to Python m.def("ThrowStdException", &ThrowStdException); diff --git a/mssql_python/pybind/ddbc_bindings.h b/mssql_python/pybind/ddbc_bindings.h index 32d9f8067..34085ebe6 100644 --- a/mssql_python/pybind/ddbc_bindings.h +++ b/mssql_python/pybind/ddbc_bindings.h @@ -607,6 +607,9 @@ inline void ProcessChar(PyObject* row, ColumnBuffers& buffers, const void* colIn } if (colInfo->useWideChar) { + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } // Wide-char path: data was bound as SQL_C_WCHAR, lives in wcharBuffers uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR); if (!colInfo->isLob && numCharsInData < colInfo->fetchBufferSize) { @@ -720,6 +723,9 @@ inline void ProcessWChar(PyObject* row, ColumnBuffers& buffers, const void* colI } uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR); + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } // Fast path: Data fits in buffer (not LOB or truncated) // fetchBufferSize includes null-terminator, numCharsInData doesn't. Hence // '<' diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index afed788e4..91a18f098 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -17,6 +17,8 @@ import platform import threading import os +import subprocess +import sys # Import ddbc_bindings with error handling try: @@ -51,6 +53,41 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ddbc.DDBCSQLFetchArrowBatch(None, [], batch_size, 0) +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +@pytest.mark.parametrize( + ("scenario", "message"), + [ + ("oversized_rows", "more rows than the allocated fetch buffers"), + ("odd_wchar", "invalid byte length"), + ("oversized_indicator", "exceeds the allocated fetch buffer"), + ], +) +def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): + with pytest.raises(RuntimeError, match=message): + ddbc._test_fetch_validation(scenario) + + +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +def test_sql_no_total_fetch_makes_progress_in_isolated_process(): + code = """ +import mssql_python.ddbc_bindings as ddbc +ret, indicator, calls, size = ddbc._test_fetch_validation("sql_no_total_progress") +assert ret == 0 +assert indicator == 2 +assert calls == 2 +assert size == 2 +""" + result = subprocess.run( + [sys.executable, "-c", code], + cwd=os.path.dirname(os.path.dirname(__file__)), + capture_output=True, + text=True, + timeout=30, + check=False, + ) + assert result.returncode == 0, result.stderr + + @pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") class TestPybindModuleInfo: """Test module information and architecture detection.""" From b441a0d2a63f1525ac66a29d1f115f6cb9827c28 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 12:08:39 +0530 Subject: [PATCH 06/21] FIX: Harden malformed fetch handling Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 29 ++++++++++++++++----------- tests/test_010_pybind_functions.py | 1 + 2 files changed, 18 insertions(+), 12 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 77e02d537..e4a9c577f 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -4708,6 +4708,9 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& void ValidateFetchedRowCount(const ColumnBuffers& buffers, SQLUSMALLINT numCols, SQLULEN numRowsFetched) { + if (numRowsFetched == 0) { + ThrowStdException("Driver reported a successful fetch with zero rows"); + } for (SQLUSMALLINT col = 0; col < numCols; ++col) { if (numRowsFetched > buffers.indicators[col].size()) { ThrowStdException("Driver returned more rows than the allocated fetch buffers"); @@ -5397,7 +5400,11 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch template SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, std::vector& dataVec, SQLLEN* indicator, size_t& reservedBytes, - py::handle messages, bool captureDiagnostics = true) { + py::handle messages, bool captureDiagnostics = true, + SQLGetDataFunc getData = nullptr) { + if (getData == nullptr) { + getData = SQLGetData_ptr; + } size_t start = 0; size_t end = 0; @@ -5431,7 +5438,7 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, ThrowStdException("Variable-length fetch buffer is too large"); } SQLLEN localInd = 0; - SQLRETURN ret = SQLGetData_ptr( + SQLRETURN ret = getData( hStmt, colNumber, cType, reinterpret_cast(dataVec.data() + start), static_cast(availableBytes), &localInd); @@ -5540,6 +5547,9 @@ py::object RunFetchValidationTest(const std::string& scenario) { if (scenario == "oversized_rows") { ColumnBuffers buffers(1, 1); ValidateFetchedRowCount(buffers, 1, 2); + } else if (scenario == "zero_rows") { + ColumnBuffers buffers(1, 1); + ValidateFetchedRowCount(buffers, 1, 0); } else if (scenario == "odd_wchar") { ColumnBuffers buffers(1, 1); buffers.wcharBuffers[0].resize(2); @@ -5554,22 +5564,17 @@ py::object RunFetchValidationTest(const std::string& scenario) { std::vector buffer(4); CheckedArrowSourceOffset(buffer, 0, buffer.size(), buffer.size() + 1); } else if (scenario == "sql_no_total_progress") { - struct RestoreSQLGetData { - SQLGetDataFunc original = SQLGetData_ptr; - ~RestoreSQLGetData() { SQLGetData_ptr = original; } - } restore; testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, {static_cast(SQL_SUCCESS), static_cast(1)}, }; testGetDataResultIndex = 0; - SQLGetData_ptr = TestSQLGetData; std::vector buffer; SQLLEN indicator = 0; size_t reservedBytes = 0; const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, - reservedBytes, py::none(), false); + reservedBytes, py::none(), false, TestSQLGetData); return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); } else { throw py::value_error("Unknown fetch validation test scenario"); @@ -6171,14 +6176,14 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, ? buffers.charBuffers[idxCol].size() : CheckedFetchColumnSize( processedColumnSize); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; EnsureNativeFetchBufferSize( *target_vec, CheckedAddSize(start, dataLen, "Arrow value buffer is too large"), reservedBytes); - const size_t sourceOffset = CheckedArrowSourceOffset( - buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -6211,14 +6216,14 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, 1, "Column fetch stride is too large"); #endif + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; EnsureNativeFetchBufferSize( *target_vec, CheckedAddSize(start, dataLen, "Arrow value buffer is too large"), reservedBytes); - const size_t sourceOffset = CheckedArrowSourceOffset( - buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 91a18f098..c33655a6e 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -58,6 +58,7 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("scenario", "message"), [ ("oversized_rows", "more rows than the allocated fetch buffers"), + ("zero_rows", "successful fetch with zero rows"), ("odd_wchar", "invalid byte length"), ("oversized_indicator", "exceeds the allocated fetch buffer"), ], From 7dac7c141d759bea3413032e814e2097968e735b Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 12:18:22 +0530 Subject: [PATCH 07/21] FIX: Reject empty Arrow fetch batches Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 28 +++++++++++++++++++++++---- tests/test_010_pybind_functions.py | 2 ++ 2 files changed, 26 insertions(+), 4 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 10f1c69fb..fed7ffd24 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -4757,6 +4757,17 @@ void ValidateFetchedRowCount(const ColumnBuffers& buffers, SQLUSMALLINT numCols, } } +void ValidateArrowFetchedRowCount(SQLULEN numRowsFetched, int currentFetchSize, + int spaceLeftInArrowBatch) { + if (numRowsFetched == 0) { + ThrowStdException("Driver reported a successful Arrow fetch with zero rows"); + } + if (numRowsFetched > static_cast(currentFetchSize) || + numRowsFetched > static_cast(spaceLeftInArrowBatch)) { + ThrowStdException("Driver returned more rows than the allocated Arrow buffers"); + } +} + // Fetch rows in batches // TODO: Move to anonymous namespace, since it is not used outside this file template @@ -5599,9 +5610,21 @@ py::object RunFetchValidationTest(const std::string& scenario) { py::list row; row.append(py::none()); ColumnProcessors::ProcessWChar(row.ptr(), buffers, &columnInfo, 1, 0, nullptr); + } else if (scenario == "odd_char_as_wchar") { + ColumnBuffers buffers(1, 1); + buffers.wcharBuffers[0].resize(2); + buffers.indicators[0][0] = 3; + ColumnInfoExt columnInfo{}; + columnInfo.useWideChar = true; + columnInfo.fetchBufferSize = 2; + py::list row; + row.append(py::none()); + ColumnProcessors::ProcessChar(row.ptr(), buffers, &columnInfo, 1, 0, nullptr); } else if (scenario == "oversized_indicator") { std::vector buffer(4); CheckedArrowSourceOffset(buffer, 0, buffer.size(), buffer.size() + 1); + } else if (scenario == "zero_arrow_rows") { + ValidateArrowFetchedRowCount(0, 1, 1); } else if (scenario == "sql_no_total_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, @@ -5924,10 +5947,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, } // numRowsFetched is the SQL_ATTR_ROWS_FETCHED_PTR attribute. // It'll be populated by SQLFetch - if (numRowsFetched > static_cast(currentFetchSize) || - numRowsFetched > static_cast(spaceLeftInArrowBatch)) { - ThrowStdException("Driver returned more rows than the allocated Arrow buffers"); - } + ValidateArrowFetchedRowCount(numRowsFetched, currentFetchSize, spaceLeftInArrowBatch); for (SQLULEN idxRowSql = 0; idxRowSql < numRowsFetched; idxRowSql++) { for (SQLUSMALLINT idxCol = 0; idxCol < numCols; idxCol++) { auto& arrowColumnProducer = arrowArrayPrivateData[idxCol]; diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index c33655a6e..63fbe538e 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -60,7 +60,9 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("oversized_rows", "more rows than the allocated fetch buffers"), ("zero_rows", "successful fetch with zero rows"), ("odd_wchar", "invalid byte length"), + ("odd_char_as_wchar", "invalid byte length"), ("oversized_indicator", "exceeds the allocated fetch buffer"), + ("zero_arrow_rows", "successful Arrow fetch with zero rows"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): From 81b63dba39dd3c470b355d1fd3c85faabd2656c1 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 12:31:40 +0530 Subject: [PATCH 08/21] FIX: Reset fetched row counts per call Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 6 ++++++ tests/test_010_pybind_functions.py | 2 ++ 2 files changed, 8 insertions(+) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 353f66ca5..e2c50ca4e 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -4779,6 +4779,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& PERF_TIMER("FetchBatchData"); LOG("FetchBatchData: Fetching data in batches"); SQLRETURN ret; + numRowsFetched = 0; { // Release the GIL during the blocking ODBC fetch py::gil_scoped_release release; @@ -5625,6 +5626,10 @@ py::object RunFetchValidationTest(const std::string& scenario) { CheckedArrowSourceOffset(buffer, 0, buffer.size(), buffer.size() + 1); } else if (scenario == "zero_arrow_rows") { ValidateArrowFetchedRowCount(0, 1, 1); + } else if (scenario == "oversized_arrow_fetch") { + ValidateArrowFetchedRowCount(2, 1, 2); + } else if (scenario == "oversized_arrow_batch") { + ValidateArrowFetchedRowCount(2, 2, 1); } else if (scenario == "sql_no_total_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, @@ -5931,6 +5936,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, currentFetchSize = spaceLeftInArrowBatch; fetchStateGuard.setRowArraySize(currentFetchSize); } + numRowsFetched = 0; { // Release GIL during the blocking ODBC fetch py::gil_scoped_release release; diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 63fbe538e..00eb0e62e 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -63,6 +63,8 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("odd_char_as_wchar", "invalid byte length"), ("oversized_indicator", "exceeds the allocated fetch buffer"), ("zero_arrow_rows", "successful Arrow fetch with zero rows"), + ("oversized_arrow_fetch", "more rows than the allocated Arrow buffers"), + ("oversized_arrow_batch", "more rows than the allocated Arrow buffers"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): From 5cac4e91285fb41c0d86292c4777b3c8ebd31ff5 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 12:42:58 +0530 Subject: [PATCH 09/21] FIX: Exclude Arrow text terminators Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 18 ++++++++++++++++++ tests/test_010_pybind_functions.py | 2 ++ 2 files changed, 20 insertions(+) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 9f28974cd..7c6cdb688 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -382,6 +382,18 @@ size_t CheckedArrowSourceOffset(const std::vector& buffer, size_t r return offset; } +template +void ValidateArrowTextPayloadLength(size_t stride, size_t dataBytes) { + if (stride == 0) { + ThrowStdException("Arrow text fetch buffer has no terminator storage"); + } + const size_t payloadBytes = CheckedMultiplySize( + stride - 1, sizeof(ElementType), "Arrow text source size is too large"); + if (dataBytes > payloadBytes) { + ThrowStdException("Driver data length exceeds the Arrow text payload capacity"); + } +} + constexpr int MAX_NATIVE_ROW_COUNT = 1000000; constexpr size_t MAX_NATIVE_FETCH_BYTES = 256ULL * 1024 * 1024; constexpr size_t MAX_NATIVE_PARAMETER_BYTES = 256ULL * 1024 * 1024; @@ -5628,6 +5640,10 @@ py::object RunFetchValidationTest(const std::string& scenario) { ValidateArrowFetchedRowCount(2, 1, 2); } else if (scenario == "oversized_arrow_batch") { ValidateArrowFetchedRowCount(2, 2, 1); + } else if (scenario == "char_terminator_indicator") { + ValidateArrowTextPayloadLength(4, 4); + } else if (scenario == "wchar_terminator_indicator") { + ValidateArrowTextPayloadLength(4, 4 * sizeof(SQLWCHAR)); } else if (scenario == "sql_no_total_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, @@ -6279,6 +6295,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, 1, "Column fetch stride is too large"); #endif + ValidateArrowTextPayloadLength(fetchBufferSize, dataLen); const size_t sourceOffset = CheckedArrowSourceOffset( buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); auto target_vec = &arrowColumnProducer->varData; @@ -6313,6 +6330,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, processedColumnSize), 1, "Column fetch stride is too large"); + ValidateArrowTextPayloadLength(fetchBufferSize, dataLen); const size_t sourceOffset = CheckedArrowSourceOffset( buffers.wcharBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); auto wcharSource = &buffers.wcharBuffers[idxCol][sourceOffset]; diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 00eb0e62e..e643d0222 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -65,6 +65,8 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("zero_arrow_rows", "successful Arrow fetch with zero rows"), ("oversized_arrow_fetch", "more rows than the allocated Arrow buffers"), ("oversized_arrow_batch", "more rows than the allocated Arrow buffers"), + ("char_terminator_indicator", "exceeds the Arrow text payload capacity"), + ("wchar_terminator_indicator", "exceeds the Arrow text payload capacity"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): From 732231939807d8a28a5612aa866b9b1552865670 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 13:03:22 +0530 Subject: [PATCH 10/21] FIX: Validate progressive fetch completion Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 65 ++++++++++++++++++++++++++- tests/test_010_pybind_functions.py | 11 +++++ 2 files changed, 74 insertions(+), 2 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index fa577d739..60be23b0a 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -2312,6 +2312,24 @@ static void CaptureFetchDiagnostics(SQLHSTMT hStmt, SQLRETURN ret, py::handle me AppendDiagRecords(hStmt, SQL_HANDLE_STMT, messages, internalTruncation); } +static bool HasDataTruncationDiagnostic(SQLHSTMT hStmt) { + SQLWCHAR state[6] = {}; + SQLINTEGER nativeError = 0; + SQLWCHAR message[SQL_MAX_MESSAGE_LENGTH] = {}; + SQLSMALLINT messageLength = 0; + for (SQLSMALLINT record = 1;; ++record) { + const SQLRETURN ret = + SQLGetDiagRec_ptr(SQL_HANDLE_STMT, hStmt, record, state, &nativeError, message, + SQL_MAX_MESSAGE_LENGTH, &messageLength); + if (ret == SQL_NO_DATA) return false; + if (!SQL_SUCCEEDED(ret)) return false; + if (state[0] == '0' && state[1] == '1' && state[2] == '0' && state[3] == '0' && + state[4] == '4') { + return true; + } + } +} + static void CheckFetchError(const SqlHandlePtr& handle, SQLRETURN ret) { if (ret < 0) py::module_::import("mssql_python.helpers") @@ -4774,6 +4792,9 @@ void ValidateFetchedRowCount(const ColumnBuffers& buffers, SQLUSMALLINT numCols, if (numRowsFetched == 0) { ThrowStdException("Driver reported a successful fetch with zero rows"); } + if (numCols == 0) { + ThrowStdException("Driver returned rows for a result set with no columns"); + } for (SQLUSMALLINT col = 0; col < numCols; ++col) { if (numRowsFetched > buffers.indicators[col].size()) { ThrowStdException("Driver returned more rows than the allocated fetch buffers"); @@ -5526,14 +5547,24 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, return ret; } + if (ret == SQL_NO_DATA) { + const size_t prefixBytes = CheckedMultiplySize( + start, sizeof(T), "Variable-length fetch result is too large"); + if (prefixBytes > static_cast(std::numeric_limits::max())) { + ThrowStdException("Variable-length fetch result is too large"); + } + *indicator = static_cast(prefixBytes); + break; + } + // Handle NULL data if (localInd == SQL_NULL_DATA) { *indicator = SQL_NULL_DATA; return SQL_SUCCESS; } - // SQL_SUCCESS or SQL_NO_DATA means we got all the data - if (ret == SQL_SUCCESS || ret == SQL_NO_DATA) { + // SQL_SUCCESS means we got all the data + if (ret == SQL_SUCCESS) { if (localInd >= 0) { const size_t prefixBytes = CheckedMultiplySize( start, sizeof(T), "Variable-length fetch result is too large"); @@ -5575,6 +5606,11 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // The next read starts where the null terminator would have been placed if (end <= dataVec.size()) { + const bool isTruncation = + !captureDiagnostics || HasDataTruncationDiagnostic(hStmt); + if (isTruncation) { + ThrowStdException("Variable-length fetch truncation made no progress"); + } const size_t prefixBytes = CheckedMultiplySize( start, sizeof(T), "Variable-length fetch result is too large"); if (localInd < 0 || @@ -5622,6 +5658,9 @@ py::object RunFetchValidationTest(const std::string& scenario) { if (scenario == "oversized_rows") { ColumnBuffers buffers(1, 1); ValidateFetchedRowCount(buffers, 1, 2); + } else if (scenario == "rows_without_columns") { + ColumnBuffers buffers(0, 1); + ValidateFetchedRowCount(buffers, 0, 1); } else if (scenario == "zero_rows") { ColumnBuffers buffers(1, 1); ValidateFetchedRowCount(buffers, 1, 0); @@ -5665,6 +5704,28 @@ py::object RunFetchValidationTest(const std::string& scenario) { }; testGetDataResultIndex = 0; + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, + reservedBytes, py::none(), false, TestSQLGetData); + return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); + } else if (scenario == "truncation_no_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(0)}, + }; + testGetDataResultIndex = 0; + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, py::none(), + false, TestSQLGetData); + } else if (scenario == "sql_no_total_no_data") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, + {static_cast(SQL_NO_DATA), std::numeric_limits::max()}, + }; + testGetDataResultIndex = 0; std::vector buffer; SQLLEN indicator = 0; size_t reservedBytes = 0; diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index e643d0222..d99f8e9ab 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -58,6 +58,7 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("scenario", "message"), [ ("oversized_rows", "more rows than the allocated fetch buffers"), + ("rows_without_columns", "result set with no columns"), ("zero_rows", "successful fetch with zero rows"), ("odd_wchar", "invalid byte length"), ("odd_char_as_wchar", "invalid byte length"), @@ -67,6 +68,7 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("oversized_arrow_batch", "more rows than the allocated Arrow buffers"), ("char_terminator_indicator", "exceeds the Arrow text payload capacity"), ("wchar_terminator_indicator", "exceeds the Arrow text payload capacity"), + ("truncation_no_progress", "truncation made no progress"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): @@ -95,6 +97,15 @@ def test_sql_no_total_fetch_makes_progress_in_isolated_process(): assert result.returncode == 0, result.stderr +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +def test_sql_no_data_ignores_undefined_indicator(): + ret, indicator, calls, size = ddbc._test_fetch_validation("sql_no_total_no_data") + assert ret == 0 + assert indicator == 1 + assert calls == 2 + assert size == 2 + + @pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") class TestPybindModuleInfo: """Test module information and architecture detection.""" From d6ec5ada4440f5260d097da2adc9f0931191613c Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 13:12:27 +0530 Subject: [PATCH 11/21] FIX: Reject invalid fetch completion Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 29 +++++++++++++++++++++------ tests/test_010_pybind_functions.py | 2 ++ 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 60be23b0a..49902b6bd 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -394,6 +394,12 @@ void ValidateArrowTextPayloadLength(size_t stride, size_t dataBytes) { } } +void ValidateDecimalDataLength(uint64_t dataLength) { + if (dataLength > MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } +} + constexpr int MAX_NATIVE_ROW_COUNT = 1000000; constexpr size_t MAX_NATIVE_FETCH_BYTES = 256ULL * 1024 * 1024; constexpr size_t MAX_NATIVE_PARAMETER_BYTES = 256ULL * 1024 * 1024; @@ -5057,9 +5063,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& case SQL_NUMERIC: { try { SQLLEN decimalDataLen = buffers.indicators[col - 1][i]; - if (decimalDataLen > MAX_DIGITS_IN_NUMERIC) { - ThrowStdException("Decimal data exceeds the allocated fetch buffer"); - } + ValidateDecimalDataLength(static_cast(decimalDataLen)); const char* rawData = reinterpret_cast( &buffers.charBuffers[col - 1][i * MAX_DIGITS_IN_NUMERIC]); @@ -5548,6 +5552,9 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, } if (ret == SQL_NO_DATA) { + if (start == 0) { + ThrowStdException("Variable-length fetch returned no data before making progress"); + } const size_t prefixBytes = CheckedMultiplySize( start, sizeof(T), "Variable-length fetch result is too large"); if (prefixBytes > static_cast(std::numeric_limits::max())) { @@ -5710,6 +5717,18 @@ py::object RunFetchValidationTest(const std::string& scenario) { const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, py::none(), false, TestSQLGetData); return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); + } else if (scenario == "first_call_no_data") { + testGetDataResults = { + {static_cast(SQL_NO_DATA), std::numeric_limits::max()}, + }; + testGetDataResultIndex = 0; + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, py::none(), + false, TestSQLGetData); + } else if (scenario == "oversized_decimal_indicator") { + ValidateDecimalDataLength(MAX_DIGITS_IN_NUMERIC + 1); } else if (scenario == "truncation_no_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(0)}, @@ -6485,9 +6504,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_DECIMAL: case SQL_NUMERIC: { // Relies on overloaded operators defined in Int128_t struct - if (dataLen > MAX_DIGITS_IN_NUMERIC) { - ThrowStdException("Decimal data exceeds the allocated fetch buffer"); - } + ValidateDecimalDataLength(dataLen); Int128_t decimalValue(0, 0); auto start = idxRowSql * MAX_DIGITS_IN_NUMERIC; int sign = 1; diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index d99f8e9ab..2efe4f9c9 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -69,6 +69,8 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("char_terminator_indicator", "exceeds the Arrow text payload capacity"), ("wchar_terminator_indicator", "exceeds the Arrow text payload capacity"), ("truncation_no_progress", "truncation made no progress"), + ("first_call_no_data", "no data before making progress"), + ("oversized_decimal_indicator", "Decimal data exceeds the allocated fetch buffer"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): From 6b7524df3ad6471fdbefcb2eec4b8949c2ba40e1 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 13:36:26 +0530 Subject: [PATCH 12/21] FIX: Tighten decimal fetch capacity Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 14 ++++++++++++-- tests/test_010_pybind_functions.py | 1 + 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 83803e4fd..f18398ca5 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -395,7 +395,7 @@ void ValidateArrowTextPayloadLength(size_t stride, size_t dataBytes) { } void ValidateDecimalDataLength(uint64_t dataLength) { - if (dataLength > MAX_DIGITS_IN_NUMERIC) { + if (dataLength >= MAX_DIGITS_IN_NUMERIC) { ThrowStdException("Decimal data exceeds the allocated fetch buffer"); } } @@ -5729,7 +5729,17 @@ py::object RunFetchValidationTest(const std::string& scenario) { GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, py::none(), false, TestSQLGetData); } else if (scenario == "oversized_decimal_indicator") { - ValidateDecimalDataLength(MAX_DIGITS_IN_NUMERIC + 1); + ValidateDecimalDataLength(MAX_DIGITS_IN_NUMERIC); + } else if (scenario == "odd_streamed_wchar") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(3)}, + }; + testGetDataResultIndex = 0; + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + GetDataVar(nullptr, 1, SQL_C_WCHAR, buffer, &indicator, reservedBytes, py::none(), + false, TestSQLGetData); } else if (scenario == "truncation_no_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(0)}, diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 2efe4f9c9..9893ea1db 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -71,6 +71,7 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("truncation_no_progress", "truncation made no progress"), ("first_call_no_data", "no data before making progress"), ("oversized_decimal_indicator", "Decimal data exceeds the allocated fetch buffer"), + ("odd_streamed_wchar", "invalid byte length"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): From 05470e68ac986851ee22f9cb5b4469b7243642d6 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 13:38:00 +0530 Subject: [PATCH 13/21] FIX: Retain Arrow column size metadata Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 2 ++ 1 file changed, 2 insertions(+) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 1dbea3f6e..26af89aa3 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -5814,6 +5814,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, bool hasLobColumns = false; std::vector dataTypes(numCols); + std::vector columnSizes(numCols); std::vector columnNullable(numCols); std::vector columnVarLen(numCols, false); std::vector nullCounts(numCols, 0); @@ -5832,6 +5833,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, SQLSMALLINT nullable = colMeta["Nullable"].cast(); dataTypes[i] = dataType; + columnSizes[i] = columnSize; columnNullable[i] = (nullable != SQL_NO_NULLS); if ((dataType == SQL_WVARCHAR || dataType == SQL_WLONGVARCHAR || dataType == SQL_VARCHAR || From 4d721428f862dab88048df5f4ceb8df018620a6a Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 14:03:46 +0530 Subject: [PATCH 14/21] TEST: Exercise fetch diagnostic classification Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 27 ++++++++++++++++++++++----- tests/test_010_pybind_functions.py | 9 +++++++++ 2 files changed, 31 insertions(+), 5 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 26af89aa3..d4c5f7547 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -5501,7 +5501,9 @@ template SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, std::vector& dataVec, SQLLEN* indicator, size_t& reservedBytes, py::handle messages, bool captureDiagnostics = true, - SQLGetDataFunc getData = nullptr) { + SQLGetDataFunc getData = nullptr, + bool (*hasTruncationDiagnostic)(SQLHSTMT) = + HasDataTruncationDiagnostic) { if (getData == nullptr) { getData = SQLGetData_ptr; } @@ -5613,8 +5615,7 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // The next read starts where the null terminator would have been placed if (end <= dataVec.size()) { - const bool isTruncation = - !captureDiagnostics || HasDataTruncationDiagnostic(hStmt); + const bool isTruncation = hasTruncationDiagnostic(hStmt); if (isTruncation) { ThrowStdException("Variable-length fetch truncation made no progress"); } @@ -5661,6 +5662,10 @@ SQLRETURN SQL_API TestSQLGetData(SQLHANDLE, SQLUSMALLINT, SQLSMALLINT, SQLPOINTE return ret; } +bool TestHasTruncationDiagnostic(SQLHSTMT) { return true; } + +bool TestHasUnrelatedWarning(SQLHSTMT) { return false; } + py::object RunFetchValidationTest(const std::string& scenario) { if (scenario == "oversized_rows") { ColumnBuffers buffers(1, 1); @@ -5726,7 +5731,19 @@ py::object RunFetchValidationTest(const std::string& scenario) { SQLLEN indicator = 0; size_t reservedBytes = 0; GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, py::none(), - false, TestSQLGetData); + false, TestSQLGetData, TestHasTruncationDiagnostic); + } else if (scenario == "unrelated_warning_no_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(0)}, + }; + testGetDataResultIndex = 0; + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + const SQLRETURN ret = + GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, + py::none(), false, TestSQLGetData, TestHasUnrelatedWarning); + return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); } else if (scenario == "oversized_decimal_indicator") { ValidateDecimalDataLength(MAX_DIGITS_IN_NUMERIC); } else if (scenario == "odd_streamed_wchar") { @@ -5748,7 +5765,7 @@ py::object RunFetchValidationTest(const std::string& scenario) { SQLLEN indicator = 0; size_t reservedBytes = 0; GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, py::none(), - false, TestSQLGetData); + false, TestSQLGetData, TestHasTruncationDiagnostic); } else if (scenario == "sql_no_total_no_data") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(SQL_NO_TOTAL)}, diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 9893ea1db..4df53a479 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -109,6 +109,15 @@ def test_sql_no_data_ignores_undefined_indicator(): assert size == 2 +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +def test_unrelated_warning_without_growth_is_not_treated_as_truncation(): + ret, indicator, calls, size = ddbc._test_fetch_validation("unrelated_warning_no_progress") + assert ret == 0 + assert indicator == 0 + assert calls == 1 + assert size == 1 + + @pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") class TestPybindModuleInfo: """Test module information and architecture detection.""" From 0ecce543c9087d758ca5c78f63c36576914f7c7e Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 14:41:29 +0530 Subject: [PATCH 15/21] FIX: Validate direct streamed fetches Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 69 ++++++++++++++++++++++++--- tests/test_010_pybind_functions.py | 4 ++ 2 files changed, 66 insertions(+), 7 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 95a827b89..a7107eacc 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3554,10 +3554,16 @@ SQLRETURN SQLFetch_wrap(SqlHandlePtr StatementHandle) { return ret; } -// Non-static so it can be called from inline functions in header -py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT cType, - bool isWideChar, bool isBinary, const std::string& charEncoding, - py::handle messages) { +inline void ValidateWideCharByteLength(SQLLEN dataLen) { + if (dataLen > 0 && dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } +} + +static py::object FetchLobColumnDataImpl( + SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT cType, bool isWideChar, bool isBinary, + const std::string& charEncoding, py::handle messages, bool captureDiagnostics, + SQLGetDataFunc getData, bool (*hasTruncationDiagnostic)(SQLHSTMT)) { PERF_TIMER("FetchLobColumnData"); std::vector buffer; size_t reservedBytes = 0; @@ -3572,9 +3578,11 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT { // Release the GIL during blocking SQLGetData LOB streaming py::gil_scoped_release release; - ret = SQLGetData_ptr(hStmt, colIndex, cType, chunk.data(), DAE_CHUNK_SIZE, &actualRead); + ret = getData(hStmt, colIndex, cType, chunk.data(), DAE_CHUNK_SIZE, &actualRead); + } + if (captureDiagnostics) { + CaptureFetchDiagnostics(hStmt, ret, messages, true); } - CaptureFetchDiagnostics(hStmt, ret, messages, true); if (ret == SQL_ERROR || !SQL_SUCCEEDED(ret) && ret != SQL_SUCCESS_WITH_INFO) { std::ostringstream oss; @@ -3587,6 +3595,12 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT LOG("FetchLobColumnData: Column %d is NULL at loop %d", colIndex, loopCount); return py::none(); } + if (actualRead < 0 && actualRead != SQL_NO_TOTAL) { + ThrowStdException("Unexpected negative LOB data indicator"); + } + if (isWideChar) { + ValidateWideCharByteLength(actualRead); + } size_t bytesRead = 0; if (actualRead >= 0) { @@ -3595,9 +3609,14 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT bytesRead = DAE_CHUNK_SIZE; } } else { - // fallback: use full buffer size if actualRead is unknown bytesRead = DAE_CHUNK_SIZE; } + if (ret == SQL_SUCCESS_WITH_INFO && bytesRead == 0) { + if (hasTruncationDiagnostic(hStmt)) { + ThrowStdException("LOB fetch truncation made no progress"); + } + break; + } // For character data, trim trailing null terminators if (!isBinary && bytesRead > 0) { @@ -3628,6 +3647,7 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT loopCount); } } + } } if (bytesRead > 0) { @@ -3688,6 +3708,15 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT } } +// Non-static so it can be called from inline functions in header +py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT cType, + bool isWideChar, bool isBinary, const std::string& charEncoding, + py::handle messages) { + return FetchLobColumnDataImpl(hStmt, colIndex, cType, isWideChar, isBinary, charEncoding, + messages, true, SQLGetData_ptr, + HasDataTruncationDiagnostic); +} + // Helper function to map sql_variant's underlying C type to SQL data type // This allows sql_variant to reuse existing fetch logic for each data type SQLSMALLINT MapVariantCTypeToSQLType(SQLLEN variantCType) { @@ -3922,6 +3951,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p dataLen >= static_cast(fetchBufferSize))); if (SQL_SUCCEEDED(ret)) { if (dataLen > 0) { + ValidateWideCharByteLength(dataLen); uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR); if (numCharsInData < dataBuffer.size()) { // Construct with explicit length: SQLGetData reports the @@ -4101,6 +4131,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p dataLen >= static_cast(fetchBufferSize))); if (SQL_SUCCEEDED(ret)) { if (dataLen > 0) { + ValidateWideCharByteLength(dataLen); uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR); if (numCharsInData < dataBuffer.size()) { // Construct with explicit length: SQLGetData reports the @@ -5760,6 +5791,30 @@ py::object RunFetchValidationTest(const std::string& scenario) { size_t reservedBytes = 0; GetDataVar(nullptr, 1, SQL_C_WCHAR, buffer, &indicator, reservedBytes, py::none(), false, TestSQLGetData); + } else if (scenario == "unexpected_lob_indicator") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(-2)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), + false, TestSQLGetData, TestHasTruncationDiagnostic); + } else if (scenario == "lob_truncation_no_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(0)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), + false, TestSQLGetData, TestHasTruncationDiagnostic); + } else if (scenario == "odd_direct_wchar") { + ValidateWideCharByteLength(3); + } else if (scenario == "odd_lob_wchar") { + testGetDataResults = { + {static_cast(SQL_SUCCESS), static_cast(3)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_WCHAR, true, false, "utf-16le", + py::none(), false, TestSQLGetData, + TestHasTruncationDiagnostic); } else if (scenario == "truncation_no_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(0)}, diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 4df53a479..f69e43748 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -72,6 +72,10 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("first_call_no_data", "no data before making progress"), ("oversized_decimal_indicator", "Decimal data exceeds the allocated fetch buffer"), ("odd_streamed_wchar", "invalid byte length"), + ("unexpected_lob_indicator", "Unexpected negative LOB data indicator"), + ("lob_truncation_no_progress", "LOB fetch truncation made no progress"), + ("odd_direct_wchar", "invalid byte length"), + ("odd_lob_wchar", "invalid byte length"), ], ) def test_driver_fetch_validation_rejects_malformed_lengths(scenario, message): From 3c5e5b6abd16c3d435a59354dc9bbde572e37e9a Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 14:54:31 +0530 Subject: [PATCH 16/21] FIX: Reject malformed LOB success lengths Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 37 ++++++++++++++++++++++----- tests/test_010_pybind_functions.py | 8 ++++++ 2 files changed, 39 insertions(+), 6 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index a7107eacc..b119bcafc 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3602,6 +3602,14 @@ static py::object FetchLobColumnDataImpl( ValidateWideCharByteLength(actualRead); } + const size_t terminatorBytes = isBinary ? 0 : (isWideChar ? sizeof(SQLWCHAR) : 1); + const size_t payloadCapacity = DAE_CHUNK_SIZE - terminatorBytes; + if (ret == SQL_SUCCESS && actualRead >= 0 && + static_cast(actualRead) > payloadCapacity) { + ThrowStdException("LOB data indicator exceeds the fetch buffer capacity"); + } + const bool continueForTruncation = + ret == SQL_SUCCESS_WITH_INFO && hasTruncationDiagnostic(hStmt); size_t bytesRead = 0; if (actualRead >= 0) { bytesRead = static_cast(actualRead); @@ -3611,11 +3619,8 @@ static py::object FetchLobColumnDataImpl( } else { bytesRead = DAE_CHUNK_SIZE; } - if (ret == SQL_SUCCESS_WITH_INFO && bytesRead == 0) { - if (hasTruncationDiagnostic(hStmt)) { - ThrowStdException("LOB fetch truncation made no progress"); - } - break; + if (continueForTruncation && bytesRead == 0) { + ThrowStdException("LOB fetch truncation made no progress"); } // For character data, trim trailing null terminators @@ -3658,7 +3663,8 @@ static py::object FetchLobColumnDataImpl( std::copy_n(chunk.data(), bytesRead, buffer.data() + previousSize); LOG("FetchLobColumnData: Appended %zu bytes at loop %d", bytesRead, loopCount); } - if (ret == SQL_SUCCESS) { + if (ret == SQL_SUCCESS || + (ret == SQL_SUCCESS_WITH_INFO && !continueForTruncation)) { LOG("FetchLobColumnData: SQL_SUCCESS - no more data at loop %d", loopCount); break; } @@ -5805,6 +5811,25 @@ py::object RunFetchValidationTest(const std::string& scenario) { testGetDataResultIndex = 0; FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), false, TestSQLGetData, TestHasTruncationDiagnostic); + } else if (scenario == "oversized_lob_success") { + testGetDataResults = { + {static_cast(SQL_SUCCESS), + static_cast(DAE_CHUNK_SIZE + 1)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), + false, TestSQLGetData, TestHasTruncationDiagnostic); + } else if (scenario == "lob_unrelated_warning_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(1)}, + }; + testGetDataResultIndex = 0; + py::bytes value = + FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", + py::none(), false, TestSQLGetData, + TestHasUnrelatedWarning) + .cast(); + return py::make_tuple(py::len(value), testGetDataResultIndex); } else if (scenario == "odd_direct_wchar") { ValidateWideCharByteLength(3); } else if (scenario == "odd_lob_wchar") { diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index f69e43748..5cf3116df 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -74,6 +74,7 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("odd_streamed_wchar", "invalid byte length"), ("unexpected_lob_indicator", "Unexpected negative LOB data indicator"), ("lob_truncation_no_progress", "LOB fetch truncation made no progress"), + ("oversized_lob_success", "LOB data indicator exceeds the fetch buffer capacity"), ("odd_direct_wchar", "invalid byte length"), ("odd_lob_wchar", "invalid byte length"), ], @@ -122,6 +123,13 @@ def test_unrelated_warning_without_growth_is_not_treated_as_truncation(): assert size == 1 +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +def test_lob_unrelated_warning_with_progress_completes(): + size, calls = ddbc._test_fetch_validation("lob_unrelated_warning_progress") + assert size == 1 + assert calls == 1 + + @pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") class TestPybindModuleInfo: """Test module information and architecture detection.""" From 1666deb64098831caebb0b5404f4c52129d1a9d6 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 15:12:11 +0530 Subject: [PATCH 17/21] FIX: Bound unrelated fetch warnings Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 75 +++++++++++++++++++-------- tests/test_010_pybind_functions.py | 2 + 2 files changed, 56 insertions(+), 21 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 1db984d69..a180182d0 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3612,12 +3612,12 @@ static py::object FetchLobColumnDataImpl( const size_t terminatorBytes = isBinary ? 0 : (isWideChar ? sizeof(SQLWCHAR) : 1); const size_t payloadCapacity = DAE_CHUNK_SIZE - terminatorBytes; - if (ret == SQL_SUCCESS && actualRead >= 0 && - static_cast(actualRead) > payloadCapacity) { - ThrowStdException("LOB data indicator exceeds the fetch buffer capacity"); - } const bool continueForTruncation = ret == SQL_SUCCESS_WITH_INFO && hasTruncationDiagnostic(hStmt); + if (actualRead >= 0 && static_cast(actualRead) > payloadCapacity && + !continueForTruncation) { + ThrowStdException("LOB data indicator exceeds the fetch buffer capacity"); + } size_t bytesRead = 0; if (actualRead >= 0) { bytesRead = static_cast(actualRead); @@ -5643,6 +5643,32 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // SQL_SUCCESS_WITH_INFO means buffer was too small, need to continue fetching if (ret == SQL_SUCCESS_WITH_INFO) { + const bool isTruncation = hasTruncationDiagnostic(hStmt); + if (!isTruncation) { + if (localInd < 0) { + ThrowStdException( + "Unexpected negative variable-length data indicator"); + } + const size_t terminatorBytes = + CheckedMultiplySize(sizeNullTerminator, sizeof(T), + "Variable-length fetch result is too large"); + const size_t payloadCapacity = + availableBytes >= terminatorBytes ? availableBytes - terminatorBytes : 0; + if (static_cast(localInd) > payloadCapacity) { + ThrowStdException( + "Variable-length data indicator exceeds the fetch buffer capacity"); + } + const size_t prefixBytes = CheckedMultiplySize( + start, sizeof(T), "Variable-length fetch result is too large"); + if (prefixBytes > static_cast(std::numeric_limits::max()) || + localInd > std::numeric_limits::max() - + static_cast(prefixBytes)) { + ThrowStdException("Variable-length fetch result is too large"); + } + *indicator = static_cast(prefixBytes) + localInd; + return SQL_SUCCESS; + } + // Determine how much more space we need if (localInd == SQL_NO_TOTAL) { // SQL_NO_TOTAL: driver doesn't know total size, double the buffer @@ -5664,20 +5690,7 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // The next read starts where the null terminator would have been placed if (end <= dataVec.size()) { - const bool isTruncation = hasTruncationDiagnostic(hStmt); - if (isTruncation) { - ThrowStdException("Variable-length fetch truncation made no progress"); - } - const size_t prefixBytes = CheckedMultiplySize( - start, sizeof(T), "Variable-length fetch result is too large"); - if (localInd < 0 || - prefixBytes > static_cast(std::numeric_limits::max()) || - localInd > std::numeric_limits::max() - - static_cast(prefixBytes)) { - ThrowStdException("Variable-length fetch made no progress"); - } - *indicator = static_cast(prefixBytes) + localInd; - return SQL_SUCCESS; + ThrowStdException("Variable-length fetch truncation made no progress"); } if (end > dataVec.max_size()) { ThrowStdException("Variable-length fetch buffer is too large"); @@ -5769,7 +5782,8 @@ py::object RunFetchValidationTest(const std::string& scenario) { SQLLEN indicator = 0; size_t reservedBytes = 0; const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, - reservedBytes, py::none(), false, TestSQLGetData); + reservedBytes, py::none(), false, TestSQLGetData, + TestHasTruncationDiagnostic); return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); } else if (scenario == "first_call_no_data") { testGetDataResults = { @@ -5804,7 +5818,7 @@ py::object RunFetchValidationTest(const std::string& scenario) { SQLLEN indicator = 0; size_t reservedBytes = 0; GetDataVar(nullptr, 1, SQL_C_WCHAR, buffer, &indicator, reservedBytes, py::none(), - false, TestSQLGetData); + false, TestSQLGetData, TestHasTruncationDiagnostic); } else if (scenario == "unexpected_lob_indicator") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(-2)}, @@ -5838,6 +5852,24 @@ py::object RunFetchValidationTest(const std::string& scenario) { TestHasUnrelatedWarning) .cast(); return py::make_tuple(py::len(value), testGetDataResultIndex); + } else if (scenario == "lob_unrelated_warning_oversized") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), + static_cast(DAE_CHUNK_SIZE + 1)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), + false, TestSQLGetData, TestHasUnrelatedWarning); + } else if (scenario == "unrelated_warning_oversized") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(2)}, + }; + testGetDataResultIndex = 0; + std::vector buffer; + SQLLEN indicator = 0; + size_t reservedBytes = 0; + GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, reservedBytes, + py::none(), false, TestSQLGetData, TestHasUnrelatedWarning); } else if (scenario == "odd_direct_wchar") { ValidateWideCharByteLength(3); } else if (scenario == "odd_lob_wchar") { @@ -5868,7 +5900,8 @@ py::object RunFetchValidationTest(const std::string& scenario) { SQLLEN indicator = 0; size_t reservedBytes = 0; const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, - reservedBytes, py::none(), false, TestSQLGetData); + reservedBytes, py::none(), false, TestSQLGetData, + TestHasTruncationDiagnostic); return py::make_tuple(ret, indicator, testGetDataResultIndex, buffer.size()); } else { throw py::value_error("Unknown fetch validation test scenario"); diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 5cf3116df..e0a4480d5 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -75,6 +75,8 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("unexpected_lob_indicator", "Unexpected negative LOB data indicator"), ("lob_truncation_no_progress", "LOB fetch truncation made no progress"), ("oversized_lob_success", "LOB data indicator exceeds the fetch buffer capacity"), + ("lob_unrelated_warning_oversized", "LOB data indicator exceeds the fetch buffer capacity"), + ("unrelated_warning_oversized", "data indicator exceeds the fetch buffer capacity"), ("odd_direct_wchar", "invalid byte length"), ("odd_lob_wchar", "invalid byte length"), ], From 851d319fbb67b9cba87ade88eecbac4ef719e48f Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 15:38:46 +0530 Subject: [PATCH 18/21] FIX: Reject non-progressing LOB chunks Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 45 +++++++++++++++++++++++++++ tests/test_010_pybind_functions.py | 3 ++ 2 files changed, 48 insertions(+) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index bff5cf1b3..9323658f6 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3614,6 +3614,9 @@ static py::object FetchLobColumnDataImpl( const size_t payloadCapacity = DAE_CHUNK_SIZE - terminatorBytes; const bool continueForTruncation = ret == SQL_SUCCESS_WITH_INFO && hasTruncationDiagnostic(hStmt); + if (actualRead == SQL_NO_TOTAL && !continueForTruncation) { + ThrowStdException("LOB SQL_NO_TOTAL requires a truncation diagnostic"); + } if (actualRead >= 0 && static_cast(actualRead) > payloadCapacity && !continueForTruncation) { ThrowStdException("LOB data indicator exceeds the fetch buffer capacity"); @@ -3663,6 +3666,9 @@ static py::object FetchLobColumnDataImpl( } } + if (continueForTruncation && bytesRead == 0) { + ThrowStdException("LOB fetch truncation made no progress"); + } if (bytesRead > 0) { const size_t previousSize = buffer.size(); const size_t requiredSize = @@ -5724,6 +5730,20 @@ SQLRETURN SQL_API TestSQLGetData(SQLHANDLE, SQLUSMALLINT, SQLSMALLINT, SQLPOINTE return ret; } +SQLRETURN SQL_API TestSQLGetDataZeroFill(SQLHANDLE, SQLUSMALLINT, SQLSMALLINT, + SQLPOINTER target, SQLLEN targetLength, + SQLLEN* indicator) { + if (testGetDataResultIndex >= testGetDataResults.size()) { + return SQL_ERROR; + } + const auto [ret, value] = testGetDataResults[testGetDataResultIndex++]; + if (target != nullptr && targetLength > 0) { + std::memset(target, 0, static_cast(targetLength)); + } + *indicator = value; + return ret; +} + bool TestHasTruncationDiagnostic(SQLHSTMT) { return true; } bool TestHasUnrelatedWarning(SQLHSTMT) { return false; } @@ -5860,6 +5880,31 @@ py::object RunFetchValidationTest(const std::string& scenario) { testGetDataResultIndex = 0; FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), false, TestSQLGetData, TestHasUnrelatedWarning); + } else if (scenario == "lob_unrelated_warning_no_total") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), + static_cast(SQL_NO_TOTAL)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), + false, TestSQLGetData, TestHasUnrelatedWarning); + } else if (scenario == "lob_narrow_terminator_no_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(1)}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_CHAR, false, false, "utf-8", + py::none(), false, TestSQLGetDataZeroFill, + TestHasTruncationDiagnostic); + } else if (scenario == "lob_wide_terminator_no_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), + static_cast(sizeof(SQLWCHAR))}, + }; + testGetDataResultIndex = 0; + FetchLobColumnDataImpl(nullptr, 1, SQL_C_WCHAR, true, false, "utf-16le", + py::none(), false, TestSQLGetDataZeroFill, + TestHasTruncationDiagnostic); } else if (scenario == "unrelated_warning_oversized") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(2)}, diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index e0a4480d5..ff9ed2f6c 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -76,6 +76,9 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("lob_truncation_no_progress", "LOB fetch truncation made no progress"), ("oversized_lob_success", "LOB data indicator exceeds the fetch buffer capacity"), ("lob_unrelated_warning_oversized", "LOB data indicator exceeds the fetch buffer capacity"), + ("lob_unrelated_warning_no_total", "SQL_NO_TOTAL requires a truncation diagnostic"), + ("lob_narrow_terminator_no_progress", "LOB fetch truncation made no progress"), + ("lob_wide_terminator_no_progress", "LOB fetch truncation made no progress"), ("unrelated_warning_oversized", "data indicator exceeds the fetch buffer capacity"), ("odd_direct_wchar", "invalid byte length"), ("odd_lob_wchar", "invalid byte length"), From 97184578555f75535c3798381b71f20cd9f7b800 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 17:07:06 +0530 Subject: [PATCH 19/21] FIX: validate fixed-width fetch indicators Reject short, zero, or oversized driver indicators for fixed C bindings before standard-row or Arrow processors read the reusable rowset buffers. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 47 +++++++++++++++++++++++++++ tests/test_010_pybind_functions.py | 2 ++ 2 files changed, 49 insertions(+) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 01e5cf38a..8fc7dcfaa 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -400,6 +400,46 @@ void ValidateDecimalDataLength(uint64_t dataLength) { } } +size_t FixedFetchValueSize(SQLSMALLINT dataType) { + switch (dataType) { + case SQL_INTEGER: + return sizeof(SQLINTEGER); + case SQL_SMALLINT: + return sizeof(SQLSMALLINT); + case SQL_TINYINT: + case SQL_BIT: + return sizeof(SQLCHAR); + case SQL_REAL: + return sizeof(SQLREAL); + case SQL_FLOAT: + case SQL_DOUBLE: + return sizeof(SQLDOUBLE); + case SQL_BIGINT: + return sizeof(SQLBIGINT); + case SQL_TIMESTAMP: + case SQL_TYPE_TIMESTAMP: + case SQL_DATETIME: + return sizeof(SQL_TIMESTAMP_STRUCT); + case SQL_TYPE_DATE: + return sizeof(SQL_DATE_STRUCT); + case SQL_SS_TIME2: + return sizeof(SQL_SS_TIME2_STRUCT); + case SQL_GUID: + return sizeof(SQLGUID); + case SQL_SS_TIMESTAMPOFFSET: + return sizeof(DateTimeOffset); + default: + return 0; + } +} + +void ValidateFixedFetchDataLength(SQLSMALLINT dataType, uint64_t dataLength) { + const size_t expectedSize = FixedFetchValueSize(dataType); + if (expectedSize != 0 && dataLength != expectedSize) { + ThrowStdException("Fixed-width data indicator does not match the bound buffer size"); + } +} + constexpr int MAX_NATIVE_ROW_COUNT = 1000000; constexpr size_t MAX_NATIVE_FETCH_BYTES = 256ULL * 1024 * 1024; constexpr size_t MAX_NATIVE_PARAMETER_BYTES = 256ULL * 1024 * 1024; @@ -5106,6 +5146,8 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& if (dataLen < 0) { ThrowStdException("Unexpected negative data length"); } + ValidateFixedFetchDataLength(columnInfos[col - 1].dataType, + static_cast(dataLen)); // Performance: Use function pointer dispatch for simple types (fast // path) This eliminates the switch statement from hot loop - @@ -5804,6 +5846,10 @@ py::object RunFetchValidationTest(const std::string& scenario) { } else if (scenario == "oversized_indicator") { std::vector buffer(4); CheckedArrowSourceOffset(buffer, 0, buffer.size(), buffer.size() + 1); + } else if (scenario == "short_fixed_indicator") { + ValidateFixedFetchDataLength(SQL_INTEGER, sizeof(SQLINTEGER) - 1); + } else if (scenario == "oversized_fixed_indicator") { + ValidateFixedFetchDataLength(SQL_GUID, sizeof(SQLGUID) + 1); } else if (scenario == "zero_arrow_rows") { ValidateArrowFetchedRowCount(0, 1, 1); } else if (scenario == "oversized_arrow_fetch") { @@ -6560,6 +6606,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, ThrowStdException("Unexpected negative data length."); } auto dataLen = static_cast(indicator); + ValidateFixedFetchDataLength(dataType, dataLen); switch (dataType) { case SQL_SS_UDT: diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index aa109d43f..8d81913c8 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -63,6 +63,8 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("odd_wchar", "invalid byte length"), ("odd_char_as_wchar", "invalid byte length"), ("oversized_indicator", "exceeds the allocated fetch buffer"), + ("short_fixed_indicator", "does not match the bound buffer size"), + ("oversized_fixed_indicator", "does not match the bound buffer size"), ("zero_arrow_rows", "successful Arrow fetch with zero rows"), ("oversized_arrow_fetch", "more rows than the allocated Arrow buffers"), ("oversized_arrow_batch", "more rows than the allocated Arrow buffers"), From 941e5b789422f32087bf87e9c64b6f25c87722c7 Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Wed, 7 Oct 2026 17:20:06 +0530 Subject: [PATCH 20/21] FIX: validate direct fetch indicators Apply fixed-width and decimal indicator bounds before SQLGetData conversions used by fetchone and LOB-containing batches, with injected malformed-driver regressions. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 35 +++++++++++++++++++++++++++ tests/test_010_pybind_functions.py | 2 ++ 2 files changed, 37 insertions(+) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 8fc7dcfaa..b4ab2a1a3 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -422,6 +422,7 @@ size_t FixedFetchValueSize(SQLSMALLINT dataType) { return sizeof(SQL_TIMESTAMP_STRUCT); case SQL_TYPE_DATE: return sizeof(SQL_DATE_STRUCT); + case SQL_TYPE_TIME: case SQL_SS_TIME2: return sizeof(SQL_SS_TIME2_STRUCT); case SQL_GUID: @@ -440,6 +441,15 @@ void ValidateFixedFetchDataLength(SQLSMALLINT dataType, uint64_t dataLength) { } } +void ValidateDirectFetchDataLength(SQLRETURN ret, SQLSMALLINT dataType, SQLLEN indicator) { + if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { + if (indicator < 0) { + ThrowStdException("Unexpected negative data length"); + } + ValidateFixedFetchDataLength(dataType, static_cast(indicator)); + } +} + constexpr int MAX_NATIVE_ROW_COUNT = 1000000; constexpr size_t MAX_NATIVE_FETCH_BYTES = 256ULL * 1024 * 1024; constexpr size_t MAX_NATIVE_PARAMETER_BYTES = 256ULL * 1024 * 1024; @@ -4275,6 +4285,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_LONG, &intValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { row.append(static_cast(intValue)); } else { @@ -4287,6 +4298,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_SHORT, &smallIntValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4306,6 +4318,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_FLOAT, &realValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4330,6 +4343,14 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p CaptureFetchDiagnostics(hStmt, ret, messages); if (SQL_SUCCEEDED(ret)) { + if (indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } + if (indicator < 0) { + ThrowStdException("Unexpected negative data length"); + } + ValidateDecimalDataLength(static_cast(indicator)); try { // Validate 'indicator' to avoid buffer overflow and // fallback to a safe null-terminated read when length @@ -4386,6 +4407,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_DOUBLE, &doubleValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4405,6 +4427,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_SBIGINT, &bigintValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4425,6 +4448,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p ret = SQLGetData_ptr(hStmt, i, SQL_C_TYPE_DATE, &dateValue, sizeof(dateValue), &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { row.append( FetchTemporal::date(dateValue.year, dateValue.month, dateValue.day)); @@ -4439,6 +4463,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_SS_TIME2, &t2, sizeof(t2), &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { row.append(FetchTemporal::time( t2.hour, t2.minute, t2.second, t2.fraction / 1000)); // ns to µs @@ -4460,6 +4485,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p ret = SQLGetData_ptr(hStmt, i, SQL_C_TYPE_TIMESTAMP, ×tampValue, sizeof(timestampValue), &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4484,6 +4510,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p ret = SQLGetData_ptr(hStmt, i, SQL_C_SS_TIMESTAMPOFFSET, &dtoValue, sizeof(dtoValue), &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { LOG("SQLGetData: Retrieved DATETIMEOFFSET for column %d - " "%d-%d-%d %d:%d:%d, fraction_ns=%u, tz_hour=%d, " @@ -4576,6 +4603,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_TINYINT, &tinyIntValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4595,6 +4623,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_BIT, &bitValue, 0, &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { row.append(py::none()); break; @@ -4616,6 +4645,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p ret = SQLGetData_ptr(hStmt, i, SQL_C_GUID, &guidValue, sizeof(guidValue), &indicator); CaptureFetchDiagnostics(hStmt, ret, messages); + ValidateDirectFetchDataLength(ret, effectiveDataType, indicator); if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { std::vector guid_bytes(16); @@ -5850,6 +5880,11 @@ py::object RunFetchValidationTest(const std::string& scenario) { ValidateFixedFetchDataLength(SQL_INTEGER, sizeof(SQLINTEGER) - 1); } else if (scenario == "oversized_fixed_indicator") { ValidateFixedFetchDataLength(SQL_GUID, sizeof(SQLGUID) + 1); + } else if (scenario == "short_direct_fixed_indicator") { + ValidateDirectFetchDataLength(SQL_SUCCESS, SQL_TYPE_TIMESTAMP, + sizeof(SQL_TIMESTAMP_STRUCT) - 1); + } else if (scenario == "oversized_direct_decimal_indicator") { + ValidateDecimalDataLength(MAX_DIGITS_IN_NUMERIC); } else if (scenario == "zero_arrow_rows") { ValidateArrowFetchedRowCount(0, 1, 1); } else if (scenario == "oversized_arrow_fetch") { diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 8d81913c8..e1e028080 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -65,6 +65,8 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("oversized_indicator", "exceeds the allocated fetch buffer"), ("short_fixed_indicator", "does not match the bound buffer size"), ("oversized_fixed_indicator", "does not match the bound buffer size"), + ("short_direct_fixed_indicator", "does not match the bound buffer size"), + ("oversized_direct_decimal_indicator", "Decimal data exceeds"), ("zero_arrow_rows", "successful Arrow fetch with zero rows"), ("oversized_arrow_fetch", "more rows than the allocated Arrow buffers"), ("oversized_arrow_batch", "more rows than the allocated Arrow buffers"), From 2552d923bf70841a3480bf0f247d068847a1bc3d Mon Sep 17 00:00:00 2001 From: gargsaumya Date: Thu, 8 Oct 2026 12:36:04 +0530 Subject: [PATCH 21/21] FIX: preserve LOB NUL payload performance Use ODBC indicators and payload capacity instead of trimming data bytes, preserve embedded NUL values, and precompute fixed-width sizes outside row and Arrow cell loops. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql_python/pybind/ddbc_bindings.cpp | 81 ++++++++++----------------- tests/test_010_pybind_functions.py | 10 +++- 2 files changed, 39 insertions(+), 52 deletions(-) diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index ba69279c2..e6855747c 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -3726,47 +3726,11 @@ static py::object FetchLobColumnDataImpl( size_t bytesRead = 0; if (actualRead >= 0) { bytesRead = static_cast(actualRead); - if (bytesRead > DAE_CHUNK_SIZE) { - bytesRead = DAE_CHUNK_SIZE; + if (continueForTruncation && bytesRead > payloadCapacity) { + bytesRead = payloadCapacity; } } else { - bytesRead = DAE_CHUNK_SIZE; - } - if (continueForTruncation && bytesRead == 0) { - ThrowStdException("LOB fetch truncation made no progress"); - } - - // For character data, trim trailing null terminators - if (!isBinary && bytesRead > 0) { - if (!isWideChar) { - // Narrow characters - while (bytesRead > 0 && chunk[bytesRead - 1] == '\0') { - --bytesRead; - } - if (bytesRead < DAE_CHUNK_SIZE) { - LOG("FetchLobColumnData: Trimmed null terminator from " - "narrow char data - loop=%d", - loopCount); - } - } else { - // Wide characters - size_t wcharSize = sizeof(SQLWCHAR); - if (bytesRead >= wcharSize && (bytesRead % wcharSize == 0)) { - size_t wcharCount = bytesRead / wcharSize; - std::vector alignedBuf(wcharCount); - std::memcpy(alignedBuf.data(), chunk.data(), bytesRead); - while (wcharCount > 0 && alignedBuf[wcharCount - 1] == 0) { - --wcharCount; - bytesRead -= wcharSize; - } - if (bytesRead < DAE_CHUNK_SIZE) { - LOG("FetchLobColumnData: Trimmed null terminator from " - "wide char data - loop=%d", - loopCount); - } - } - - } + bytesRead = payloadCapacity; } if (continueForTruncation && bytesRead == 0) { ThrowStdException("LOB fetch truncation made no progress"); @@ -5035,6 +4999,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& SQLULEN columnSize; SQLULEN processedColumnSize; uint64_t fetchBufferSize; + size_t fixedValueSize; bool isLob; }; const bool useWideChar = (charCtype == SQL_C_WCHAR); @@ -5053,6 +5018,8 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& const auto& columnMeta = GetFetchColumnMetadata(columnNames, col); columnInfos[col].dataType = GetFetchColumnType(columnMeta); columnInfos[col].columnSize = GetFetchColumnSize(columnMeta); + columnInfos[col].fixedValueSize = + FixedFetchValueSize(columnInfos[col].dataType); columnInfos[col].isLob = std::find(lobColumns.begin(), lobColumns.end(), col + 1) != lobColumns.end(); columnInfos[col].processedColumnSize = columnInfos[col].columnSize; @@ -5204,8 +5171,12 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& if (dataLen < 0) { ThrowStdException("Unexpected negative data length"); } - ValidateFixedFetchDataLength(columnInfos[col - 1].dataType, - static_cast(dataLen)); + const size_t fixedValueSize = columnInfos[col - 1].fixedValueSize; + if (fixedValueSize != 0 && + static_cast(dataLen) != fixedValueSize) { + ThrowStdException( + "Fixed-width data indicator does not match the bound buffer size"); + } // Performance: Use function pointer dispatch for simple types (fast // path) This eliminates the switch statement from hot loop - @@ -6020,23 +5991,27 @@ py::object RunFetchValidationTest(const std::string& scenario) { testGetDataResultIndex = 0; FetchLobColumnDataImpl(nullptr, 1, SQL_C_BINARY, false, true, "", py::none(), false, TestSQLGetData, TestHasUnrelatedWarning); - } else if (scenario == "lob_narrow_terminator_no_progress") { + } else if (scenario == "lob_narrow_nul_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(1)}, + {static_cast(SQL_SUCCESS), static_cast(0)}, }; testGetDataResultIndex = 0; - FetchLobColumnDataImpl(nullptr, 1, SQL_C_CHAR, false, false, "utf-8", - py::none(), false, TestSQLGetDataZeroFill, - TestHasTruncationDiagnostic); - } else if (scenario == "lob_wide_terminator_no_progress") { + py::object value = FetchLobColumnDataImpl( + nullptr, 1, SQL_C_CHAR, false, false, "utf-8", py::none(), false, + TestSQLGetDataZeroFill, TestHasTruncationDiagnostic); + return py::make_tuple(py::len(value), testGetDataResultIndex); + } else if (scenario == "lob_wide_nul_progress") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(sizeof(SQLWCHAR))}, + {static_cast(SQL_SUCCESS), static_cast(0)}, }; testGetDataResultIndex = 0; - FetchLobColumnDataImpl(nullptr, 1, SQL_C_WCHAR, true, false, "utf-16le", - py::none(), false, TestSQLGetDataZeroFill, - TestHasTruncationDiagnostic); + py::object value = FetchLobColumnDataImpl( + nullptr, 1, SQL_C_WCHAR, true, false, "utf-16le", py::none(), false, + TestSQLGetDataZeroFill, TestHasTruncationDiagnostic); + return py::make_tuple(py::len(value), testGetDataResultIndex); } else if (scenario == "unrelated_warning_oversized") { testGetDataResults = { {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(2)}, @@ -6136,6 +6111,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, std::vector dataTypes(numCols); std::vector columnSizes(numCols); + std::vector fixedValueSizes(numCols); std::vector columnNullable(numCols); std::vector columnVarLen(numCols, false); std::vector nullCounts(numCols, 0); @@ -6155,6 +6131,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, dataTypes[i] = dataType; columnSizes[i] = columnSize; + fixedValueSizes[i] = FixedFetchValueSize(dataType); columnNullable[i] = (nullable != SQL_NO_NULLS); if ((dataType == SQL_WVARCHAR || dataType == SQL_WLONGVARCHAR || dataType == SQL_VARCHAR || @@ -6669,7 +6646,11 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, ThrowStdException("Unexpected negative data length."); } auto dataLen = static_cast(indicator); - ValidateFixedFetchDataLength(dataType, dataLen); + if (fixedValueSizes[idxCol] != 0 && + dataLen != fixedValueSizes[idxCol]) { + ThrowStdException( + "Fixed-width data indicator does not match the bound buffer size"); + } switch (dataType) { case SQL_SS_UDT: diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index b8b13aeeb..d18cdcfd5 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -81,8 +81,6 @@ def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): ("oversized_lob_success", "LOB data indicator exceeds the fetch buffer capacity"), ("lob_unrelated_warning_oversized", "LOB data indicator exceeds the fetch buffer capacity"), ("lob_unrelated_warning_no_total", "SQL_NO_TOTAL requires a truncation diagnostic"), - ("lob_narrow_terminator_no_progress", "LOB fetch truncation made no progress"), - ("lob_wide_terminator_no_progress", "LOB fetch truncation made no progress"), ("unrelated_warning_oversized", "data indicator exceeds the fetch buffer capacity"), ("odd_direct_wchar", "invalid byte length"), ("odd_lob_wchar", "invalid byte length"), @@ -139,6 +137,14 @@ def test_lob_unrelated_warning_with_progress_completes(): assert calls == 1 +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +@pytest.mark.parametrize("scenario", ["lob_narrow_nul_progress", "lob_wide_nul_progress"]) +def test_lob_truncation_preserves_nul_payload(scenario): + size, calls = ddbc._test_fetch_validation(scenario) + assert size == 1 + assert calls == 2 + + @pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") @pytest.mark.parametrize("param_set_size", [-1, 0, 1_000_001, True]) def test_executemany_rejects_unsafe_size_before_handle_access(param_set_size):