Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
36 commits
Select commit Hold shift + click to select a range
15e4092
FIX: Validate caller-controlled native buffer sizes
gargsaumya Sep 21, 2026
39dea1b
FIX: Update negative Arrow batch size test
gargsaumya Sep 21, 2026
0110fef
Merge branch 'main' into saumya/native-size-validation
gargsaumya Sep 21, 2026
a02fdf4
FIX: Address native size review findings
gargsaumya Sep 21, 2026
44982a9
Merge remote-tracking branch 'origin/saumya/native-size-validation' i…
gargsaumya Sep 21, 2026
7519aa4
Merge branch 'main' into saumya/native-size-validation
gargsaumya Sep 22, 2026
3038841
FIX: Bound native parameter and Arrow source buffers
gargsaumya Sep 22, 2026
710ff12
FIX: Address native size validation review
gargsaumya Sep 25, 2026
597d2c2
Merge remote-tracking branch 'origin/main' into saumya/native-size-va…
Copilot Sep 25, 2026
a096689
FIX: Close remaining native size validation gaps
gargsaumya Oct 7, 2026
231ae90
FIX: Measure override payloads in native units
gargsaumya Oct 7, 2026
f75aef7
Merge main into native size validation
gargsaumya Oct 7, 2026
444fb7f
FIX: Enforce peak native fetch memory limit
gargsaumya Oct 7, 2026
d5b4431
FIX: Bound streamed LOB native memory
gargsaumya Oct 7, 2026
3b60b8c
FIX: Validate executemany payload sizes
gargsaumya Oct 7, 2026
9f5c11a
FIX: Copy encoded wide text as bytes
gargsaumya Oct 7, 2026
5b26c36
FIX: Stream DAE payloads in bounded chunks
gargsaumya Oct 7, 2026
3a57124
FIX: Complete bounded DAE streaming
gargsaumya Oct 7, 2026
c64fbdd
FIX: Align DAE and array metadata
gargsaumya Oct 7, 2026
b8a8873
FIX: Satisfy strict DAE return analysis
gargsaumya Oct 7, 2026
9a55372
FIX: Close remaining native size gaps
gargsaumya Oct 7, 2026
7bb3a1a
FIX: Restore driver-compatible DAE binding
gargsaumya Oct 7, 2026
37ec527
FIX: Bound wide LOB conversion memory
gargsaumya Oct 7, 2026
7b5a3b5
FIX: Align streamed parameter metadata
gargsaumya Oct 7, 2026
190760f
FIX: Preserve MAX parameter type semantics
gargsaumya Oct 7, 2026
5d0a2e3
FIX: Size text by active C binding
gargsaumya Oct 7, 2026
e98d558
FIX: Bind NULL rows independently in DAE batches
gargsaumya Oct 7, 2026
99e837c
FIX: Validate row-local DAE tokens
gargsaumya Oct 7, 2026
a72841c
FIX: Use MAX sentinel for oversized DAE
gargsaumya Oct 7, 2026
5919670
FIX: Size DAE metadata through MAX boundary
gargsaumya Oct 7, 2026
db18317
FIX: align DAE metadata with wide bindings
gargsaumya Oct 7, 2026
400bd0f
FIX: preserve actual parameter metadata size
gargsaumya Oct 7, 2026
26c4763
FIX: address native validation review findings
gargsaumya Oct 7, 2026
f533230
FIX: reset DAE bindings after each row
gargsaumya Oct 7, 2026
48a2f23
FIX: restore geometric LOB buffer growth
gargsaumya Oct 8, 2026
d0912f6
FIX: bound LOB growth test helper
gargsaumya Oct 8, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
162 changes: 135 additions & 27 deletions mssql_python/cursor.py

Large diffs are not rendered by default.

1,041 changes: 775 additions & 266 deletions mssql_python/pybind/ddbc_bindings.cpp

Large diffs are not rendered by default.

220 changes: 195 additions & 25 deletions mssql_python/pybind/param_detect.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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();
Expand All @@ -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<SQLSMALLINT>::min() ||
sqlType > std::numeric_limits<SQLSMALLINT>::max() ||
cType < std::numeric_limits<SQLSMALLINT>::min() ||
cType > std::numeric_limits<SQLSMALLINT>::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<unsigned long long>(-1) && PyErr_Occurred()) {
PyErr_Clear();
throw py::value_error("column size is out of range");
}
if (requestedSize > std::numeric_limits<SQLULEN>::max()) {
throw py::value_error("column size is out of range");
}
const unsigned long long requestedDigits =
PyLong_AsUnsignedLongLong(decimalDigits);
if (requestedDigits == static_cast<unsigned long long>(-1) && PyErr_Occurred()) {
PyErr_Clear();
throw py::value_error("decimal digits are out of range");
}
if (requestedDigits >
static_cast<unsigned long long>(std::numeric_limits<SQLSMALLINT>::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<py::tuple>(inputSize);
info.paramSQLType = values[0].cast<SQLSMALLINT>();
info.paramCType = values[1].cast<SQLSMALLINT>();
Expand Down Expand Up @@ -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;
Comment thread
gargsaumya marked this conversation as resolved.
if (!isNumeric && PyUnicode_Check(obj)) {
const SQLULEN actualSize = static_cast<SQLULEN>(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<SQLULEN>(
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);
Expand Down Expand Up @@ -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<ParamInfo> DetectParamTypes(PyObject* params, PyObject* inputSizes) {
inline std::vector<ParamInfo> 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<ParamInfo> infos(n);
Expand All @@ -312,7 +492,7 @@ inline std::vector<ParamInfo> DetectParamTypes(PyObject* params, PyObject* input
info.isDAE = false;

if (i < inputSizeCount) {
Comment thread
gargsaumya marked this conversation as resolved.
ApplyInputSizeOverride(params, PyList_GET_ITEM(inputSizes, i), i, info);
ApplyInputSizeOverride(params, PyList_GET_ITEM(inputSizes, i), i, info, charEncoding);
continue;
}

Expand Down Expand Up @@ -398,16 +578,7 @@ inline std::vector<ParamInfo> 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.
Expand Down Expand Up @@ -439,16 +610,14 @@ inline std::vector<ParamInfo> 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;
Expand All @@ -467,7 +636,8 @@ inline std::vector<ParamInfo> 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<SQLULEN>(length));
info.dataPtr = borrow(obj);
} else {
info.columnSize = std::max<SQLULEN>(length, 1);
Expand Down
58 changes: 58 additions & 0 deletions tests/test_004_cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
22 changes: 18 additions & 4 deletions tests/test_004_cursor_arrow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
Loading
Loading