From 5d9b64b99f146eaa165ad50da6d9cccfd595ff8d Mon Sep 17 00:00:00 2001 From: Jahnvi Thakkar Date: Fri, 25 Sep 2026 11:49:24 +0530 Subject: [PATCH] FIX: Resolve UDT parameter metadata before binding Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- mssql_python/cursor.py | 22 ++++- mssql_python/pybind/ddbc_bindings.cpp | 43 +++++++++ tests/test_017_spatial_types.py | 130 ++++++++++++++++++++++++++ 3 files changed, 191 insertions(+), 4 deletions(-) diff --git a/mssql_python/cursor.py b/mssql_python/cursor.py index 37551d1e0..0825ea1b5 100644 --- a/mssql_python/cursor.py +++ b/mssql_python/cursor.py @@ -1184,6 +1184,14 @@ def setinputsizes(self, sizes: List[Union[int, tuple]]) -> None: cursor.executemany(sql, params) Note: + SQL_SS_UDT parameters require a statement whose parameter metadata + SQL Server can describe. The driver discovers the UDT identity before + binding; this may require a metadata round trip. Temporary tables and + table variables may not be describable, and discovery errors are + propagated. For built-in spatial types, bind serialized bytes without + a SQL_SS_UDT override or use a SQL constructor such as + hierarchyid::Parse(?) with a text parameter instead. + When inserting NULL into BINARY/VARBINARY columns in temp tables (#table) or table variables, SQLDescribeParam cannot resolve the column type and falls back to SQL_VARCHAR. This causes an implicit conversion error from @@ -2686,15 +2694,21 @@ 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: - any_dae = True + if paraminfo.isDAE: + any_dae = True if any_dae: logger.debug( "DAE parameters detected. Falling back to row-by-row execution with streaming.", ) - for row in seq_of_parameters: - self.execute(operation, row) + input_sizes = self._inputsizes + try: + for row in seq_of_parameters: + # execute() consumes overrides; every streamed row needs them. + self._inputsizes = input_sizes + self.execute(operation, row) + finally: + self._reset_inputsizes() return # Process parameters into column-wise format with possible type conversions diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index ebd60fe4c..f410631d4 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -452,6 +452,45 @@ static void PreResolveUnknownNullTypes(SqlHandle& handle, SQLHANDLE hStmt, } } +static SQLRETURN PreResolveUdtTypes(SQLHANDLE hStmt, std::vector& paramInfos) { + bool hasUdt = false; + for (const auto& info : paramInfos) { + if (info.paramSQLType == SQL_SS_UDT) { + hasUdt = true; + break; + } + } + if (!hasUdt) return SQL_SUCCESS; + + // UDT identity lives in the IPD, not in the scalar describe cache. Describe + // unbound records on each execution, including reused statements and DAE rows. + SQLRETURN rc = SQLFreeStmt_ptr(hStmt, SQL_RESET_PARAMS); + if (!SQL_SUCCEEDED(rc)) { + LOG("PreResolveUdtTypes: SQL_RESET_PARAMS failed, rc=%d", rc); + return rc; + } + for (size_t i = 0; i < paramInfos.size(); ++i) { + if (paramInfos[i].paramSQLType != SQL_SS_UDT) continue; + SQLSMALLINT type, digits, nullable; + SQLULEN size; + { + py::gil_scoped_release release; + rc = SQLDescribeParam_ptr(hStmt, static_cast(i + 1), + &type, &size, &digits, &nullable); + } + if (!SQL_SUCCEEDED(rc)) { + LOG("PreResolveUdtTypes: SQLDescribeParam failed for param[%zu], rc=%d", i, rc); + return rc; + } + // ODBC requires SQL_SS_LENGTH_UNLIMITED (0), not a byte count above + // 8000, for large UDTs. Type detection has already selected streaming. + if (paramInfos[i].columnSize > MAX_INLINE_BINARY) { + paramInfos[i].columnSize = 0; + } + } + return SQL_SUCCESS; +} + // Given a list of parameters and their ParamInfo, calls SQLBindParameter on // each of them with appropriate arguments SQLRETURN BindParameters(SqlHandle& handle, SQLHANDLE hStmt, const py::list& params, @@ -463,6 +502,8 @@ SQLRETURN BindParameters(SqlHandle& handle, SQLHANDLE hStmt, const py::list& par "with %zu parameters", (void*)hStmt, params.size()); + SQLRETURN describeRc = PreResolveUdtTypes(hStmt, paramInfos); + if (!SQL_SUCCEEDED(describeRc)) return describeRc; // GH-627: resolve unknown NULL param SQL types before binding any param. PreResolveUnknownNullTypes(handle, hStmt, paramInfos, ¶ms); for (int paramIndex = 0; paramIndex < params.size(); paramIndex++) { @@ -2257,6 +2298,8 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list& std::vector> tempBuffers; try { + SQLRETURN describeRc = PreResolveUdtTypes(hStmt, paramInfos); + if (!SQL_SUCCEEDED(describeRc)) return describeRc; // GH-627: resolve unknown NULL array param SQL types before binding any param. PreResolveUnknownNullTypes(handle, hStmt, paramInfos); for (int paramIndex = 0; paramIndex < columnwise_params.size(); ++paramIndex) { diff --git a/tests/test_017_spatial_types.py b/tests/test_017_spatial_types.py index 7fcb09c68..53352f301 100644 --- a/tests/test_017_spatial_types.py +++ b/tests/test_017_spatial_types.py @@ -1,8 +1,10 @@ """Tests for SQL Server spatial types (geography, geometry, hierarchyid).""" import pytest +import uuid from decimal import Decimal import mssql_python +from mssql_python.constants import ConstantsDDBC as C # ==================== GEOGRAPHY TYPE TESTS ==================== @@ -824,3 +826,131 @@ def test_hierarchyid_invalid_parsing(cursor, db_connection): "1/2/", ) db_connection.rollback() + + +# ==================== UDT PARAMETER BINDING (GH-816) ==================== + + +@pytest.fixture +def udt_table(db_connection): + name = f"dbo.udt_parameters_{uuid.uuid4().hex}" + with db_connection.cursor() as cursor: + cursor.execute( + f"CREATE TABLE {name} " + "(id int, h hierarchyid NULL, g geometry NULL, geo geography NULL, b varbinary(10))" + ) + try: + yield name + finally: + with db_connection.cursor() as cursor: + cursor.execute(f"DROP TABLE {name}") + db_connection.commit() + + +@pytest.fixture +def udt_payloads(db_connection): + with db_connection.cursor() as cursor: + cursor.execute( + "SELECT CONVERT(varbinary(max), hierarchyid::Parse('/1/')), " + "CONVERT(varbinary(max), geometry::STGeomFromText('POINT(1 2)', 0)), " + "CONVERT(varbinary(max), geography::STGeomFromText('POINT(1 2)', 4326))" + ) + return tuple(cursor.fetchone()) + + +@pytest.mark.parametrize("method", ["execute", "executemany"]) +@pytest.mark.parametrize("value_kind", ["bytes", "bytearray", "null"]) +def test_explicit_udt_parameters(db_connection, udt_table, udt_payloads, method, value_kind): + if value_kind == "null": + values = [None] * 3 + elif value_kind == "bytearray": + values = [bytearray(value) for value in udt_payloads] + else: + values = list(udt_payloads) + sql = f"INSERT INTO {udt_table} (id, h, g, geo) VALUES (?, ?, ?, ?)" + sizes = [(C.SQL_INTEGER.value, 0, 0)] + [(C.SQL_SS_UDT.value, 8000, 0)] * 3 + with db_connection.cursor() as cursor: + for i in range(2): + cursor.setinputsizes(sizes) + params = [i, *values] + getattr(cursor, method)(sql, params if method == "execute" else [params]) + cursor.execute( + f"SELECT CONVERT(varbinary(max), h), CONVERT(varbinary(max), g), " + f"CONVERT(varbinary(max), geo) FROM {udt_table} ORDER BY id" + ) + expected = (None, None, None) if value_kind == "null" else udt_payloads + assert [tuple(row) for row in cursor.fetchall()] == [expected, expected] + + +def test_udt_array_mixed_nulls(db_connection, udt_table, udt_payloads): + with db_connection.cursor() as cursor: + cursor.setinputsizes([(C.SQL_SS_UDT.value, 8000, 0)] * 3) + cursor.executemany( + f"INSERT INTO {udt_table} (h, g, geo) VALUES (?, ?, ?)", + [(None, None, None), udt_payloads, (None, None, None), udt_payloads], + ) + cursor.execute(f"SELECT COUNT(*), COUNT(h), COUNT(g), COUNT(geo) FROM {udt_table}") + assert tuple(cursor.fetchone()) == (4, 2, 2, 2) + + +def test_udt_changed_statement_and_type(db_connection, udt_table, udt_payloads): + with db_connection.cursor() as cursor: + for column, payload in zip(("h", "g", "geo", "h"), (*udt_payloads, udt_payloads[0])): + cursor.setinputsizes([(C.SQL_SS_UDT.value, 8000, 0)]) + cursor.execute(f"INSERT INTO {udt_table} ({column}) VALUES (?)", [payload]) + cursor.execute(f"SELECT COUNT(h), COUNT(g), COUNT(geo) FROM {udt_table}") + assert tuple(cursor.fetchone()) == (2, 1, 1) + + +def test_udt_with_inferred_binary_null(db_connection, udt_table, udt_payloads): + with db_connection.cursor() as cursor: + for _ in range(2): + cursor.setinputsizes([(C.SQL_SS_UDT.value, 8000, 0)]) + with pytest.warns(Warning, match="Number of input sizes"): + cursor.execute( + f"INSERT INTO {udt_table} (h, b) VALUES (?, ?)", [udt_payloads[0], None] + ) + cursor.execute(f"SELECT h.ToString(), b FROM {udt_table}") + assert [tuple(row) for row in cursor.fetchall()] == [("/1/", None)] * 2 + + +@pytest.mark.parametrize("method", ["execute", "executemany"]) +def test_large_udt_parameters(db_connection, udt_table, method, monkeypatch): + wkt = "LINESTRING(" + ", ".join(f"{i} {i % 7}" for i in range(1000)) + ")" + with db_connection.cursor() as cursor: + cursor.execute("SELECT CONVERT(varbinary(max), geometry::STGeomFromText(?, 0))", [wkt]) + payload = cursor.fetchone()[0] + assert len(payload) > 8000 + cursor.setinputsizes([(C.SQL_SS_UDT.value, len(payload), 0)]) + sql = f"INSERT INTO {udt_table} (g) VALUES (?)" + params = [payload] if method == "execute" else [(payload,), (payload,)] + sizes = cursor._inputsizes + observed_sizes = [] + execute = cursor.execute + + def record_overrides(operation, parameters): + observed_sizes.append(cursor._inputsizes) + return execute(operation, parameters) + + with monkeypatch.context() as patch: + patch.setattr(cursor, "execute", record_overrides) + getattr(cursor, method)(sql, params) + expected_rows = 1 if method == "execute" else 2 + assert observed_sizes == [sizes] * expected_rows + assert cursor._inputsizes is None + cursor.execute(f"SELECT CONVERT(varbinary(max), g) FROM {udt_table}") + assert [row[0] for row in cursor.fetchall()] == [payload] * expected_rows + + +@pytest.mark.parametrize("method", ["execute", "executemany"]) +def test_udt_discovery_error_and_recovery(db_connection, udt_table, udt_payloads, method): + with db_connection.cursor() as cursor: + cursor.setinputsizes([(C.SQL_SS_UDT.value, 8000, 0)]) + missing = f"dbo.udt_missing_{uuid.uuid4().hex}" + params = [udt_payloads[0]] if method == "execute" else [(udt_payloads[0],)] + with pytest.raises(mssql_python.DatabaseError, match="Invalid object name"): + getattr(cursor, method)(f"INSERT INTO {missing} (h) VALUES (?)", params) + cursor.setinputsizes([(C.SQL_SS_UDT.value, 8000, 0)]) + getattr(cursor, method)(f"INSERT INTO {udt_table} (h) VALUES (?)", params) + cursor.execute(f"SELECT h.ToString() FROM {udt_table}") + assert cursor.fetchone()[0] == "/1/"