diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index b89ca350c..7d57e0dc5 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -366,6 +366,90 @@ 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; +} + +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"); + } +} + +void ValidateDecimalDataLength(uint64_t dataLength) { + if (dataLength >= MAX_DIGITS_IN_NUMERIC) { + ThrowStdException("Decimal data exceeds the allocated fetch buffer"); + } +} + +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_TYPE_TIME: + 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"); + } +} + +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; @@ -447,20 +531,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)) { @@ -2339,6 +2409,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") @@ -3584,10 +3672,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 dataSize = 0; @@ -3603,9 +3697,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; @@ -3618,48 +3714,35 @@ 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); + } + const size_t terminatorBytes = isBinary ? 0 : (isWideChar ? sizeof(SQLWCHAR) : 1); + 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"); + } 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 { - // fallback: use full buffer size if actualRead is unknown - bytesRead = DAE_CHUNK_SIZE; + bytesRead = payloadCapacity; } - - // 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); - } - } - } + if (continueForTruncation && bytesRead == 0) { + ThrowStdException("LOB fetch truncation made no progress"); } if (bytesRead > 0) { const size_t previousSize = dataSize; @@ -3667,7 +3750,8 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT 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; } @@ -3717,6 +3801,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) { @@ -3951,6 +4044,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 @@ -4130,6 +4224,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 @@ -4191,6 +4286,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 { @@ -4203,6 +4299,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; @@ -4222,6 +4319,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; @@ -4246,6 +4344,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 @@ -4302,6 +4408,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; @@ -4321,6 +4428,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; @@ -4341,6 +4449,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)); @@ -4355,6 +4464,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 @@ -4376,6 +4486,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; @@ -4400,6 +4511,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, " @@ -4492,6 +4604,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; @@ -4511,6 +4624,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; @@ -4532,6 +4646,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); @@ -4825,6 +4940,32 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& return ret; } +void ValidateFetchedRowCount(const ColumnBuffers& buffers, SQLUSMALLINT numCols, + SQLULEN numRowsFetched) { + 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"); + } + } +} + +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 @@ -4836,6 +4977,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& PERF_TIMER("FetchBatchData"); LOG("FetchBatchData: Fetching data in batches"); SQLRETURN ret; + numRowsFetched = 0; { numRowsFetched = 0; // Release the GIL during the blocking ODBC fetch @@ -4854,6 +4996,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 @@ -4865,6 +5008,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); @@ -4883,6 +5027,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; @@ -5031,6 +5177,15 @@ 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"); + } + 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 - @@ -5056,13 +5211,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"); @@ -5072,6 +5220,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& case SQL_NUMERIC: { try { SQLLEN decimalDataLen = buffers.indicators[col - 1][i]; + ValidateDecimalDataLength(static_cast(decimalDataLen)); const char* rawData = reinterpret_cast( &buffers.charBuffers[col - 1][i * MAX_DIGITS_IN_NUMERIC]); @@ -5491,7 +5640,6 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, py::handl } fetchStateGuard.close(); - return ret; } @@ -5510,7 +5658,13 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, py::handl 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, + SQLGetDataFunc getData = nullptr, + bool (*hasTruncationDiagnostic)(SQLHSTMT) = + HasDataTruncationDiagnostic) { + if (getData == nullptr) { + getData = SQLGetData_ptr; + } size_t start = 0; size_t end = 0; @@ -5528,18 +5682,47 @@ 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( + SQLRETURN ret = getData( 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; + } + + 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())) { + ThrowStdException("Variable-length fetch result is too large"); + } + *indicator = static_cast(prefixBytes); + break; + } // Handle NULL data if (localInd == SQL_NULL_DATA) { @@ -5547,31 +5730,67 @@ 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) { + // SQL_SUCCESS means we got all the data + if (ret == SQL_SUCCESS) { 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; } // 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 < 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"), @@ -5579,6 +5798,12 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, } // The next read starts where the null terminator would have been placed + if (end <= dataVec.size()) { + ThrowStdException("Variable-length fetch truncation made no progress"); + } + if (end > dataVec.max_size()) { + ThrowStdException("Variable-length fetch buffer is too large"); + } start = dataVec.size() - sizeNullTerminator; // Resize buffer for next iteration @@ -5592,6 +5817,259 @@ 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; +} + +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; } + +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); + } 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 == "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 == "short_fixed_indicator") { + 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") { + 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)}, + {static_cast(SQL_SUCCESS), static_cast(1)}, + }; + 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, + TestHasTruncationDiagnostic); + 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, 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") { + 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, TestHasTruncationDiagnostic); + } 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 == "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 == "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 == "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_nul_progress") { + testGetDataResults = { + {static_cast(SQL_SUCCESS_WITH_INFO), static_cast(1)}, + {static_cast(SQL_SUCCESS), static_cast(0)}, + }; + testGetDataResultIndex = 0; + 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; + 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)}, + }; + 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") { + 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)}, + }; + 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, TestHasTruncationDiagnostic); + } 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; + const SQLRETURN ret = GetDataVar(nullptr, 1, SQL_C_BINARY, buffer, &indicator, + 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"); + } + 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) @@ -5641,6 +6119,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, bool hasLobColumns = false; 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); @@ -5659,6 +6139,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, SQLSMALLINT nullable = colMeta["Nullable"].cast(); 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 || @@ -5671,6 +6153,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); @@ -5872,10 +6356,13 @@ 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); } + numRowsFetched = 0; { // Release GIL during the blocking ODBC fetch py::gil_scoped_release release; @@ -5892,11 +6379,12 @@ 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)); + ValidateArrowFetchedRowCount(numRowsFetched, currentFetchSize, spaceLeftInArrowBatch); 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"); @@ -6167,25 +6655,32 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, ThrowStdException("Unexpected negative data length."); } auto dataLen = static_cast(indicator); + 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: 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]; if (dataLen > 0) { + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); 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); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -6197,6 +6692,31 @@ 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 + ValidateArrowTextPayloadLength(fetchBufferSize, dataLen); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.charBuffers[idxCol], idxRowSql, fetchBufferSize, dataLen); auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; if (dataLen > 0) { @@ -6205,13 +6725,6 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, 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); std::memcpy(&(*target_vec)[start], &buffers.charBuffers[idxCol][sourceOffset], dataLen); @@ -6229,15 +6742,21 @@ 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"); + ValidateArrowTextPayloadLength(fetchBufferSize, dataLen); 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)); @@ -6262,7 +6781,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( @@ -6313,7 +6832,7 @@ 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); + ValidateDecimalDataLength(dataLen); Int128_t decimalValue(0, 0); auto start = idxRowSql * MAX_DIGITS_IN_NUMERIC; int sign = 1; @@ -6764,7 +7283,6 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, } fetchStateGuard.close(); - return ret; } @@ -6921,6 +7439,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 41e175af4..13bc468d1 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,98 @@ 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"), + ("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"), + ("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"), + ("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"), + ("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"), + ("lob_unrelated_warning_oversized", "LOB data indicator exceeds the fetch buffer capacity"), + ("lob_unrelated_warning_no_total", "SQL_NO_TOTAL requires a truncation diagnostic"), + ("unrelated_warning_oversized", "data indicator exceeds the fetch buffer capacity"), + ("odd_direct_wchar", "invalid byte length"), + ("odd_lob_wchar", "invalid byte length"), + ], +) +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") +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") +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") +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") +@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):