diff --git a/mssql_python/cursor.py b/mssql_python/cursor.py index 79436ab2e..bafc19de6 100644 --- a/mssql_python/cursor.py +++ b/mssql_python/cursor.py @@ -14,6 +14,7 @@ import decimal import logging import uuid +import codecs import datetime import warnings from typing import List, Mapping, Union, Any, Optional, Tuple, Sequence, TYPE_CHECKING, Iterable @@ -47,10 +48,30 @@ MAX_INLINE_CHAR: int = ( 4000 # NVARCHAR/VARCHAR inline limit; this triggers NVARCHAR(MAX)/VARCHAR(MAX) + DAE ) +MAX_INLINE_BINARY: int = 8000 SMALLMONEY_MIN: decimal.Decimal = decimal.Decimal("-214748.3648") SMALLMONEY_MAX: decimal.Decimal = decimal.Decimal("214748.3647") MONEY_MIN: decimal.Decimal = decimal.Decimal("-922337203685477.5808") MONEY_MAX: decimal.Decimal = decimal.Decimal("922337203685477.5807") +# Bound each native fetch allocation; Arrow initially reserves 42 bytes per variable-width row. +MAX_NATIVE_ROW_COUNT: int = 1_000_000 +MAX_NATIVE_PARAMETER_SIZE: int = 256 * 1024 * 1024 + + +def _encoded_length_exceeds(value: str, encoding: str, limit: int) -> bool: + encoder = codecs.getincrementalencoder(encoding)(errors="strict") + total = 0 + chunk_size = 4096 + for offset in range(0, len(value), chunk_size): + end = min(offset + chunk_size, len(value)) + total += len(encoder.encode(value[offset:end], final=end == len(value))) + if total > limit: + return True + if not value: + total += len(encoder.encode("", final=True)) + return total > limit + + # SQL BIGINT is a signed 64-bit integer. Ints outside this range have no BIGINT # encoding and must be rejected at detect time on both paths (see _map_sql_type). BIGINT_MIN: int = -(2**63) @@ -406,9 +427,7 @@ def __init__(self, connection: "Connection", timeout: int = 0) -> None: ] ] = None self.rowcount: int = -1 - self.arraysize: int = ( - 1 # Default number of rows to fetch at a time is 1, user can change it - ) + self.arraysize = 1 self.buffer_length: int = 1024 # Default buffer length for string data self._result_set_empty: bool = False # Add this initialization self.last_executed_stmt: str = "" # Stores the last statement executed by this cursor @@ -962,7 +981,7 @@ def _map_sql_type( # pylint: disable=too-many-arguments,too-many-positional-arg if isinstance(param, (bytes, bytearray)): length = len(param) - if length > 8000: # Use VARBINARY(MAX) for large blobs + if length > MAX_INLINE_BINARY: # Use VARBINARY(MAX) for large blobs return ( ddbc_sql_const.SQL_VARBINARY.value, ddbc_sql_const.SQL_C_BINARY.value, @@ -1170,6 +1189,25 @@ def _ensure_pyarrow(self) -> Any: "pyarrow is required for Arrow fetch methods. Please install pyarrow." ) from e + @staticmethod + def _validate_native_row_count(value: int, name: str, allow_zero: bool) -> int: + minimum = 0 if allow_zero else 1 + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"{name} must be an integer") + if value < minimum or value > MAX_NATIVE_ROW_COUNT: + raise ValueError( + f"{name} must be between {minimum} and {MAX_NATIVE_ROW_COUNT}, got {value}" + ) + return value + + @property + def arraysize(self) -> int: + return self._arraysize + + @arraysize.setter + def arraysize(self, value: int) -> None: + self._arraysize = self._validate_native_row_count(value, "arraysize", allow_zero=False) + def setinputsizes(self, sizes: List[Union[int, tuple]]) -> None: """ Sets the type information to be used for parameters in execute and executemany. @@ -1230,7 +1268,11 @@ def setinputsizes(self, sizes: List[Union[int, tuple]]) -> None: sql_type, column_size, decimal_digits = size_info # Validate SQL type - if not isinstance(sql_type, int) or sql_type not in valid_sql_types: + if ( + isinstance(sql_type, bool) + or not isinstance(sql_type, int) + or sql_type not in valid_sql_types + ): raise ValueError( f"Invalid SQL type: {sql_type}. Must be a valid SQL type constant." ) @@ -1238,12 +1280,29 @@ def setinputsizes(self, sizes: List[Union[int, tuple]]) -> None: sql_type = ODBC3_TEMPORAL_SQL_TYPES.get(sql_type, sql_type) # Validate size and precision - if not isinstance(column_size, int) or column_size < 0: + if ( + isinstance(column_size, bool) + or not isinstance(column_size, int) + or column_size < 0 + or ( + sql_type + not in ( + ddbc_sql_const.SQL_DECIMAL.value, + ddbc_sql_const.SQL_NUMERIC.value, + ) + and column_size > MAX_NATIVE_PARAMETER_SIZE + ) + ): raise ValueError( - f"Invalid column size: {column_size}. Must be a non-negative integer." + f"Invalid column size: {column_size}. Must be a non-negative integer " + f"no greater than {MAX_NATIVE_PARAMETER_SIZE}." ) - if not isinstance(decimal_digits, int) or decimal_digits < 0: + if ( + isinstance(decimal_digits, bool) + or not isinstance(decimal_digits, int) + or decimal_digits < 0 + ): raise ValueError( f"Invalid decimal digits: {decimal_digits}. " f"Must be a non-negative integer." @@ -1262,7 +1321,11 @@ def setinputsizes(self, sizes: List[Union[int, tuple]]) -> None: sql_type = size_info # Validate SQL type - if not isinstance(sql_type, int) or sql_type not in valid_sql_types: + if ( + isinstance(sql_type, bool) + or not isinstance(sql_type, int) + or sql_type not in valid_sql_types + ): raise ValueError( f"Invalid SQL type: {sql_type}. Must be a valid SQL type constant." ) @@ -2535,6 +2598,7 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s param_info = ddbc_bindings.ParamInfo parameters_type = [] any_dae = False + encoding_settings = self._get_encoding_settings() # Check if we have explicit input sizes set if self._inputsizes: @@ -2549,6 +2613,7 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s # Prepare parameter type information with perf_phase("py::executemany::param_type_detection"): for col_index in range(param_count): + requires_row_fallback = False column = ( [row[col_index] for row in seq_of_parameters] if hasattr(seq_of_parameters, "__getitem__") @@ -2570,12 +2635,48 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s ): c_type = ddbc_sql_const.SQL_C_CHAR.value - # Check if this should be a DAE (data at execution) parameter based on column size - if sample_value is not None: - if isinstance(sample_value, str) and column_size > MAX_INLINE_CHAR: - is_dae = True - elif isinstance(sample_value, (bytes, bytearray)) and column_size > 8000: - is_dae = True + # Text values stay on the driver's wide binding path. Bytes supplied + # for narrow SQL text types use real SQL_C_CHAR in scalar execution. + narrow_text_type = sql_type in ( + ddbc_sql_const.SQL_CHAR.value, + ddbc_sql_const.SQL_VARCHAR.value, + ddbc_sql_const.SQL_LONGVARCHAR.value, + ) + wide_text_type = sql_type in ( + ddbc_sql_const.SQL_WCHAR.value, + ddbc_sql_const.SQL_WVARCHAR.value, + ddbc_sql_const.SQL_WLONGVARCHAR.value, + ddbc_sql_const.SQL_SS_XML.value, + ) + if narrow_text_type or wide_text_type: + text_values = [value for value in column if isinstance(value, str)] + binary_values = [ + value for value in column if isinstance(value, (bytes, bytearray)) + ] + if c_type == ddbc_sql_const.SQL_CHAR.value: + text_is_large = any( + _encoded_length_exceeds( + value, + encoding_settings["encoding"], + MAX_INLINE_CHAR, + ) + for value in text_values + ) + else: + text_is_large = any( + sum(2 if ord(char) > 0xFFFF else 1 for char in value) + > MAX_INLINE_CHAR + for value in text_values + ) + binary_is_large = narrow_text_type and any( + len(value) > MAX_INLINE_BINARY for value in binary_values + ) + requires_row_fallback = bool( + narrow_text_type and text_values and binary_values + ) + if narrow_text_type and binary_values and not text_values: + c_type = ddbc_sql_const.SQL_CHAR.value + is_dae = text_is_large or binary_is_large # Sanitize precision/scale for numeric types if sql_type in ( @@ -2590,6 +2691,7 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s ddbc_sql_const.SQL_BINARY.value, ddbc_sql_const.SQL_VARBINARY.value, ddbc_sql_const.SQL_LONGVARBINARY.value, + ddbc_sql_const.SQL_SS_UDT.value, ): # Find the maximum size needed for any row's binary data max_binary_size = 0 @@ -2599,14 +2701,18 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s max_binary_size = max(max_binary_size, len(value)) # For SQL Server VARBINARY(MAX), we need to use large object binding - if column_size > 8000 or max_binary_size > 8000: - sql_type = ddbc_sql_const.SQL_LONGVARBINARY.value + if max_binary_size > MAX_INLINE_BINARY: + if sql_type != ddbc_sql_const.SQL_SS_UDT.value: + sql_type = ddbc_sql_const.SQL_LONGVARBINARY.value is_dae = True # Update column_size to actual maximum size if it's larger # Always ensure at least a minimum size of 1 for empty strings column_size = max(max_binary_size, 1) + if is_dae: + column_size = 0 + paraminfo = param_info() paraminfo.paramCType = c_type paraminfo.paramSQLType = sql_type @@ -2685,7 +2791,7 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s max_binary_size = max(max_binary_size, len(value)) # For SQL Server VARBINARY(MAX), we need to use large object binding - if max_binary_size > 8000: + if max_binary_size > MAX_INLINE_BINARY: paraminfo.paramSQLType = ddbc_sql_const.SQL_LONGVARBINARY.value paraminfo.isDAE = True @@ -2694,7 +2800,7 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s paraminfo.columnSize = max(max_binary_size, 1) parameters_type.append(paraminfo) - if paraminfo.isDAE: + if paraminfo.isDAE or requires_row_fallback: any_dae = True if any_dae: @@ -2771,9 +2877,6 @@ def executemany( # pylint: disable=too-many-locals,too-many-branches,too-many-s processed_parameters ) - # Get encoding settings - encoding_settings = self._get_encoding_settings() - # Debug logging: emit batch metadata only. Never log parameter values or # row representations here -- rows may contain PII (SSNs, emails, # balances) that would leak into log files and APM/log shippers even @@ -2918,8 +3021,10 @@ def fetchmany(self, size: Optional[int] = None) -> List[Row]: if size is None: size = self.arraysize + else: + size = self._validate_native_row_count(size, "size", allow_zero=True) - if size <= 0: + if size == 0: return [] if self._cached_decoding_generation != self._connection._decoding_generation: @@ -3073,6 +3178,7 @@ def arrow_batch(self, batch_size: int = 8192) -> "pyarrow.RecordBatch": """ self._check_closed() # Check if the cursor is closed pyarrow = self._ensure_pyarrow() + batch_size = self._validate_native_row_count(batch_size, "batch_size", allow_zero=True) if not self._has_result_set and self.description: self._reset_rownumber() @@ -3081,7 +3187,7 @@ def arrow_batch(self, batch_size: int = 8192) -> "pyarrow.RecordBatch": char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) char_c_type = char_decoding.get("ctype", ddbc_sql_const.SQL_WCHAR.value) ret = ddbc_bindings.DDBCSQLFetchArrowBatch( - self.hstmt, capsules, max(batch_size, 0), char_c_type, self.messages + self.hstmt, capsules, batch_size, char_c_type, self.messages ) check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret) @@ -3468,12 +3574,14 @@ def _bulkcopy_core_and_validate(table_name, batch_size, timeout): logger.error("bulkcopy: Invalid table_name parameter") raise ValueError("table_name must be a non-empty string") - if not isinstance(batch_size, int): + if not isinstance(batch_size, int) or isinstance(batch_size, bool): raise TypeError( f"batch_size must be a non-negative integer, got {type(batch_size).__name__}" ) - if batch_size < 0: - raise ValueError(f"batch_size must be non-negative, got {batch_size}") + if batch_size < 0 or batch_size > MAX_NATIVE_ROW_COUNT: + raise ValueError( + f"batch_size must be between 0 and {MAX_NATIVE_ROW_COUNT}, got {batch_size}" + ) if not isinstance(timeout, int) or isinstance(timeout, bool): raise TypeError(f"timeout must be a non-negative integer, got {type(timeout).__name__}") diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 4e4523955..b89ca350c 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -23,6 +23,7 @@ #include // For std::memcpy #include #include +#include #include // std::forward #include // CPython datetime API (PyDateTime_IMPORT, PyDateTime_GET_*, etc.) @@ -325,12 +326,230 @@ ParamType* AllocateParamBuffer(std::vector>& paramBuffers, template ParamType* AllocateParamBufferArray(std::vector>& paramBuffers, size_t count) { + if (count > std::numeric_limits::max() / sizeof(ParamType)) { + ThrowStdException("Parameter buffer size is too large"); + } std::shared_ptr buffer(new ParamType[count], std::default_delete()); ParamType* raw = buffer.get(); paramBuffers.push_back(buffer); return raw; } +size_t CheckedAddSize(size_t left, size_t right, const char* errorMessage) { + if (left > std::numeric_limits::max() - right) { + ThrowStdException(errorMessage); + } + return left + right; +} + +size_t CheckedMultiplySize(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) { + const size_t result = static_cast(columnSize); + if (static_cast(result) != columnSize) { + ThrowStdException("Column size is too large"); + } + return result; +} + +SQLLEN CheckedFetchBufferLength(size_t elementCount, size_t elementSize) { + const size_t byteCount = + CheckedMultiplySize(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); +} + +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; + +void ValidateNativeRowCount(int value, const char* name, bool allowZero) { + const int minimum = allowZero ? 0 : 1; + if (value < minimum || value > MAX_NATIVE_ROW_COUNT) { + ThrowStdException(std::string(name) + " must be between " + std::to_string(minimum) + + " and " + std::to_string(MAX_NATIVE_ROW_COUNT)); + } +} + +int ValidateNativeRowCountArgument(py::handle value, const char* name, bool allowZero) { + if (PyBool_Check(value.ptr())) { + ThrowStdException(std::string(name) + " must be an integer, not bool"); + } + const int converted = value.cast(); + ValidateNativeRowCount(converted, name, allowZero); + return converted; +} + +void ReserveNativeParameterBytes(size_t& reservedBytes, size_t count, size_t elementSize) { + const size_t allocationBytes = + CheckedMultiplySize(count, elementSize, "Parameter buffer size is too large"); + reservedBytes = + CheckedAddSize(reservedBytes, allocationBytes, "Parameter buffer size is too large"); + if (reservedBytes > MAX_NATIVE_PARAMETER_BYTES) { + ThrowStdException("Parameter buffers exceed the 256 MiB allocation limit"); + } +} + +size_t ParameterArrayElementSize(const ParamInfo& info) { + switch (info.paramCType) { + case SQL_C_LONG: + return sizeof(int); + case SQL_C_DOUBLE: + return sizeof(double); + case SQL_C_WCHAR: + return CheckedMultiplySize( + CheckedAddSize(info.columnSize, 1, "Wide-character parameter size is too large"), + sizeof(SQLWCHAR), "Wide-character parameter size is too large"); + case SQL_C_TINYINT: + case SQL_C_UTINYINT: + return sizeof(unsigned char); + case SQL_C_SHORT: + return sizeof(short); + case SQL_C_CHAR: + case SQL_C_BINARY: + return CheckedAddSize(info.columnSize, 1, "Character parameter size is too large"); + case SQL_C_BIT: + return sizeof(char); + case SQL_C_STINYINT: + case SQL_C_USHORT: + return sizeof(unsigned short); + case SQL_C_SBIGINT: + case SQL_C_SLONG: + case SQL_C_UBIGINT: + case SQL_C_ULONG: + return sizeof(int64_t); + case SQL_C_FLOAT: + return sizeof(float); + case SQL_C_TYPE_DATE: + return sizeof(SQL_DATE_STRUCT); + case SQL_C_TYPE_TIME: + return sizeof(SQL_TIME_STRUCT); + case SQL_C_TYPE_TIMESTAMP: + return sizeof(SQL_TIMESTAMP_STRUCT); + case SQL_C_SS_TIMESTAMPOFFSET: + return sizeof(DateTimeOffset); + case SQL_C_NUMERIC: + return sizeof(SQL_NUMERIC_STRUCT); + case SQL_C_GUID: + return sizeof(SQLGUID); + case SQL_C_DEFAULT: + return sizeof(char); + default: + ThrowStdException("Unsupported C type for parameter array allocation"); + } + 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)) { + ThrowStdException(errorMessage); + } + return std::make_unique(count); +} + +void ReserveNativeFetchBytes(size_t& reservedBytes, size_t count, size_t elementSize) { + const size_t allocationBytes = + CheckedMultiplySize(count, elementSize, "Native fetch buffer size is too large"); + const size_t totalBytes = + CheckedAddSize(reservedBytes, allocationBytes, "Native fetch buffer size is too large"); + if (totalBytes > MAX_NATIVE_FETCH_BYTES) { + ThrowStdException("Native fetch buffers exceed the 256 MiB allocation limit"); + } + reservedBytes = totalBytes; +} + +template +void ResizeNativeFetchBuffer(std::vector& buffer, size_t count, + size_t& reservedBytes) { + if (buffer.size() >= count) { + return; + } + + const size_t oldBytes = CheckedMultiplySize( + buffer.size(), sizeof(ElementType), "Native fetch buffer size is too large"); + ReserveNativeFetchBytes(reservedBytes, count, sizeof(ElementType)); + { + std::vector resizedBuffer(count); + std::copy(buffer.begin(), buffer.end(), resizedBuffer.begin()); + buffer.swap(resizedBuffer); + } + reservedBytes -= oldBytes; +} + +template +void EnsureNativeFetchBufferSize(std::vector& buffer, size_t requiredSize, + size_t& reservedBytes) { + if (buffer.size() >= requiredSize) { + return; + } + size_t newSize = std::max(buffer.size(), 1); + while (newSize < requiredSize) { + newSize = CheckedMultiplySize(newSize, 2, "Native fetch buffer size is too large"); + } + ResizeNativeFetchBuffer(buffer, newSize, reservedBytes); +} + +size_t ExtendNativeFetchBuffer(std::vector& buffer, size_t currentSize, size_t appendSize, + size_t& reservedBytes) { + const size_t requiredSize = + CheckedAddSize(currentSize, appendSize, "Native fetch buffer size is too large"); + EnsureNativeFetchBufferSize(buffer, requiredSize, reservedBytes); + return requiredSize; +} + +py::tuple TestLobBufferGrowth(py::handle chunkCountArg, py::handle chunkSizeArg) { + const size_t chunkCount = static_cast( + ValidateNativeRowCountArgument(chunkCountArg, "LOB growth chunk count", false)); + if (!PyLong_Check(chunkSizeArg.ptr()) || PyBool_Check(chunkSizeArg.ptr())) { + ThrowStdException("LOB growth chunk size must be a positive integer"); + } + const size_t chunkSize = chunkSizeArg.cast(); + if (chunkSize == 0 || chunkSize > MAX_NATIVE_FETCH_BYTES) { + ThrowStdException("LOB growth chunk size must be between 1 and 256 MiB"); + } + std::vector buffer; + size_t dataSize = 0; + size_t reservedBytes = 0; + size_t growthCount = 0; + for (size_t chunk = 0; chunk < chunkCount; ++chunk) { + const size_t previousAllocation = buffer.size(); + dataSize = ExtendNativeFetchBuffer(buffer, dataSize, chunkSize, reservedBytes); + if (buffer.size() != previousAllocation) { + ++growthCount; + } + } + return py::make_tuple(growthCount, dataSize, buffer.size(), reservedBytes); +} + +template +std::unique_ptr AllocateArrowArray(size_t count, size_t& reservedBytes, + const char* errorMessage) { + ReserveNativeFetchBytes(reservedBytes, count, sizeof(ElementType)); + return AllocateUniqueArray(count, errorMessage); +} + std::string DescribeChar(unsigned char ch) { if (ch >= 32 && ch <= 126) { return std::string("'") + static_cast(ch) + "'"; @@ -347,6 +566,7 @@ std::string DescribeChar(unsigned char ch) { template // The callable hides whether the caller wraps SQLPutData with GIL management; chunk sizing stays shared. static SQLRETURN stream_dae_chunks(const void* data, size_t total_bytes, PutDataFn put_data_fn) { + if (total_bytes == 0) return put_data_fn(nullptr, 0); const char* bytes = static_cast(data); for (size_t offset = 0; offset < total_bytes; offset += DAE_CHUNK_SIZE) { size_t len = std::min(static_cast(DAE_CHUNK_SIZE), total_bytes - offset); @@ -357,6 +577,88 @@ static SQLRETURN stream_dae_chunks(const void* data, size_t total_bytes, PutData return SQL_SUCCESS; } +template +static SQLRETURN stream_unicode_dae_chunks(PyObject* value, const std::string& encoding, + PutDataFn put_data_fn) { + py::object encoderFactory = + py::module_::import("codecs").attr("getincrementalencoder")(encoding); + py::object encoder = encoderFactory("strict"); + const Py_ssize_t length = PyUnicode_GET_LENGTH(value); + constexpr Py_ssize_t chunkSize = 4096; + for (Py_ssize_t offset = 0; offset < length; offset += chunkSize) { + const Py_ssize_t end = std::min(offset + chunkSize, length); + py::object chunk = steal(PyUnicode_Substring(value, offset, end)); + if (!chunk) throw py::error_already_set(); + py::object encoded = encoder.attr("encode")(chunk, end == length); + char* data = nullptr; + Py_ssize_t size = 0; + if (PyBytes_AsStringAndSize(encoded.ptr(), &data, &size) != 0) { + throw py::error_already_set(); + } + if (size > 0) { + SQLRETURN rc = put_data_fn(data, static_cast(size)); + if (!SQL_SUCCEEDED(rc)) return rc; + } + } + if (length == 0) { + py::object encoded = encoder.attr("encode")(py::str(), true); + char* data = nullptr; + Py_ssize_t size = 0; + if (PyBytes_AsStringAndSize(encoded.ptr(), &data, &size) != 0) { + throw py::error_already_set(); + } + if (size > 0) return put_data_fn(data, static_cast(size)); + return put_data_fn(nullptr, 0); + } + return SQL_SUCCESS; +} + +static SQLRETURN StreamDAEParameter(SQLHSTMT hStmt, const ParamInfo& info, + const std::string& charEncoding) { + PyObject* value = info.dataPtr.ptr(); + if (!value || value == Py_None) { + py::gil_scoped_release release; + return SQLPutData_ptr(hStmt, nullptr, 0); + } + + auto putImmutableData = [&](SQLPOINTER data, SQLLEN length) { + py::gil_scoped_release release; + return SQLPutData_ptr(hStmt, data, length); + }; + if (PyUnicode_Check(value)) { + if (info.paramCType == SQL_C_WCHAR) { + return stream_unicode_dae_chunks(value, "utf-16-le", putImmutableData); + } + if (info.paramCType == SQL_C_CHAR) { + return stream_unicode_dae_chunks(value, charEncoding, putImmutableData); + } + ThrowStdException("DAE only supports text C types for str values"); + } + if (PyBytes_Check(value)) { + return stream_dae_chunks(PyBytes_AS_STRING(value), + static_cast(PyBytes_GET_SIZE(value)), + putImmutableData); + } + if (PyByteArray_Check(value)) { + const size_t totalBytes = static_cast(PyByteArray_GET_SIZE(value)); + if (totalBytes == 0) return putImmutableData(nullptr, 0); + std::vector chunk(std::min(static_cast(DAE_CHUNK_SIZE), totalBytes)); + for (size_t offset = 0; offset < totalBytes; offset += chunk.size()) { + const size_t currentSize = static_cast(PyByteArray_GET_SIZE(value)); + const size_t length = std::min(chunk.size(), totalBytes - offset); + if (currentSize < offset + length) { + ThrowStdException("bytearray changed size during DAE streaming"); + } + std::copy_n(PyByteArray_AS_STRING(value) + offset, length, chunk.data()); + SQLRETURN rc = putImmutableData(chunk.data(), static_cast(length)); + if (!SQL_SUCCEEDED(rc)) return rc; + } + return SQL_SUCCESS; + } + ThrowStdException("DAE only supports str, bytes, or bytearray values"); + return SQL_ERROR; +} + // GH-610: Resolve SQL type for a NULL parameter using per-handle cache. // On cache miss, calls SQLDescribeParam and stores the result. static DescribedParamInfo ResolveNullParamType(SqlHandle& handle, SQLHANDLE hStmt, int paramIndex) { @@ -502,6 +804,10 @@ SQLRETURN BindParameters(SqlHandle& handle, SQLHANDLE hStmt, const py::list& par "with %zu parameters", (void*)hStmt, params.size()); + if (params.size() != paramInfos.size()) { + ThrowStdException("Parameter count does not match parameter metadata count"); + } + SQLRETURN describeRc = PreResolveUdtTypes(hStmt, paramInfos); if (!SQL_SUCCEEDED(describeRc)) return describeRc; // GH-627: resolve unknown NULL param SQL types before binding any param. @@ -509,6 +815,9 @@ SQLRETURN BindParameters(SqlHandle& handle, SQLHANDLE hStmt, const py::list& par for (int paramIndex = 0; paramIndex < params.size(); paramIndex++) { const auto& param = params[paramIndex]; ParamInfo& paramInfo = paramInfos[paramIndex]; + if (paramInfo.isDAE) { + paramInfo.dataPtr = py::reinterpret_borrow(param); + } LOG("BindParameters: Processing param[%d] - C_Type=%d, SQL_Type=%d, " "ColumnSize=%lu, DecimalDigits=%d, InputOutputType=%d", paramIndex, paramInfo.paramCType, paramInfo.paramSQLType, @@ -581,8 +890,7 @@ SQLRETURN BindParameters(SqlHandle& handle, SQLHANDLE hStmt, const py::list& par break; } case SQL_C_BINARY: { - if (!py::isinstance(param) && !py::isinstance(param) && - !py::isinstance(param)) { + if (!py::isinstance(param) && !py::isinstance(param)) { ThrowStdException(MakeParamMismatchErrorStr(paramInfo.paramCType, paramIndex)); } if (paramInfo.isDAE) { @@ -616,8 +924,7 @@ SQLRETURN BindParameters(SqlHandle& handle, SQLHANDLE hStmt, const py::list& par break; } case SQL_C_WCHAR: { - if (!py::isinstance(param) && !py::isinstance(param) && - !py::isinstance(param)) { + if (!py::isinstance(param)) { ThrowStdException(MakeParamMismatchErrorStr(paramInfo.paramCType, paramIndex)); } if (paramInfo.isDAE) { @@ -2121,6 +2428,7 @@ SQLRETURN SQLExecute_wrap(const SqlHandlePtr statementHandle, bool use_prepare, const py::dict& encoding_settings) { PERF_TIMER("SQLExecute_wrap"); + ValidateInputSizes(input_sizes.ptr()); if (!statementHandle || !statementHandle->get()) { return SQL_INVALID_HANDLE; } @@ -2156,7 +2464,8 @@ SQLRETURN SQLExecute_wrap(const SqlHandlePtr statementHandle, // Run DetectParamTypes BEFORE SQLPrepare so that type-detection errors // (unsupported type, NaN Decimal, precision overflow) don't leave the // cursor in a half-prepared state. - std::vector paramInfos = DetectParamTypes(params.ptr(), input_sizes.ptr()); + std::vector paramInfos = + DetectParamTypes(params.ptr(), input_sizes.ptr(), charEncoding); RETCODE rc; bool already_prepared = is_stmt_prepared[0].cast(); @@ -2195,10 +2504,6 @@ SQLRETURN SQLExecute_wrap(const SqlHandlePtr statementHandle, // GIL is released around each ODBC call to match slow-path concurrency. if (rc == SQL_NEED_DATA) { SQLPOINTER paramToken = nullptr; - auto putData = [&](SQLPOINTER data, SQLLEN len) { - py::gil_scoped_release release; - return SQLPutData_ptr(hStmt, data, len); - }; while (true) { { py::gil_scoped_release release; @@ -2217,58 +2522,8 @@ SQLRETURN SQLExecute_wrap(const SqlHandlePtr statementHandle, if (matchedInfo < first || matchedInfo >= last) { ThrowStdException("SQLExecute: unrecognized paramToken from SQLParamData"); } - PyObject* pyObj = matchedInfo->dataPtr.ptr(); - if (!pyObj || pyObj == Py_None) { - py::gil_scoped_release release; - SQLPutData_ptr(hStmt, nullptr, 0); - continue; - } - - if (PyUnicode_Check(pyObj)) { - if (matchedInfo->paramCType == SQL_C_WCHAR) { - std::u16string u16 = - borrow(pyObj).cast(); - rc = stream_dae_chunks( - reinterpretU16stringAsSqlWChar(u16), - u16.size() * sizeof(SQLWCHAR), - putData); - if (!SQL_SUCCEEDED(rc)) return rc; - } else if (matchedInfo->paramCType == SQL_C_CHAR) { - std::string encodedStr; - py::object encoded = borrow(pyObj) - .attr("encode")(charEncoding, "strict"); - encodedStr = encoded.cast(); - rc = stream_dae_chunks(encodedStr.data(), encodedStr.size(), putData); - if (!SQL_SUCCEEDED(rc)) return rc; - } else { - ThrowStdException("SQLExecute: unsupported C type for str in DAE"); - } - } else if (PyBytes_Check(pyObj) || PyByteArray_Check(pyObj)) { - // matchedInfo->dataPtr holds a strong ref to pyObj for the whole loop. - const char* dataPtr = nullptr; - size_t totalBytes = 0; - std::string bytesStorage; // only used for the bytearray copy below - - if (PyBytes_Check(pyObj)) { - // bytes is immutable and kept alive by the strong ref above, so stream - // straight from its internal buffer with no copy. This is the large-blob - // DAE path, so skipping a full payload copy is the whole point. - dataPtr = PyBytes_AS_STRING(pyObj); - totalBytes = static_cast(PyBytes_GET_SIZE(pyObj)); - } else { - // bytearray is mutable and the GIL is released mid-stream, so copy to a - // stable buffer before streaming. - bytesStorage.assign(PyByteArray_AS_STRING(pyObj), - static_cast(PyByteArray_GET_SIZE(pyObj))); - dataPtr = bytesStorage.data(); - totalBytes = bytesStorage.size(); - } - - rc = stream_dae_chunks(dataPtr, totalBytes, putData); - if (!SQL_SUCCEEDED(rc)) return rc; - } else { - ThrowStdException("SQLExecute: DAE only supported for str or bytes"); - } + rc = StreamDAEParameter(hStmt, *matchedInfo, charEncoding); + if (!SQL_SUCCEEDED(rc)) return rc; } if (!SQL_SUCCEEDED(rc) && rc != SQL_NO_DATA) return rc; } @@ -2290,6 +2545,9 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& LOG("BindParameterArray: Starting column-wise array binding - " "param_count=%zu, param_set_size=%zu", columnwise_params.size(), paramSetSize); + if (columnwise_params.size() != paramInfos.size()) { + ThrowStdException("Parameter count does not match parameter metadata count"); + } std::vector> tempBuffers; @@ -2298,6 +2556,12 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& if (!SQL_SUCCEEDED(describeRc)) return describeRc; // GH-627: resolve unknown NULL array param SQL types before binding any param. PreResolveUnknownNullTypes(handle, hStmt, paramInfos); + size_t reservedParameterBytes = 0; + for (const ParamInfo& info : paramInfos) { + ReserveNativeParameterBytes(reservedParameterBytes, paramSetSize, + ParameterArrayElementSize(info)); + ReserveNativeParameterBytes(reservedParameterBytes, paramSetSize, sizeof(SQLLEN)); + } for (int paramIndex = 0; paramIndex < columnwise_params.size(); ++paramIndex) { const py::list& columnValues = columnwise_params[paramIndex].cast(); ParamInfo& info = paramInfos[paramIndex]; @@ -2359,8 +2623,18 @@ 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( + info.columnSize, 1, "Wide-character parameter size is too large"); + const size_t bufferBytes = CheckedMultiplySize( + elementWidth, sizeof(SQLWCHAR), + "Wide-character parameter length is too large"); + if (bufferBytes > static_cast(std::numeric_limits::max())) { + ThrowStdException("Wide-character parameter length is too large"); + } SQLWCHAR* wcharArray = AllocateParamBufferArray( - tempBuffers, paramSetSize * (info.columnSize + 1)); + tempBuffers, + CheckedMultiplySize(paramSetSize, elementWidth, + "Wide-character parameter buffer is too large")); strLenOrIndArray = AllocateParamBufferArray(tempBuffers, paramSetSize); for (size_t i = 0; i < paramSetSize; ++i) { if (columnValues[i].is_none()) { @@ -2368,16 +2642,41 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& std::memset(wcharArray + i * (info.columnSize + 1), 0, (info.columnSize + 1) * sizeof(SQLWCHAR)); } else { - std::u16string wstr = columnValues[i].cast(); - // u16string is already UTF-16, so the - // original check is sufficient - if (wstr.length() > info.columnSize) { + if (!py::isinstance(columnValues[i])) { + ThrowStdException( + MakeParamMismatchErrorStr(info.paramCType, paramIndex)); + } + if (static_cast( + PyUnicode_GET_LENGTH(columnValues[i].ptr())) > + info.columnSize) { ThrowStdException("Input string exceeds allowed column size " "at parameter index " + std::to_string(paramIndex)); } - std::memcpy(wcharArray + i * (info.columnSize + 1), wstr.c_str(), - (wstr.length() + 1) * sizeof(SQLWCHAR)); + py::object encoded = + columnValues[i].attr("encode")("utf-16-le", "strict"); + char* encodedData = nullptr; + Py_ssize_t encodedSize = 0; + if (PyBytes_AsStringAndSize(encoded.ptr(), &encodedData, + &encodedSize) != 0) { + throw py::error_already_set(); + } + if (encodedSize < 0 || + static_cast(encodedSize) % sizeof(SQLWCHAR) != 0) { + ThrowStdException( + "Wide-character parameter has an invalid byte length"); + } + const size_t wcharLength = + static_cast(encodedSize) / sizeof(SQLWCHAR); + if (wcharLength > info.columnSize) { + ThrowStdException("Input string exceeds allowed column size " + "at parameter index " + + std::to_string(paramIndex)); + } + SQLWCHAR* destination = wcharArray + i * elementWidth; + std::copy_n(encodedData, static_cast(encodedSize), + reinterpret_cast(destination)); + destination[wcharLength] = 0; strLenOrIndArray[i] = SQL_NTS; } } @@ -2385,7 +2684,7 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& "param_index=%d", paramIndex); dataPtr = wcharArray; - bufferLength = (info.columnSize + 1) * sizeof(SQLWCHAR); + bufferLength = static_cast(bufferBytes); break; } case SQL_C_TINYINT: @@ -2456,8 +2755,15 @@ 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( + 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, paramSetSize * (info.columnSize + 1)); + tempBuffers, + CheckedMultiplySize(paramSetSize, elementWidth, + "Character parameter buffer is too large")); strLenOrIndArray = AllocateParamBufferArray(tempBuffers, paramSetSize); for (size_t i = 0; i < paramSetSize; ++i) { if (columnValues[i].is_none()) { @@ -2465,19 +2771,30 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& std::memset(charArray + i * (info.columnSize + 1), 0, info.columnSize + 1); } else { - std::string encodedStr; - + if (info.paramCType == SQL_C_BINARY && + !py::isinstance(columnValues[i]) && + !py::isinstance(columnValues[i])) { + ThrowStdException(MakeParamMismatchErrorStr(info.paramCType, + paramIndex)); + } + py::object encoded; + char* encodedData = nullptr; + Py_ssize_t encodedSize = 0; if (py::isinstance(columnValues[i])) { // Use Python's codec system to encode the string with specified // encoding try { - py::object encoded = + encoded = columnValues[i].attr("encode")(charEncoding, "strict"); - encodedStr = encoded.cast(); + if (PyBytes_AsStringAndSize(encoded.ptr(), &encodedData, + &encodedSize) != 0) { + throw py::error_already_set(); + } LOG("BindParameterArray: param[%d] row[%zu] SQL_C_CHAR - " "Encoded with '%s', " "size=%zu bytes", - paramIndex, i, charEncoding.c_str(), encodedStr.size()); + paramIndex, i, charEncoding.c_str(), + static_cast(encodedSize)); } catch (const py::error_already_set& e) { LOG_ERROR("BindParameterArray: param[%d] row[%zu] SQL_C_CHAR - " "Failed to encode " @@ -2488,29 +2805,41 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& std::to_string(paramIndex) + " row " + std::to_string(i) + " with encoding '" + charEncoding + "': " + e.what()); } + } else if (py::isinstance(columnValues[i])) { + encoded = py::reinterpret_borrow(columnValues[i]); + if (PyBytes_AsStringAndSize(encoded.ptr(), &encodedData, + &encodedSize) != 0) { + throw py::error_already_set(); + } + } else if (py::isinstance(columnValues[i])) { + encodedData = PyByteArray_AsString(columnValues[i].ptr()); + encodedSize = PyByteArray_Size(columnValues[i].ptr()); + if (encodedData == nullptr || encodedSize < 0) { + throw py::error_already_set(); + } } else { - // bytes/bytearray - use as-is (already encoded) - encodedStr = columnValues[i].cast(); + ThrowStdException( + MakeParamMismatchErrorStr(info.paramCType, paramIndex)); } - if (encodedStr.size() > info.columnSize) { + const size_t dataSize = static_cast(encodedSize); + if (dataSize > info.columnSize) { LOG("BindParameterArray: String/binary too " "long - param_index=%d, row=%zu, size=%zu, " "max=%zu", - paramIndex, i, encodedStr.size(), info.columnSize); + paramIndex, i, dataSize, info.columnSize); ThrowStdException("Input exceeds column size at index " + std::to_string(i)); } - std::memcpy(charArray + i * (info.columnSize + 1), encodedStr.c_str(), - encodedStr.size()); - strLenOrIndArray[i] = static_cast(encodedStr.size()); + std::copy_n(encodedData, dataSize, charArray + i * elementWidth); + strLenOrIndArray[i] = static_cast(dataSize); } } LOG("BindParameterArray: SQL_C_CHAR/BINARY bound - " "param_index=%d", paramIndex); dataPtr = charArray; - bufferLength = info.columnSize + 1; + bufferLength = static_cast(elementWidth); break; } case SQL_C_BIT: { @@ -2922,9 +3251,14 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& SQLRETURN SQLExecuteMany_wrap(const SqlHandlePtr statementHandle, const std::u16string& query, const py::list& columnwise_params, - std::vector& paramInfos, size_t paramSetSize, + std::vector& paramInfos, py::handle paramSetSizeArg, const py::dict& encodingSettings) { PERF_TIMER("SQLExecuteMany_wrap"); + const size_t paramSetSize = static_cast( + ValidateNativeRowCountArgument(paramSetSizeArg, "Parameter set size", false)); + if (!statementHandle || !statementHandle->get()) { + return SQL_INVALID_HANDLE; + } statementHandle->resultMetadata.clear(); LOG("SQLExecuteMany: Starting batch execution - param_count=%zu, " "param_set_size=%zu", @@ -2992,14 +3326,34 @@ SQLRETURN SQLExecuteMany_wrap(const SqlHandlePtr statementHandle, const std::u16 return rc; } else { LOG("SQLExecuteMany: Using DAE (data-at-execution) - row_count=%zu", - columnwise_params.size()); - size_t rowCount = columnwise_params.size(); + paramSetSize); + if (columnwise_params.size() != paramInfos.size()) { + ThrowStdException("Parameter count does not match parameter metadata count"); + } + rc = SQLSetStmtAttr_ptr(hStmt, SQL_ATTR_PARAMSET_SIZE, (SQLPOINTER)1, 0); + if (!SQL_SUCCEEDED(rc)) return rc; + size_t rowCount = paramSetSize; for (size_t rowIndex = 0; rowIndex < rowCount; ++rowIndex) { LOG("SQLExecuteMany: Processing DAE row %zu of %zu", rowIndex + 1, rowCount); - py::list rowParams = columnwise_params[rowIndex]; + py::list rowParams; + for (py::handle column : columnwise_params) { + py::sequence values = py::reinterpret_borrow(column); + if (static_cast(py::len(values)) != paramSetSize) { + ThrowStdException("Parameter column length does not match parameter set size"); + } + rowParams.append(values[rowIndex]); + } + std::vector rowParamInfos = paramInfos; + for (size_t paramIndex = 0; paramIndex < rowParamInfos.size(); ++paramIndex) { + if (rowParams[paramIndex].is_none()) { + rowParamInfos[paramIndex].paramCType = SQL_C_DEFAULT; + rowParamInfos[paramIndex].isDAE = false; + rowParamInfos[paramIndex].dataPtr = py::none(); + } + } std::vector> paramBuffers; - rc = BindParameters(*statementHandle, hStmt, rowParams, paramInfos, + rc = BindParameters(*statementHandle, hStmt, rowParams, rowParamInfos, paramBuffers, charEncoding); if (!SQL_SUCCEEDED(rc)) { LOG("SQLExecuteMany: BindParameters failed for row %zu - rc=%d", rowIndex, rc); @@ -3030,48 +3384,17 @@ SQLRETURN SQLExecuteMany_wrap(const SqlHandlePtr statementHandle, const std::u16 dae_chunk_count, rc); return rc; } - - py::object* py_obj_ptr = reinterpret_cast(token); - if (!py_obj_ptr) { - LOG("SQLExecuteMany: NULL token pointer in DAE - chunk=%zu", dae_chunk_count); - return SQL_ERROR; - } - - if (py::isinstance(*py_obj_ptr)) { - std::string data = py_obj_ptr->cast(); - SQLLEN data_len = static_cast(data.size()); - LOG("SQLExecuteMany: Sending string DAE data - chunk=%zu, " - "length=%lld", - dae_chunk_count, static_cast(data_len)); - rc = [&] { - py::gil_scoped_release release; - return SQLPutData_ptr(hStmt, (SQLPOINTER)data.c_str(), data_len); - }(); - if (!SQL_SUCCEEDED(rc) && rc != SQL_NEED_DATA) { - LOG("SQLExecuteMany: SQLPutData(string) failed - " - "chunk=%zu, rc=%d", - dae_chunk_count, rc); - } - } else if (py::isinstance(*py_obj_ptr) || - py::isinstance(*py_obj_ptr)) { - std::string data = py_obj_ptr->cast(); - SQLLEN data_len = static_cast(data.size()); - LOG("SQLExecuteMany: Sending bytes/bytearray DAE data - " - "chunk=%zu, length=%lld", - dae_chunk_count, static_cast(data_len)); - rc = [&] { - py::gil_scoped_release release; - return SQLPutData_ptr(hStmt, (SQLPOINTER)data.c_str(), data_len); - }(); - if (!SQL_SUCCEEDED(rc) && rc != SQL_NEED_DATA) { - LOG("SQLExecuteMany: SQLPutData(bytes) failed - " - "chunk=%zu, rc=%d", - dae_chunk_count, rc); - } - } else { - LOG("SQLExecuteMany: Unsupported DAE data type - chunk=%zu", dae_chunk_count); - return SQL_ERROR; + if (rc != SQL_NEED_DATA) break; + + const ParamInfo* matchedInfo = reinterpret_cast(token); + const ParamInfo* first = rowParamInfos.data(); + const ParamInfo* last = first + rowParamInfos.size(); + if (matchedInfo < first || matchedInfo >= last) { + ThrowStdException( + "SQLExecuteMany: unrecognized paramToken from SQLParamData"); } + rc = StreamDAEParameter(hStmt, *matchedInfo, charEncoding); + if (!SQL_SUCCEEDED(rc)) return rc; dae_chunk_count++; } LOG("SQLExecuteMany: DAE completed for row %zu - total_chunks=%zu, " @@ -3082,6 +3405,16 @@ SQLRETURN SQLExecuteMany_wrap(const SqlHandlePtr statementHandle, const std::u16 LOG("SQLExecuteMany: DAE row %zu failed - rc=%d", rowIndex, rc); return rc; } + + const SQLRETURN rowRc = rc; + rc = SQLFreeStmt_ptr(hStmt, SQL_RESET_PARAMS); + if (!SQL_SUCCEEDED(rc)) { + LOG("SQLExecuteMany: SQL_RESET_PARAMS failed for row %zu - " + "rc=%d", + rowIndex, rc); + return rc; + } + rc = rowRc; } LOG("SQLExecuteMany: All DAE rows processed successfully - " "total_rows=%zu", @@ -3257,6 +3590,9 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT py::handle messages) { PERF_TIMER("FetchLobColumnData"); std::vector buffer; + size_t dataSize = 0; + size_t reservedBytes = 0; + ReserveNativeFetchBytes(reservedBytes, DAE_CHUNK_SIZE, sizeof(char)); SQLRETURN ret = SQL_SUCCESS_WITH_INFO; int loopCount = 0; @@ -3326,7 +3662,9 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT } } if (bytesRead > 0) { - buffer.insert(buffer.end(), chunk.begin(), chunk.begin() + bytesRead); + const size_t previousSize = dataSize; + dataSize = ExtendNativeFetchBuffer(buffer, dataSize, bytesRead, reservedBytes); + std::copy_n(chunk.data(), bytesRead, buffer.data() + previousSize); LOG("FetchLobColumnData: Appended %zu bytes at loop %d", bytesRead, loopCount); } if (ret == SQL_SUCCESS) { @@ -3334,35 +3672,42 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT break; } } - LOG("FetchLobColumnData: Total bytes collected=%zu for column %d", buffer.size(), colIndex); + LOG("FetchLobColumnData: Total bytes collected=%zu for column %d", dataSize, colIndex); - if (buffer.empty()) { + if (dataSize == 0) { if (isBinary) { return py::bytes(""); } return py::str(""); } if (isWideChar) { - size_t wcharCount = buffer.size() / sizeof(SQLWCHAR); - std::vector alignedBuf(wcharCount); - std::memcpy(alignedBuf.data(), buffer.data(), buffer.size()); - return py::cast(dupeSqlWCharAsUtf16Le(alignedBuf.data(), wcharCount)); + if (dataSize % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character LOB data has an invalid byte length"); + } + int byteOrder = -1; + PyObject* decoded = + PyUnicode_DecodeUTF16(buffer.data(), static_cast(dataSize), + "strict", &byteOrder); + if (decoded == nullptr) { + throw py::error_already_set(); + } + return py::reinterpret_steal(decoded); } if (isBinary) { LOG("FetchLobColumnData: Returning binary data - %zu bytes for column " "%d", - buffer.size(), colIndex); - return py::bytes(buffer.data(), buffer.size()); + dataSize, colIndex); + return py::bytes(buffer.data(), dataSize); } // For SQL_C_CHAR data, decode using the appropriate encoding. const std::string effectiveCharEncoding = GetEffectiveCharDecoding(charEncoding); - py::bytes raw_bytes(buffer.data(), buffer.size()); + py::bytes raw_bytes(buffer.data(), dataSize); try { py::object decoded = raw_bytes.attr("decode")(effectiveCharEncoding, "strict"); LOG("FetchLobColumnData: Decoded narrow string with '%s' - %zu bytes -> %zu chars for " "column %d", - effectiveCharEncoding.c_str(), buffer.size(), py::len(decoded), colIndex); + effectiveCharEncoding.c_str(), dataSize, py::len(decoded), colIndex); return decoded; } catch (const py::error_already_set& e) { LOG_ERROR("FetchLobColumnData: Failed to decode with '%s' for column %d: %s", @@ -4277,9 +4622,18 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri // TODO: Move to anonymous namespace, since it is not used outside this file template SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames, - SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR, - py::handle messages = {}) { + SQLUSMALLINT numCols, int fetchSize, size_t& reservedBytes, + int charCtype = SQL_C_WCHAR, py::handle messages = {}) { PERF_TIMER("SQLBindColums"); + struct UnbindOnFailure { + SQLHSTMT hStmt; + bool active = true; + ~UnbindOnFailure() { + if (active) { + SQLFreeStmt_ptr(hStmt, SQL_UNBIND); + } + } + } unbindOnFailure{hStmt}; SQLRETURN ret = SQL_SUCCESS; const bool useWideChar = (charCtype == SQL_C_WCHAR); // Bind columns based on their data types @@ -4293,25 +4647,39 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& 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 = CheckedAddSize( + baseColumnSize, 1, "Column fetch stride is too large"); + ResizeNativeFetchBuffer(buffers.wcharBuffers[col - 1], + CheckedMultiplySize(fetchSize, fetchBufferSize, + "Native fetch buffer is too large"), + 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 = CheckedAddSize( + CheckedMultiplySize(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 = CheckedAddSize( + baseColumnSize, 1, "Column fetch stride is too large"); #endif - buffers.charBuffers[col - 1].resize(fetchSize * fetchBufferSize); + ResizeNativeFetchBuffer(buffers.charBuffers[col - 1], + CheckedMultiplySize(fetchSize, fetchBufferSize, + "Native fetch buffer is too large"), + 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; } @@ -4321,49 +4689,57 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // 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 = CheckedAddSize( + CheckedFetchColumnSize(columnSize), 1, "Column fetch stride is too large"); + ResizeNativeFetchBuffer(buffers.wcharBuffers[col - 1], + CheckedMultiplySize(fetchSize, fetchBufferSize, + "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; } case SQL_INTEGER: - buffers.intBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.intBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_SLONG, buffers.intBuffers[col - 1].data(), sizeof(SQLINTEGER), buffers.indicators[col - 1].data()); break; case SQL_SMALLINT: - buffers.smallIntBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.smallIntBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_SSHORT, buffers.smallIntBuffers[col - 1].data(), sizeof(SQLSMALLINT), buffers.indicators[col - 1].data()); break; case SQL_TINYINT: - buffers.charBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.charBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_TINYINT, buffers.charBuffers[col - 1].data(), sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; case SQL_BIT: - buffers.charBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.charBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_BIT, buffers.charBuffers[col - 1].data(), sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; case SQL_REAL: - buffers.realBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.realBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_FLOAT, buffers.realBuffers[col - 1].data(), sizeof(SQLREAL), buffers.indicators[col - 1].data()); break; case SQL_DECIMAL: case SQL_NUMERIC: - buffers.charBuffers[col - 1].resize(fetchSize * MAX_DIGITS_IN_NUMERIC); + ResizeNativeFetchBuffer( + buffers.charBuffers[col - 1], + CheckedMultiplySize(fetchSize, MAX_DIGITS_IN_NUMERIC, + "Native fetch buffer is too large"), + reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_CHAR, buffers.charBuffers[col - 1].data(), MAX_DIGITS_IN_NUMERIC * sizeof(SQLCHAR), buffers.indicators[col - 1].data()); break; case SQL_DOUBLE: case SQL_FLOAT: - buffers.doubleBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.doubleBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_DOUBLE, buffers.doubleBuffers[col - 1].data(), sizeof(SQLDOUBLE), buffers.indicators[col - 1].data()); @@ -4371,50 +4747,58 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& case SQL_TIMESTAMP: case SQL_TYPE_TIMESTAMP: case SQL_DATETIME: - buffers.timestampBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.timestampBuffers[col - 1], fetchSize, + reservedBytes); ret = SQLBindCol_ptr( hStmt, col, SQL_C_TYPE_TIMESTAMP, buffers.timestampBuffers[col - 1].data(), sizeof(SQL_TIMESTAMP_STRUCT), buffers.indicators[col - 1].data()); break; case SQL_BIGINT: - buffers.bigIntBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.bigIntBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_SBIGINT, buffers.bigIntBuffers[col - 1].data(), sizeof(SQLBIGINT), buffers.indicators[col - 1].data()); break; case SQL_TYPE_DATE: - buffers.dateBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.dateBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_TYPE_DATE, buffers.dateBuffers[col - 1].data(), sizeof(SQL_DATE_STRUCT), buffers.indicators[col - 1].data()); break; case SQL_SS_TIME2: - buffers.timeBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.timeBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_SS_TIME2, buffers.timeBuffers[col - 1].data(), sizeof(SQL_SS_TIME2_STRUCT), buffers.indicators[col - 1].data()); break; case SQL_GUID: - buffers.guidBuffers[col - 1].resize(fetchSize); + ResizeNativeFetchBuffer(buffers.guidBuffers[col - 1], fetchSize, reservedBytes); ret = SQLBindCol_ptr(hStmt, col, SQL_C_GUID, buffers.guidBuffers[col - 1].data(), sizeof(SQLGUID), buffers.indicators[col - 1].data()); break; case SQL_SS_UDT: case SQL_BINARY: case SQL_VARBINARY: - case SQL_LONGVARBINARY: + case SQL_LONGVARBINARY: { // TODO: handle variable length data correctly. This logic wont // suffice HandleZeroColumnSizeAtFetch(columnSize); - buffers.charBuffers[col - 1].resize(fetchSize * columnSize); + const size_t fetchBufferSize = CheckedFetchColumnSize(columnSize); + ResizeNativeFetchBuffer(buffers.charBuffers[col - 1], + CheckedMultiplySize(fetchSize, fetchBufferSize, + "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(fetchBufferSize, 1), + buffers.indicators[col - 1].data()); break; + } case SQL_SS_TIMESTAMPOFFSET: - buffers.datetimeoffsetBuffers[col - 1].resize(fetchSize); + 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: @@ -4437,6 +4821,7 @@ SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& } CaptureFetchDiagnostics(hStmt, ret, messages); } + unbindOnFailure.active = false; return ret; } @@ -4452,6 +4837,7 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& LOG("FetchBatchData: Fetching data in batches"); SQLRETURN ret; { + numRowsFetched = 0; // Release the GIL during the blocking ODBC fetch py::gil_scoped_release release; PERF_TIMER("FetchBatchData::SQLFetchScroll_call"); @@ -4807,77 +5193,96 @@ SQLRETURN FetchBatchData(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& // Given a list of columns that are a part of single row in the result set, // calculates the max size of the row // TODO: Move to anonymous namespace, since it is not used outside this file -size_t calculateRowSize(py::list& columnNames, SQLUSMALLINT numCols) { - size_t rowSize = 0; +size_t calculateRowSize(py::list& columnNames, SQLUSMALLINT numCols, + int charCtype = SQL_C_WCHAR) { + size_t rowSize = CheckedMultiplySize(numCols, sizeof(SQLLEN), + "Native fetch indicator row size is too large"); for (SQLUSMALLINT col = 1; col <= numCols; col++) { auto columnMeta = columnNames[col - 1].cast(); SQLSMALLINT dataType = columnMeta["DataType"].cast(); SQLULEN columnSize = columnMeta["ColumnSize"].cast(); + size_t columnBytes = 0; switch (dataType) { case SQL_CHAR: case SQL_VARCHAR: - case SQL_LONGVARCHAR: - rowSize += columnSize; + case SQL_LONGVARCHAR: { + HandleZeroColumnSizeAtFetch(columnSize); + if (charCtype == SQL_C_WCHAR) { + columnBytes = CheckedMultiplySize( + CheckedAddSize(columnSize, 1, "Native fetch row size is too large"), + sizeof(SQLWCHAR), "Native fetch row size is too large"); + } else { +#if defined(__APPLE__) || defined(__linux__) + columnBytes = CheckedAddSize( + CheckedMultiplySize(columnSize, 4, + "Native fetch row size is too large"), + 1, "Native fetch row size is too large"); +#else + columnBytes = + CheckedAddSize(columnSize, 1, "Native fetch row size is too large"); +#endif + } break; + } case SQL_SS_XML: case SQL_WCHAR: case SQL_WVARCHAR: case SQL_WLONGVARCHAR: - rowSize += columnSize * sizeof(SQLWCHAR); + HandleZeroColumnSizeAtFetch(columnSize); + columnBytes = CheckedMultiplySize( + CheckedAddSize(columnSize, 1, "Native fetch row size is too large"), + sizeof(SQLWCHAR), "Native fetch row size is too large"); break; case SQL_INTEGER: - rowSize += sizeof(SQLINTEGER); + columnBytes = sizeof(SQLINTEGER); break; case SQL_SMALLINT: - rowSize += sizeof(SQLSMALLINT); + columnBytes = sizeof(SQLSMALLINT); break; case SQL_REAL: - rowSize += sizeof(SQLREAL); + columnBytes = sizeof(SQLREAL); break; case SQL_FLOAT: - rowSize += sizeof(SQLFLOAT); + columnBytes = sizeof(SQLFLOAT); break; case SQL_DOUBLE: - rowSize += sizeof(SQLDOUBLE); + columnBytes = sizeof(SQLDOUBLE); break; case SQL_DECIMAL: case SQL_NUMERIC: - rowSize += MAX_DIGITS_IN_NUMERIC; + columnBytes = MAX_DIGITS_IN_NUMERIC; break; case SQL_TIMESTAMP: case SQL_TYPE_TIMESTAMP: case SQL_DATETIME: - rowSize += sizeof(SQL_TIMESTAMP_STRUCT); + columnBytes = sizeof(SQL_TIMESTAMP_STRUCT); break; case SQL_BIGINT: - rowSize += sizeof(SQLBIGINT); + columnBytes = sizeof(SQLBIGINT); break; case SQL_TYPE_DATE: - rowSize += sizeof(SQL_DATE_STRUCT); + columnBytes = sizeof(SQL_DATE_STRUCT); break; case SQL_SS_TIME2: - rowSize += sizeof(SQL_SS_TIME2_STRUCT); + columnBytes = sizeof(SQL_SS_TIME2_STRUCT); break; case SQL_GUID: - rowSize += sizeof(SQLGUID); + columnBytes = sizeof(SQLGUID); break; case SQL_TINYINT: case SQL_BIT: - rowSize += sizeof(SQLCHAR); + columnBytes = sizeof(SQLCHAR); break; case SQL_SS_UDT: - rowSize += (static_cast(columnSize) == SQL_NO_TOTAL || columnSize == 0) - ? SQL_MAX_LOB_SIZE - : columnSize; - break; case SQL_BINARY: case SQL_VARBINARY: case SQL_LONGVARBINARY: - rowSize += columnSize; + HandleZeroColumnSizeAtFetch(columnSize); + columnBytes = columnSize; break; case SQL_SS_TIMESTAMPOFFSET: - rowSize += sizeof(DateTimeOffset); + columnBytes = sizeof(DateTimeOffset); break; default: std::string columnName = columnMeta["ColumnName"].cast(); @@ -4888,6 +5293,7 @@ size_t calculateRowSize(py::list& columnNames, SQLUSMALLINT numCols) { ThrowStdException(errorString.str()); break; } + rowSize = CheckedAddSize(rowSize, columnBytes, "Native fetch row size is too large"); } return rowSize; } @@ -4965,11 +5371,13 @@ struct FetchStateGuard { // the result set and populates the provided Python list with the row data. If // there are no more rows to fetch, it returns SQL_NO_DATA. If an error occurs // during fetching, it throws a runtime error. -SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetchSize, +SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, py::handle fetchSizeArg, const std::string& charEncoding = "utf-16le", const std::string& wcharEncoding = "utf-16le", int charCtype = SQL_C_WCHAR, py::handle messages = {}) { PERF_TIMER("FetchMany_wrap"); + const int fetchSize = + ValidateNativeRowCountArgument(fetchSizeArg, "Fetch size", false); // Issue #531: upgrade SQL_C_CHAR + utf-8 to SQL_C_WCHAR on Windows so the // driver does lossless UTF-16 conversion instead of returning ACP bytes. charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding); @@ -5055,11 +5463,18 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch } // Initialize column buffers + size_t reservedBytes = 0; + ReserveNativeFetchBytes( + reservedBytes, + CheckedMultiplySize(static_cast(numCols), static_cast(fetchSize), + "Native fetch indicator buffer is too large"), + sizeof(SQLLEN)); ColumnBuffers buffers(numCols, fetchSize); FetchStateGuard fetchStateGuard(StatementHandle, messages); // Bind columns - ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, charCtype, messages); + ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, reservedBytes, charCtype, + messages); if (!SQL_SUCCEEDED(ret)) { LOG("FetchMany_wrap: Error when binding columns - SQLRETURN=%d", ret); return ret; @@ -5094,7 +5509,8 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch // @return SQLRETURN: SQL_SUCCESS on success, or error code on failure template SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, - std::vector& dataVec, SQLLEN* indicator, py::handle messages) { + std::vector& dataVec, SQLLEN* indicator, size_t& reservedBytes, + py::handle messages) { size_t start = 0; size_t end = 0; @@ -5114,7 +5530,7 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // Ensure initial buffer has space for at least the null terminator if (dataVec.size() < sizeNullTerminator) { - dataVec.resize(sizeNullTerminator); + ResizeNativeFetchBuffer(dataVec, sizeNullTerminator, reservedBytes); } while (true) { @@ -5151,18 +5567,22 @@ SQLRETURN GetDataVar(SQLHSTMT hStmt, SQLUSMALLINT colNumber, SQLSMALLINT cType, // Determine how much more space we need if (localInd < 0) { // SQL_NO_TOTAL: driver doesn't know total size, double the buffer - end = dataVec.size() * 2; + end = CheckedMultiplySize(dataVec.size(), 2, + "Native fetch buffer size is too large"); } else { // Driver returned total size: allocate exactly what we need assert(localInd % sizeof(T) == 0); - end = start + static_cast(localInd) / sizeof(T) + sizeNullTerminator; + end = CheckedAddSize( + CheckedAddSize(start, static_cast(localInd) / sizeof(T), + "Native fetch buffer size is too large"), + sizeNullTerminator, "Native fetch buffer size is too large"); } // The next read starts where the null terminator would have been placed start = dataVec.size() - sizeNullTerminator; // Resize buffer for next iteration - dataVec.resize(end); + ResizeNativeFetchBuffer(dataVec, end, reservedBytes); } else { // Unexpected return code return ret; @@ -5183,14 +5603,23 @@ int32_t days_from_civil(int y, int m, int d) { return era * 146097 + static_cast(doe) - 719468; } -SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, int arrowBatchSize, - int charCtype, py::handle messages = {}) { +SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, + py::handle arrowBatchSizeArg, int charCtype, + py::handle messages = {}) { PERF_TIMER("FetchArrowBatch_wrap"); + const int arrowBatchSize = + ValidateNativeRowCountArgument(arrowBatchSizeArg, "Arrow batch size", true); + const size_t batchSize = static_cast(arrowBatchSize); + const size_t offsetCount = CheckedAddSize(batchSize, 1, "Arrow batch size is too large"); + const size_t initialVarDataSize = + CheckedMultiplySize(batchSize, 42, "Arrow batch size is too large"); + const size_t bitmapSize = + CheckedAddSize(batchSize, 7, "Arrow batch size 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"); - // An overly large fetch size doesn't seem to help performance - int fetchSize = 64; + // An overly large fetch size doesn't seem to help performance. + int fetchSize = arrowBatchSize > 0 ? std::min(arrowBatchSize, 64) : 1; SQLRETURN ret = SQL_ERROR; ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret); @@ -5212,10 +5641,10 @@ 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); + size_t reservedBytes = 0; std::vector> arrowArrayPrivateData(numCols); std::vector> arrowSchemaPrivateData(numCols); @@ -5230,7 +5659,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 || @@ -5259,8 +5687,11 @@ 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 = + AllocateArrowArray(offsetCount, reservedBytes, + "Arrow offset buffer is too large"); + ResizeNativeFetchBuffer(arrowColumnProducer->varData, initialVarDataSize, + reservedBytes); columnVarLen[i] = true; // start at offset 0 arrowColumnProducer->varVal[0] = 0; @@ -5271,8 +5702,11 @@ 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 = + AllocateArrowArray(offsetCount, reservedBytes, + "Arrow offset buffer is too large"); + ResizeNativeFetchBuffer(arrowColumnProducer->varData, initialVarDataSize, + reservedBytes); columnVarLen[i] = true; // start at offset 0 arrowColumnProducer->varVal[0] = 0; @@ -5280,33 +5714,45 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, break; case SQL_TINYINT: format = "C"; - arrowColumnProducer->uint8Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->uint8Val = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->uint8Val.get(); break; case SQL_SMALLINT: format = "s"; - arrowColumnProducer->int16Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->int16Val = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->int16Val.get(); break; case SQL_INTEGER: format = "i"; - arrowColumnProducer->int32Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->int32Val = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->int32Val.get(); break; case SQL_BIGINT: format = "l"; - arrowColumnProducer->int64Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->int64Val = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->int64Val.get(); break; case SQL_REAL: format = "f"; - arrowColumnProducer->float32Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->float32Val = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->float32Val.get(); break; case SQL_FLOAT: case SQL_DOUBLE: format = "g"; - arrowColumnProducer->float64Val = std::make_unique(arrowBatchSize); + arrowColumnProducer->float64Val = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->float64Val.get(); break; case SQL_DECIMAL: @@ -5319,7 +5765,9 @@ 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 = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->decimalVal.get(); break; } @@ -5327,28 +5775,38 @@ 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 = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->tsMicroVal.get(); break; case SQL_SS_TIMESTAMPOFFSET: format = "tsu:+00:00"; - arrowColumnProducer->tsMicroVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->tsMicroVal = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->tsMicroVal.get(); break; case SQL_TYPE_DATE: format = "tdD"; - arrowColumnProducer->dateVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->dateVal = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->dateVal.get(); break; case SQL_SS_TIME2: format = "ttn"; - arrowColumnProducer->timeNanoVal = std::make_unique(arrowBatchSize); + arrowColumnProducer->timeNanoVal = + AllocateArrowArray(batchSize, reservedBytes, + "Arrow value buffer is too large"); 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 = + AllocateArrowArray(bitmapSize, reservedBytes, + "Arrow bitmap is too large"); + std::memset(arrowColumnProducer->bitVal.get(), 0, bitmapSize); arrowColumnProducer->ptrValueBuffer = arrowColumnProducer->bitVal.get(); break; default: @@ -5369,18 +5827,39 @@ 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 = + AllocateArrowArray(bitmapSize, reservedBytes, "Arrow bitmap is too large"); // Initialize validity bitmap to all valid - std::memset(arrowColumnProducer->valid.get(), 0xFF, (arrowBatchSize + 7) / 8); + std::memset(arrowColumnProducer->valid.get(), 0xFF, bitmapSize); } + if (arrowBatchSize > 0 && !hasLobColumns) { + const size_t rowSize = calculateRowSize(columnNames, numCols, charCtype); + const size_t remainingBytes = MAX_NATIVE_FETCH_BYTES - reservedBytes; + const size_t rowsWithinBudget = remainingBytes / rowSize; + if (rowsWithinBudget == 0) { + ThrowStdException("Native fetch buffers exceed the 256 MiB allocation limit"); + } + fetchSize = static_cast(std::min(fetchSize, rowsWithinBudget)); + } + + const int bufferFetchSize = arrowBatchSize > 0 ? fetchSize : 0; + if (bufferFetchSize > 0) { + ReserveNativeFetchBytes( + reservedBytes, + CheckedMultiplySize(static_cast(numCols), + static_cast(bufferFetchSize), + "Native fetch indicator buffer is too large"), + sizeof(SQLLEN)); + } // Initialize column buffers - ColumnBuffers buffers(numCols, fetchSize); + ColumnBuffers buffers(numCols, bufferFetchSize); SQLULEN numRowsFetched = 0; FetchStateGuard fetchStateGuard(StatementHandle, messages); - if (!hasLobColumns && fetchSize > 0) { - ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, charCtype, messages); + if (!hasLobColumns && bufferFetchSize > 0) { + ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, reservedBytes, + charCtype, messages); if (!SQL_SUCCEEDED(ret)) { LOG("Error when binding columns"); return ret; @@ -5418,7 +5897,6 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, 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"); @@ -5430,7 +5908,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_LONGVARBINARY: { ret = GetDataVar(hStmt, idxCol + 1, SQL_C_BINARY, buffers.charBuffers[idxCol], - buffers.indicators[idxCol].data(), messages); + buffers.indicators[idxCol].data(), reservedBytes, messages); CheckFetchError(StatementHandle, ret); if (!SQL_SUCCEEDED(ret)) { LOG("Error fetching BINARY LOB for column %d", idxCol + 1); @@ -5444,7 +5922,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, if (charCtype == SQL_C_CHAR) { ret = GetDataVar(hStmt, idxCol + 1, SQL_C_CHAR, buffers.charBuffers[idxCol], - buffers.indicators[idxCol].data(), messages); + buffers.indicators[idxCol].data(), reservedBytes, messages); CheckFetchError(StatementHandle, ret); if (!SQL_SUCCEEDED(ret)) { LOG("Error fetching CHAR LOB data for column %d", idxCol + 1); @@ -5460,7 +5938,7 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_WLONGVARCHAR: { ret = GetDataVar(hStmt, idxCol + 1, SQL_C_WCHAR, buffers.wcharBuffers[idxCol], - buffers.indicators[idxCol].data(), messages); + buffers.indicators[idxCol].data(), reservedBytes, messages); CheckFetchError(StatementHandle, ret); if (!SQL_SUCCEEDED(ret)) { LOG("Error fetching WCHAR LOB data for column %d", idxCol + 1); @@ -5695,16 +6173,23 @@ 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 */; auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; - while (target_vec->size() < start + dataLen) { - target_vec->resize(target_vec->size() * 2); - } + if (dataLen > 0) { + 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][idxRowSql * fetchBufferSize], - dataLen); + std::memcpy(&(*target_vec)[start], + &buffers.charBuffers[idxCol][sourceOffset], dataLen); + } arrowColumnProducer->varVal[idxRowArrow + 1] = start + dataLen; break; } @@ -5712,20 +6197,25 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, case SQL_VARCHAR: case SQL_LONGVARCHAR: { if (charCtype == SQL_C_CHAR) { -#if defined(__APPLE__) || defined(__linux__) - uint64_t fetchBufferSize = columnSize * 4 + 1 /*null-terminator*/; -#else - uint64_t fetchBufferSize = columnSize + 1 /*null-terminator*/; -#endif auto target_vec = &arrowColumnProducer->varData; auto start = arrowColumnProducer->varVal[idxRowArrow]; - while (target_vec->size() < start + dataLen) { - target_vec->resize(target_vec->size() * 2); + if (dataLen > 0) { + 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); } - - std::memcpy(&(*target_vec)[start], - &buffers.charBuffers[idxCol][idxRowSql * fetchBufferSize], - dataLen); arrowColumnProducer->varVal[idxRowArrow + 1] = start + dataLen; break; } @@ -5736,20 +6226,29 @@ 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); + if (dataLen % sizeof(SQLWCHAR) != 0) { + ThrowStdException("Wide-character data has an invalid byte length"); + } auto dataLenW = dataLen / sizeof(SQLWCHAR); - auto wcharSource = - &buffers.wcharBuffers[idxCol][idxRowSql * (columnSize + 1)]; + const size_t sourceStride = hasLobColumns + ? buffers.wcharBuffers[idxCol].size() + : buffers.wcharBuffers[idxCol].size() / + static_cast(fetchSize); + const size_t sourceOffset = CheckedArrowSourceOffset( + buffers.wcharBuffers[idxCol], idxRowSql, sourceStride, dataLen); + auto wcharSource = &buffers.wcharBuffers[idxCol][sourceOffset]; auto 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); - } + size_t maxUtf8Size = + CheckedMultiplySize(dataLenW, 3, "Arrow value buffer is too large"); + EnsureNativeFetchBufferSize( + *target_vec, + CheckedAddSize(start, maxUtf8Size, + "Arrow value buffer is too large"), + reservedBytes); size_t bytesWritten = simdutf::convert_utf16le_to_utf8_with_replacement( utf16Source, dataLenW, @@ -5766,9 +6265,10 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules, auto 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); - } + EnsureNativeFetchBufferSize( + *target_vec, + CheckedAddSize(start, 37, "Arrow value buffer is too large"), + reservedBytes); // Get the GUID from the buffer const SQLGUID& guidValue = buffers.guidBuffers[idxCol][idxRowSql]; @@ -6193,14 +6693,12 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, } // No LOBs detected - use binding path with batch fetching - // Define a memory limit (1 GB) - const size_t memoryLimit = 1ULL * 1024 * 1024 * 1024; - size_t totalRowSize = calculateRowSize(columnNames, numCols); + size_t totalRowSize = calculateRowSize(columnNames, numCols, charCtype); - // Calculate fetch size based on the total row size and memory limit + // Calculate fetch size from the same storage budget enforced during binding. size_t numRowsInMemLimit; if (totalRowSize > 0) { - numRowsInMemLimit = static_cast(memoryLimit / totalRowSize); + numRowsInMemLimit = MAX_NATIVE_FETCH_BYTES / totalRowSize; } else { // Handle case where totalRowSize is 0 to avoid division by zero. // This can happen for NVARCHAR(MAX) cols. SQLDescribeCol returns 0 @@ -6231,14 +6729,23 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows, } else { fetchSize = 1000; } + fetchSize = static_cast( + std::min(fetchSize, std::max(numRowsInMemLimit, 1))); LOG("FetchAll_wrap: Fetching data in batch sizes of %d", fetchSize); + size_t reservedBytes = 0; + ReserveNativeFetchBytes( + reservedBytes, + CheckedMultiplySize(static_cast(numCols), static_cast(fetchSize), + "Native fetch indicator buffer is too large"), + sizeof(SQLLEN)); ColumnBuffers buffers(numCols, fetchSize); SQLULEN numRowsFetched = 0; FetchStateGuard fetchStateGuard(StatementHandle, messages); // Bind columns - ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, charCtype, messages); + ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, reservedBytes, charCtype, + messages); if (!SQL_SUCCEEDED(ret)) { LOG("FetchAll_wrap: Error when binding columns - SQLRETURN=%d", ret); return ret; @@ -6530,6 +7037,8 @@ PYBIND11_MODULE(ddbc_bindings, m) { "Fetch an arrow batch of given length from the result set", py::arg("StatementHandle"), py::arg("capsules"), py::arg("arrowBatchSize"), py::arg("charCtype"), py::arg("messages") = py::none()); + m.def("_test_lob_buffer_growth", &TestLobBufferGrowth, py::arg("chunk_count"), + py::arg("chunk_size")); m.def("DDBCSQLFreeHandle", &SQLFreeHandle_wrap, "Free a handle"); m.def("DDBCSQLResetStmt", &SQLResetStmt_wrap, "Close cursor and unbind params without freeing HSTMT"); diff --git a/mssql_python/pybind/param_detect.hpp b/mssql_python/pybind/param_detect.hpp index 172a98b51..6443d3d23 100644 --- a/mssql_python/pybind/param_detect.hpp +++ b/mssql_python/pybind/param_detect.hpp @@ -140,6 +140,26 @@ inline constexpr int MAX_INLINE_CHAR = 4000; // Binary data longer than this uses DAE streaming (SQL Server max for non-MAX types) inline constexpr int MAX_INLINE_BINARY = 8000; +inline SQLULEN DAEColumnSize(SQLSMALLINT sqlType, SQLULEN actualSize) { + switch (sqlType) { + case SQL_CHAR: + case SQL_VARCHAR: + return actualSize > MAX_INLINE_CHAR ? 0 : actualSize; + case SQL_WCHAR: + case SQL_WVARCHAR: + return actualSize > MAX_INLINE_CHAR ? 0 : actualSize; + case SQL_BINARY: + case SQL_VARBINARY: + return actualSize > MAX_INLINE_BINARY ? 0 : actualSize; + case SQL_LONGVARCHAR: + case SQL_WLONGVARCHAR: + case SQL_LONGVARBINARY: + return actualSize; + default: + return actualSize; + } +} + // SQL Server maximum numeric precision inline constexpr int MAX_NUMERIC_PRECISION = 38; @@ -190,6 +210,45 @@ inline bool PyLongGreaterThan(PyObject* value, long long threshold) { return overflow > 0 || (overflow == 0 && result > threshold); } +inline Py_ssize_t UnicodeUtf16Length(PyObject* value) { + const Py_ssize_t length = PyUnicode_GET_LENGTH(value); + if (PyUnicode_KIND(value) <= PyUnicode_2BYTE_KIND) { + return length; + } + + Py_ssize_t utf16Length = 0; + const Py_UCS4* data = PyUnicode_4BYTE_DATA(value); + for (Py_ssize_t index = 0; index < length; ++index) { + utf16Length += data[index] > 0xFFFF ? 2 : 1; + } + return utf16Length; +} + +inline Py_ssize_t EncodedUnicodeLength(PyObject* value, const std::string& encoding) { + py::object encoderFactory = + py::module_::import("codecs").attr("getincrementalencoder")(encoding); + py::object encoder = encoderFactory("strict"); + const Py_ssize_t length = PyUnicode_GET_LENGTH(value); + constexpr Py_ssize_t chunkSize = 4096; + Py_ssize_t total = 0; + for (Py_ssize_t offset = 0; offset < length; offset += chunkSize) { + const Py_ssize_t end = std::min(offset + chunkSize, length); + py::object chunk = steal(PyUnicode_Substring(value, offset, end)); + if (!chunk) throw py::error_already_set(); + py::object encoded = encoder.attr("encode")(chunk, end == length); + const Py_ssize_t encodedSize = PyBytes_GET_SIZE(encoded.ptr()); + if (encodedSize > MAX_INLINE_BINARY - total) { + return MAX_INLINE_BINARY + 1; + } + total += encodedSize; + } + if (length == 0) { + py::object encoded = encoder.attr("encode")(py::str(), true); + total = PyBytes_GET_SIZE(encoded.ptr()); + } + return total; +} + inline PyObject* FormatDecimalParam(PyObject* params, Py_ssize_t index, PyObject* value) { py::object formatted = steal(PyObject_CallMethod(value, "__format__", "s", "f")); if (!formatted) throw py::error_already_set(); @@ -215,8 +274,91 @@ inline void NormalizeTimeParam(PyObject* params, Py_ssize_t index, SQLULEN& colu } } +inline long long ValidatedInputSizeInteger(PyObject* value, const char* fieldName) { + if (!PyLong_Check(value) || PyBool_Check(value)) { + throw py::type_error(std::string(fieldName) + " must be an integer"); + } + int overflow = 0; + const long long result = PyLong_AsLongLongAndOverflow(value, &overflow); + if ((result == -1 && PyErr_Occurred()) || overflow != 0) { + PyErr_Clear(); + throw py::value_error(std::string(fieldName) + " is out of range"); + } + return result; +} + +inline void ValidateInputSizes(PyObject* inputSizes) { + if (inputSizes == Py_None) { + return; + } + if (!PyList_Check(inputSizes)) { + throw py::type_error("inputSizes must be None or a list"); + } + + const Py_ssize_t count = PyList_GET_SIZE(inputSizes); + for (Py_ssize_t index = 0; index < count; ++index) { + PyObject* entry = PyList_GET_ITEM(inputSizes, index); + if (!PyTuple_Check(entry) || PyTuple_GET_SIZE(entry) != 4) { + throw py::type_error("each inputSizes entry must be a four-item tuple"); + } + + const long long sqlType = + ValidatedInputSizeInteger(PyTuple_GET_ITEM(entry, 0), "SQL type"); + const long long cType = + ValidatedInputSizeInteger(PyTuple_GET_ITEM(entry, 1), "C type"); + PyObject* columnSize = PyTuple_GET_ITEM(entry, 2); + PyObject* decimalDigits = PyTuple_GET_ITEM(entry, 3); + if (!PyLong_Check(columnSize) || PyBool_Check(columnSize) || + !PyLong_Check(decimalDigits) || PyBool_Check(decimalDigits)) { + throw py::type_error("column size and decimal digits must be integers"); + } + + if (sqlType < std::numeric_limits::min() || + sqlType > std::numeric_limits::max() || + cType < std::numeric_limits::min() || + cType > std::numeric_limits::max()) { + throw py::value_error("SQL and C types must fit in SQLSMALLINT"); + } + const bool isNumeric = sqlType == SQL_DECIMAL || sqlType == SQL_NUMERIC; + py::int_ zero(0); + const int negativeColumnSize = + PyObject_RichCompareBool(columnSize, zero.ptr(), Py_LT); + const int negativeDecimalDigits = + PyObject_RichCompareBool(decimalDigits, zero.ptr(), Py_LT); + if (negativeColumnSize == -1 || negativeDecimalDigits == -1) { + throw py::error_already_set(); + } + if (negativeColumnSize == 1) { + throw py::value_error("column size must be non-negative"); + } + if (negativeDecimalDigits == 1) { + throw py::value_error("decimal digits must be non-negative"); + } + if (!isNumeric) { + const unsigned long long requestedSize = PyLong_AsUnsignedLongLong(columnSize); + if (requestedSize == static_cast(-1) && PyErr_Occurred()) { + PyErr_Clear(); + throw py::value_error("column size is out of range"); + } + if (requestedSize > std::numeric_limits::max()) { + throw py::value_error("column size is out of range"); + } + const unsigned long long requestedDigits = + PyLong_AsUnsignedLongLong(decimalDigits); + if (requestedDigits == static_cast(-1) && PyErr_Occurred()) { + PyErr_Clear(); + throw py::value_error("decimal digits are out of range"); + } + if (requestedDigits > + static_cast(std::numeric_limits::max())) { + throw py::value_error("decimal digits are out of range"); + } + } + } +} + inline void ApplyInputSizeOverride(PyObject* params, PyObject* inputSize, Py_ssize_t index, - ParamInfo& info) { + ParamInfo& info, const std::string& charEncoding) { py::tuple values = borrow(inputSize); info.paramSQLType = values[0].cast(); info.paramCType = values[1].cast(); @@ -257,10 +399,42 @@ inline void ApplyInputSizeOverride(PyObject* params, PyObject* inputSize, Py_ssi } } - info.isDAE = - (PyUnicode_Check(obj) && PyLongGreaterThan(columnSize, MAX_INLINE_CHAR)) || - ((PyBytes_Check(obj) || PyByteArray_Check(obj)) && - PyLongGreaterThan(columnSize, MAX_INLINE_BINARY)); + if ((PyBytes_Check(obj) || PyByteArray_Check(obj)) && + (info.paramSQLType == SQL_CHAR || info.paramSQLType == SQL_VARCHAR || + info.paramSQLType == SQL_LONGVARCHAR)) { + info.paramCType = SQL_C_CHAR; + } + + if (info.paramCType == SQL_C_WCHAR && + (PyBytes_Check(obj) || PyByteArray_Check(obj))) { + throw py::type_error("bytes values cannot be bound as SQL_C_WCHAR"); + } + + Py_ssize_t actualTextLength = 0; + if (PyUnicode_Check(obj)) { + actualTextLength = info.paramCType == SQL_C_CHAR + ? EncodedUnicodeLength(obj, charEncoding) + : UnicodeUtf16Length(obj); + } + const bool textNeedsDAE = + !isNumeric && PyUnicode_Check(obj) && actualTextLength > MAX_INLINE_CHAR; + const bool binaryNeedsDAE = + (PyBytes_Check(obj) || PyByteArray_Check(obj)) && + (PyBytes_Check(obj) ? PyBytes_GET_SIZE(obj) : PyByteArray_GET_SIZE(obj)) > + MAX_INLINE_BINARY; + info.isDAE = textNeedsDAE || binaryNeedsDAE; + if (!isNumeric && PyUnicode_Check(obj)) { + const SQLULEN actualSize = static_cast(actualTextLength); + info.columnSize = + info.isDAE ? DAEColumnSize(info.paramSQLType, actualSize) + : std::max(info.columnSize, actualSize); + } else if (PyBytes_Check(obj) || PyByteArray_Check(obj)) { + const SQLULEN actualSize = static_cast( + PyBytes_Check(obj) ? PyBytes_GET_SIZE(obj) : PyByteArray_GET_SIZE(obj)); + info.columnSize = + info.isDAE ? DAEColumnSize(info.paramSQLType, actualSize) + : std::max(info.columnSize, actualSize); + } if (PyTime_Check(obj) && info.paramCType == PARAM_C_TYPE_TEXT) { NormalizeTimeParam(params, index, info.columnSize); @@ -296,9 +470,15 @@ inline void ApplyInputSizeOverride(PyObject* params, PyObject* inputSize, Py_ssi // // Takes raw PyObject* lists. Caller guarantees params is a fresh copy (cursor.py // does list(actual_params)), so in-place mutation via PyList_SetItem is safe. -inline std::vector DetectParamTypes(PyObject* params, PyObject* inputSizes) { +inline std::vector DetectParamTypes(PyObject* params, PyObject* inputSizes, + const std::string& charEncoding = "utf-8") { PyTypeCache::initialize(); + if (!PyList_Check(params)) { + throw py::type_error("params must be a list"); + } + ValidateInputSizes(inputSizes); + const Py_ssize_t n = PyList_GET_SIZE(params); const Py_ssize_t inputSizeCount = inputSizes == Py_None ? 0 : PyList_GET_SIZE(inputSizes); std::vector infos(n); @@ -312,7 +492,7 @@ inline std::vector DetectParamTypes(PyObject* params, PyObject* input info.isDAE = false; if (i < inputSizeCount) { - ApplyInputSizeOverride(params, PyList_GET_ITEM(inputSizes, i), i, info); + ApplyInputSizeOverride(params, PyList_GET_ITEM(inputSizes, i), i, info, charEncoding); continue; } @@ -398,16 +578,7 @@ inline std::vector DetectParamTypes(PyObject* params, PyObject* input unsigned int kind = PyUnicode_KIND(obj); const void* udata = PyUnicode_DATA(obj); - Py_ssize_t utf16_len; - if (kind <= PyUnicode_2BYTE_KIND) { - utf16_len = length; - } else { - utf16_len = 0; - const Py_UCS4* data = PyUnicode_4BYTE_DATA(obj); - for (Py_ssize_t j = 0; j < length; ++j) { - utf16_len += (data[j] > 0xFFFF) ? 2 : 1; - } - } + const Py_ssize_t utf16_len = UnicodeUtf16Length(obj); // Detect whether the string needs wide-char (NVARCHAR) or narrow (VARCHAR) binding. // PyUnicode_IS_COMPACT_ASCII is a struct field check (O(1)), not a content scan. @@ -439,16 +610,14 @@ inline std::vector DetectParamTypes(PyObject* params, PyObject* input // Strings > 4000 UTF-16 code units exceed SQL Server's inline NVARCHAR(MAX) // threshold. Switch to data-at-execution (DAE) streaming: ODBC driver pulls // data in chunks via SQLPutData, avoiding a single massive buffer allocation. - // DAE path: match slow-path types exactly. - // Non-unicode (ASCII) → SQL_VARCHAR + PARAM_C_TYPE_TEXT, which is - // SQL_C_WCHAR and matches the slow path's SQL_C_CHAR (numerically - // -8 == SQL_C_WCHAR — a long-standing alias in the Python layer). - // Unicode → SQL_WVARCHAR + SQL_C_WCHAR (wide-char streaming) + // Use the validated payload size when it is legal fixed-width metadata; + // larger values use the driver's MAX-length sentinel. info.isDAE = true; - info.columnSize = 0; + const SQLSMALLINT sqlType = is_unicode ? SQL_WVARCHAR : SQL_VARCHAR; + info.columnSize = DAEColumnSize(sqlType, utf16_len); info.utf16Len = utf16_len; info.dataPtr = borrow(obj); - info.paramSQLType = is_unicode ? SQL_WVARCHAR : SQL_VARCHAR; + info.paramSQLType = sqlType; info.paramCType = is_unicode ? SQL_C_WCHAR : PARAM_C_TYPE_TEXT; } else { info.columnSize = is_unicode ? utf16_len : length; @@ -467,7 +636,8 @@ inline std::vector DetectParamTypes(PyObject* params, PyObject* input info.decimalDigits = 0; if (length > MAX_INLINE_BINARY) { info.isDAE = true; - info.columnSize = 0; + info.columnSize = + DAEColumnSize(SQL_VARBINARY, static_cast(length)); info.dataPtr = borrow(obj); } else { info.columnSize = std::max(length, 1); diff --git a/tests/test_004_cursor.py b/tests/test_004_cursor.py index dbd917879..02dec2e9f 100644 --- a/tests/test_004_cursor.py +++ b/tests/test_004_cursor.py @@ -1693,6 +1693,64 @@ def test_arraysize(cursor): assert cursor.arraysize == 5, "Arraysize mismatch after change" +@pytest.mark.parametrize("value", [0, -1, 1_000_001]) +def test_arraysize_rejects_out_of_range_values(value): + cursor = mssql_python.Cursor.__new__(mssql_python.Cursor) + with pytest.raises(ValueError, match="arraysize"): + cursor.arraysize = value + + +@pytest.mark.parametrize("value", [True, False, 1.5, "10"]) +def test_arraysize_rejects_non_integer_values(value): + cursor = mssql_python.Cursor.__new__(mssql_python.Cursor) + with pytest.raises(TypeError, match="arraysize"): + cursor.arraysize = value + + +@pytest.mark.parametrize( + "size_info", + [ + (True,), + (mssql_python.SQL_WVARCHAR, True, 0), + (mssql_python.SQL_DECIMAL, 18, True), + ], +) +def test_setinputsizes_rejects_boolean_sizes(size_info): + cursor = mssql_python.Cursor.__new__(mssql_python.Cursor) + with pytest.raises(ValueError): + cursor.setinputsizes([size_info]) + + +def test_setinputsizes_rejects_excessive_column_size(): + cursor = mssql_python.Cursor.__new__(mssql_python.Cursor) + with pytest.raises(ValueError, match="column size"): + cursor.setinputsizes([(mssql_python.SQL_VARCHAR, 1 << 40, 0)]) + + +def test_executemany_rejects_excessive_cumulative_parameter_buffers(cursor): + cursor.setinputsizes([(mssql_python.SQL_WVARCHAR, 50_000_000, 0)] * 3) + with pytest.raises(RuntimeError, match="Parameter buffers exceed the 256 MiB allocation limit"): + cursor.executemany("SELECT ?, ?, ?", [(None, None, None)]) + + +def test_executemany_rejects_text_for_binary_parameter(cursor): + cursor.setinputsizes([(mssql_python.SQL_VARBINARY, 10, 0)]) + with pytest.raises(RuntimeError, match="object type does not match"): + cursor.executemany("SELECT CAST(? AS VARBINARY(10))", [("text",)]) + + +def test_fetchmany_rejects_excessive_native_buffer(cursor): + cursor.execute("SELECT CAST('x' AS VARCHAR(8000))") + with pytest.raises(RuntimeError, match="256 MiB allocation limit"): + cursor.fetchmany(100_000) + assert cursor.fetchone()[0] == "x" + + +def test_fetchall_clamps_wide_result_batch_to_native_buffer_budget(cursor): + cursor.execute("SELECT " + ", ".join("CAST(N'x' AS NVARCHAR(4000))" for _ in range(34))) + assert [tuple(row) for row in cursor.fetchall()] == [("x",) * 34] + + def test_description(cursor): """Test description""" cursor.execute("SELECT * FROM #pytest_all_data_types WHERE id = 1") diff --git a/tests/test_004_cursor_arrow.py b/tests/test_004_cursor_arrow.py index 37a336789..859f7a630 100644 --- a/tests/test_004_cursor_arrow.py +++ b/tests/test_004_cursor_arrow.py @@ -225,11 +225,25 @@ def test_arrow_empty_fetch(cursor: mssql_python.Cursor): _test_arrow_test_data(cursor, [col_data], fetch_length=0) +def test_arrow_wide_schema_only_and_single_row_fit_native_buffer_budget( + cursor: mssql_python.Cursor, +): + columns = ", ".join(f"CAST(N'x' AS NVARCHAR(4000)) AS col_{index}" for index in range(600)) + cursor.execute(f"SELECT {columns}") + + schema_batch = cursor.arrow_batch(0) + assert schema_batch.num_rows == 0 + assert schema_batch.num_columns == 600 + + data_batch = cursor.arrow_batch(1) + assert data_batch.num_rows == 1 + assert data_batch.num_columns == 600 + + def test_arrow_table_batchsize_negative(cursor: mssql_python.Cursor): - tbl = cursor.execute("select 1 a").arrow(batch_size=-42) - assert type(tbl) is pa.Table - assert tbl.num_rows == 0 - assert tbl.num_columns == 1 + cursor.execute("select 1 a") + with pytest.raises(ValueError, match="batch_size"): + cursor.arrow(batch_size=-42) assert cursor.fetchone()[0] == 1 diff --git a/tests/test_010_pybind_functions.py b/tests/test_010_pybind_functions.py index 106b64ca3..41e175af4 100644 --- a/tests/test_010_pybind_functions.py +++ b/tests/test_010_pybind_functions.py @@ -37,6 +37,71 @@ ) +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +@pytest.mark.parametrize("fetch_size", [-1, 1_000_001, True]) +def test_fetchmany_rejects_unsafe_size_before_handle_access(fetch_size): + with pytest.raises(RuntimeError, match="Fetch size"): + ddbc.DDBCSQLFetchMany(None, [], fetch_size) + + +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +@pytest.mark.parametrize("batch_size", [-1, 1_000_001, True]) +def test_arrow_batch_rejects_unsafe_size_before_handle_access(batch_size): + with pytest.raises(RuntimeError, match="Arrow batch size"): + ddbc.DDBCSQLFetchArrowBatch(None, [], batch_size, 0) + + +@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): + with pytest.raises(RuntimeError, match="Parameter set size"): + ddbc.SQLExecuteMany(None, "", [], [], param_set_size, {}) + + +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +@pytest.mark.parametrize( + ("input_sizes", "message"), + [ + ((12, -8, 1, 0), "inputSizes must be None or a list"), + ([None], "four-item tuple"), + ([(12, -8, 1)], "four-item tuple"), + ([(12, -8, True, 0)], "must be integers"), + ([(12, -8, -1, 0)], "column size must be non-negative"), + ([(12, -8, 2**128, 0)], "column size is out of range"), + ([(12, -8, 1, 2**128)], "decimal digits are out of range"), + ], +) +def test_execute_rejects_malformed_input_sizes_before_handle_access(input_sizes, message): + with pytest.raises((TypeError, ValueError), match=message): + ddbc.DDBCSQLExecute(None, "", [], input_sizes, [False], True, {}) + + +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +def test_lob_buffer_growth_is_geometric(): + growth_count, data_size, allocation_size, reserved_bytes = ddbc._test_lob_buffer_growth( + 64, 64 * 1024 + ) + assert data_size == 4 * 1024 * 1024 + assert allocation_size == data_size + assert reserved_bytes == allocation_size + assert growth_count == 7 + + +@pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") +@pytest.mark.parametrize( + ("chunk_count", "chunk_size", "message"), + [ + (1_000_001, 1, "chunk count"), + (True, 1, "chunk count"), + (1, 0, "chunk size"), + (1, True, "chunk size"), + ], +) +def test_lob_buffer_growth_helper_rejects_unsafe_sizes(chunk_count, chunk_size, message): + with pytest.raises((RuntimeError, TypeError), match=message): + ddbc._test_lob_buffer_growth(chunk_count, chunk_size) + + @pytest.mark.skipif(not DDBC_AVAILABLE, reason="ddbc_bindings not available") class TestPybindModuleInfo: """Test module information and architecture detection.""" diff --git a/tests/test_023_execute_path_parity.py b/tests/test_023_execute_path_parity.py index d9bada323..daeb42514 100644 --- a/tests/test_023_execute_path_parity.py +++ b/tests/test_023_execute_path_parity.py @@ -296,6 +296,61 @@ def test_setinputsizes_text_binding_normalizes_time(cursor): ) +@pytest.mark.parametrize( + ("value", "sql_type"), + [ + ("x" * 4001, ddbc_sql_const.SQL_VARCHAR.value), + ("\U0001f600" * 3000, ddbc_sql_const.SQL_WVARCHAR.value), + (b"x" * 8001, ddbc_sql_const.SQL_VARBINARY.value), + ], +) +def test_setinputsizes_uses_actual_value_size_for_dae(cursor, value, sql_type): + """A too-small declared size cannot bypass streaming for a large value.""" + assert _override_roundtrip(cursor, value, sql_type, 1) == value + + +@pytest.mark.parametrize( + ("value", "sql_type"), + [ + pytest.param("x" * 4001, ddbc_sql_const.SQL_VARCHAR.value, id="varchar-encoded-bytes"), + pytest.param( + "\U0001f600" * 2001, + ddbc_sql_const.SQL_WVARCHAR.value, + id="nvarchar-utf16-units", + ), + ], +) +def test_executemany_setinputsizes_uses_actual_text_size_for_dae(cursor, value, sql_type): + """A too-small declared size cannot bypass streaming for an array-bound text value.""" + table_name = ( + "#dae_batch_varchar" + if sql_type == ddbc_sql_const.SQL_VARCHAR.value + else "#dae_batch_nvarchar" + ) + cursor.execute(f"CREATE TABLE {table_name} (value NVARCHAR(MAX))") + cursor.setinputsizes([(sql_type, 1, 0)]) + cursor.executemany(f"INSERT INTO {table_name} (value) VALUES (?)", [(value,), ("short",)]) + cursor.execute(f"SELECT value FROM {table_name} ORDER BY LEN(value) DESC") + assert [tuple(row) for row in cursor.fetchall()] == [(value,), ("short",)] + + +def test_executemany_setinputsizes_scans_mixed_text_values_for_dae(cursor): + """A bytes sample cannot hide a later oversized string from DAE selection.""" + value = "x" * 4001 + cursor.execute("CREATE TABLE #mixed_dae_batch (value VARCHAR(MAX))") + cursor.setinputsizes([(ddbc_sql_const.SQL_VARCHAR.value, 1, 0)]) + cursor.executemany("INSERT INTO #mixed_dae_batch (value) VALUES (?)", [(b"x",), (value,)]) + cursor.execute("SELECT value FROM #mixed_dae_batch ORDER BY LEN(value) DESC") + assert [tuple(row) for row in cursor.fetchall()] == [(value,), ("x",)] + + +def test_setinputsizes_rejects_bytes_for_wide_character_binding(cursor): + """Wide-character overrides require text rather than raw encoded bytes.""" + cursor.setinputsizes([(ddbc_sql_const.SQL_WVARCHAR.value, 1, 0)]) + with pytest.raises(TypeError, match="SQL_C_WCHAR"): + cursor.execute("SELECT ?", [b"x" * 8001]) + + @pytest.mark.parametrize("sql_type", [None, ddbc_sql_const.SQL_VARCHAR.value]) def test_time_isoformat_must_return_string(cursor, sql_type): """Native time normalization rejects a broken subclass contract on either path.""" diff --git a/tests/test_024_bulkcopy_arrow.py b/tests/test_024_bulkcopy_arrow.py index a840853ba..c8718284b 100644 --- a/tests/test_024_bulkcopy_arrow.py +++ b/tests/test_024_bulkcopy_arrow.py @@ -202,6 +202,15 @@ def test_batch_size_negative(self): with pytest.raises(ValueError, match="batch_size"): _bare_cursor().bulkcopy_arrow("t", pa.table({"a": [1]}), batch_size=-1) + def test_batch_size_exceeds_native_row_limit(self): + with pytest.raises(ValueError, match="batch_size"): + _bare_cursor().bulkcopy_arrow("t", pa.table({"a": [1]}), batch_size=1_000_001) + + def test_batch_size_bool_rejected(self): + for flag in (False, True): + with pytest.raises(TypeError, match="batch_size"): + _bare_cursor().bulkcopy_arrow("t", pa.table({"a": [1]}), batch_size=flag) + def test_timeout_wrong_type(self): with pytest.raises(TypeError, match="timeout"): _bare_cursor().bulkcopy_arrow("t", pa.table({"a": [1]}), timeout="30")