Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
22 changes: 18 additions & 4 deletions mssql_python/cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
43 changes: 43 additions & 0 deletions mssql_python/pybind/ddbc_bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -452,6 +452,45 @@ static void PreResolveUnknownNullTypes(SqlHandle& handle, SQLHANDLE hStmt,
}
}

static SQLRETURN PreResolveUdtTypes(SQLHANDLE hStmt, std::vector<ParamInfo>& 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<SQLUSMALLINT>(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,
Expand All @@ -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, &params);
for (int paramIndex = 0; paramIndex < params.size(); paramIndex++) {
Expand Down Expand Up @@ -2257,6 +2298,8 @@ SQLRETURN BindParameterArray(SqlHandle& handle, SQLHANDLE hStmt, const py::list&
std::vector<std::shared_ptr<void>> 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) {
Expand Down
130 changes: 130 additions & 0 deletions tests/test_017_spatial_types.py
Original file line number Diff line number Diff line change
@@ -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 ====================

Expand Down Expand Up @@ -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/"
Loading