diff --git a/mssql_python/connection.py b/mssql_python/connection.py index 59496cb1f..1984f4979 100644 --- a/mssql_python/connection.py +++ b/mssql_python/connection.py @@ -540,6 +540,7 @@ def __init__( "ctype": ConstantsDDBC.SQL_WCHAR.value, }, } + self._decoding_generation = 0 # Auth type for acquiring fresh tokens at bulk copy time. # We intentionally do NOT cache the token — a fresh one is acquired @@ -782,6 +783,7 @@ def _token_factory(): # Initialize output converters dictionary and its lock for thread safety self._output_converters = {} + self._converters_generation = 0 self._converters_lock = threading.Lock() # Initialize encoding/decoding settings lock for thread safety @@ -1320,6 +1322,9 @@ def setdecoding( """ Sets the text decoding used when reading SQL_CHAR and SQL_WCHAR from the database. + Existing cursors refresh their cached SQL_CHAR/SQL_WCHAR decoding settings + before their next fetch. + This method configures how text data is decoded when reading from the database. In Python 3, all text is Unicode (str), so this primarily affects the encoding used to decode bytes from the database. @@ -1443,6 +1448,7 @@ def setdecoding( # Store the decoding settings for the specified sqltype (thread-safe with lock) with self._encoding_lock: self._decoding_settings[sqltype] = {"encoding": encoding, "ctype": ctype} + self._decoding_generation += 1 # Log with sanitized values for security sqltype_name = { @@ -1647,6 +1653,8 @@ def add_output_converter(self, sqltype: Union[int, type], func: Callable[[Any], Thread-safe implementation that protects the converters dictionary with a lock. + Changes apply on the next fetch, including for an already executed result set. + ⚠️ WARNING: Registering an output converter will cause the supplied Python function to be executed on every matching database value. Do not register converters from untrusted sources, as this can result in arbitrary code execution and security @@ -1687,6 +1695,7 @@ def add_output_converter(self, sqltype: Union[int, type], func: Callable[[Any], """ with self._converters_lock: self._output_converters[sqltype] = func + self._converters_generation += 1 # Pass to the underlying connection if native implementation supports it if hasattr(self._conn, "add_output_converter"): self._conn.add_output_converter(sqltype, func) @@ -1717,6 +1726,8 @@ def remove_output_converter(self, sqltype: Union[int, type]) -> None: Thread-safe implementation that protects the converters dictionary with a lock. + Existing cursors use the updated converters on their next fetch. + Args: sqltype (int or type): The SQL type value to remove the converter for @@ -1726,6 +1737,7 @@ def remove_output_converter(self, sqltype: Union[int, type]) -> None: with self._converters_lock: if sqltype in self._output_converters: del self._output_converters[sqltype] + self._converters_generation += 1 # Pass to the underlying connection if native implementation supports it if hasattr(self._conn, "remove_output_converter"): self._conn.remove_output_converter(sqltype) @@ -1737,11 +1749,14 @@ def clear_output_converters(self) -> None: Thread-safe implementation that protects the converters dictionary with a lock. + Existing cursors stop applying converters on their next fetch. + Returns: None """ with self._converters_lock: self._output_converters.clear() + self._converters_generation += 1 # Pass to the underlying connection if native implementation supports it if hasattr(self._conn, "clear_output_converters"): self._conn.clear_output_converters() diff --git a/mssql_python/cursor.py b/mssql_python/cursor.py index db45462c8..ad5eab798 100644 --- a/mssql_python/cursor.py +++ b/mssql_python/cursor.py @@ -62,6 +62,17 @@ } +def _string_only_output_converter(converter): + """Gate fallback conversion on the fetched value, not its column metadata.""" + + def convert(value): + if isinstance(value, (str, bytes)): + return converter(value) + return value + + return convert + + def _normalize_time_param(value, c_type): """Convert a datetime.time to its isoformat string when bound via text C-types. @@ -406,6 +417,7 @@ def __init__(self, connection: "Connection", timeout: int = 0) -> None: self._cached_column_map = None self._cached_column_map_lower = None self._cached_converter_map = None + self._cached_converters_generation = self._connection._converters_generation # Canonical, order-preserving column names snapshotted once per result set # and handed to each Row so mapping views never read the live cursor.description # (which changes when the cursor is reused for another query). _result_columns_src @@ -423,6 +435,7 @@ def __init__(self, connection: "Connection", timeout: int = 0) -> None: self._conn_native_uuid = getattr(self.connection, "_native_uuid", None) self._next_row_index = 0 # internal: index of the next row the driver will return (0-based) self._has_result_set = False # Track if we have an active result set + self._refresh_decoding_cache() self._skip_increment_for_next_fetch = ( False # Track if we need to skip incrementing the row index ) @@ -438,11 +451,7 @@ def _is_unicode_string(self, param: str) -> bool: Returns: True if the string contains non-ASCII characters, False otherwise. """ - try: - param.encode("ascii") - return False # Can be encoded to ASCII, so not Unicode - except UnicodeEncodeError: - return True # Contains non-ASCII characters, so treat as Unicode + return not param.isascii() def _parse_date(self, param: str) -> Optional[datetime.date]: """ @@ -617,6 +626,18 @@ def _get_encoding_settings(self): # This is the only case where defaults are appropriate (method doesn't exist) return {"encoding": "utf-16le", "ctype": ddbc_sql_const.SQL_WCHAR.value} + def _refresh_decoding_cache(self): + """Read decoding settings only when the connection configuration changes.""" + generation = self._connection._decoding_generation + char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) + wchar_encoding = self._get_decoding_settings(ddbc_sql_const.SQL_WCHAR.value).get( + "encoding", "utf-16le" + ) + self._cached_char_encoding = char_decoding.get("encoding", "utf-16le") + self._cached_char_ctype = char_decoding.get("ctype", ddbc_sql_const.SQL_WCHAR.value) + self._cached_wchar_encoding = wchar_encoding + self._cached_decoding_generation = generation + def _get_decoding_settings(self, sql_type): """ Get decoding settings for a specific SQL type. @@ -1235,45 +1256,51 @@ def _reset_inputsizes(self) -> None: """Reset input sizes after execution""" self._inputsizes = None + # Pre-built constant lookup table — avoids rebuilding ~30 entries on every call. + # Used by setinputsizes fallback path (PR #549 fast path doesn't need this). + _SQL_TO_C_TYPE = None + + @classmethod + def _get_sql_to_c_type_map(cls): + if cls._SQL_TO_C_TYPE is None: + cls._SQL_TO_C_TYPE = { + ddbc_sql_const.SQL_CHAR.value: ddbc_sql_const.SQL_C_CHAR.value, + ddbc_sql_const.SQL_VARCHAR.value: ddbc_sql_const.SQL_C_CHAR.value, + ddbc_sql_const.SQL_LONGVARCHAR.value: ddbc_sql_const.SQL_C_CHAR.value, + ddbc_sql_const.SQL_WCHAR.value: ddbc_sql_const.SQL_C_WCHAR.value, + ddbc_sql_const.SQL_WVARCHAR.value: ddbc_sql_const.SQL_C_WCHAR.value, + ddbc_sql_const.SQL_WLONGVARCHAR.value: ddbc_sql_const.SQL_C_WCHAR.value, + ddbc_sql_const.SQL_DECIMAL.value: ddbc_sql_const.SQL_C_NUMERIC.value, + ddbc_sql_const.SQL_NUMERIC.value: ddbc_sql_const.SQL_C_NUMERIC.value, + ddbc_sql_const.SQL_BIT.value: ddbc_sql_const.SQL_C_BIT.value, + ddbc_sql_const.SQL_TINYINT.value: ddbc_sql_const.SQL_C_TINYINT.value, + ddbc_sql_const.SQL_SMALLINT.value: ddbc_sql_const.SQL_C_SHORT.value, + ddbc_sql_const.SQL_INTEGER.value: ddbc_sql_const.SQL_C_LONG.value, + ddbc_sql_const.SQL_BIGINT.value: ddbc_sql_const.SQL_C_SBIGINT.value, + ddbc_sql_const.SQL_REAL.value: ddbc_sql_const.SQL_C_FLOAT.value, + ddbc_sql_const.SQL_FLOAT.value: ddbc_sql_const.SQL_C_DOUBLE.value, + ddbc_sql_const.SQL_DOUBLE.value: ddbc_sql_const.SQL_C_DOUBLE.value, + ddbc_sql_const.SQL_BINARY.value: ddbc_sql_const.SQL_C_BINARY.value, + ddbc_sql_const.SQL_VARBINARY.value: ddbc_sql_const.SQL_C_BINARY.value, + ddbc_sql_const.SQL_LONGVARBINARY.value: ddbc_sql_const.SQL_C_BINARY.value, + ddbc_sql_const.SQL_SS_UDT.value: ddbc_sql_const.SQL_C_BINARY.value, + ddbc_sql_const.SQL_TYPE_DATE.value: ddbc_sql_const.SQL_C_TYPE_DATE.value, + ddbc_sql_const.SQL_TYPE_TIME.value: ddbc_sql_const.SQL_C_TYPE_TIME.value, + ddbc_sql_const.SQL_TYPE_TIMESTAMP.value: ddbc_sql_const.SQL_C_TYPE_TIMESTAMP.value, + ddbc_sql_const.SQL_SS_TIME2.value: ddbc_sql_const.SQL_C_TYPE_TIME.value, + ddbc_sql_const.SQL_DATETIMEOFFSET.value: ddbc_sql_const.SQL_C_SS_TIMESTAMPOFFSET.value, + ddbc_sql_const.SQL_DATE.value: ddbc_sql_const.SQL_C_TYPE_DATE.value, + ddbc_sql_const.SQL_TIME.value: ddbc_sql_const.SQL_C_TYPE_TIME.value, + ddbc_sql_const.SQL_TIMESTAMP.value: ddbc_sql_const.SQL_C_TYPE_TIMESTAMP.value, + ddbc_sql_const.SQL_GUID.value: ddbc_sql_const.SQL_C_GUID.value, + ddbc_sql_const.SQL_SS_XML.value: ddbc_sql_const.SQL_C_WCHAR.value, + ddbc_sql_const.SQL_SS_VARIANT.value: ddbc_sql_const.SQL_C_BINARY.value, + } + return cls._SQL_TO_C_TYPE + def _get_c_type_for_sql_type(self, sql_type: int) -> int: """Map SQL type to appropriate C type for parameter binding.""" - sql_to_c_type = { - ddbc_sql_const.SQL_CHAR.value: ddbc_sql_const.SQL_C_CHAR.value, - ddbc_sql_const.SQL_VARCHAR.value: ddbc_sql_const.SQL_C_CHAR.value, - ddbc_sql_const.SQL_LONGVARCHAR.value: ddbc_sql_const.SQL_C_CHAR.value, - ddbc_sql_const.SQL_WCHAR.value: ddbc_sql_const.SQL_C_WCHAR.value, - ddbc_sql_const.SQL_WVARCHAR.value: ddbc_sql_const.SQL_C_WCHAR.value, - ddbc_sql_const.SQL_WLONGVARCHAR.value: ddbc_sql_const.SQL_C_WCHAR.value, - ddbc_sql_const.SQL_DECIMAL.value: ddbc_sql_const.SQL_C_NUMERIC.value, - ddbc_sql_const.SQL_NUMERIC.value: ddbc_sql_const.SQL_C_NUMERIC.value, - ddbc_sql_const.SQL_BIT.value: ddbc_sql_const.SQL_C_BIT.value, - ddbc_sql_const.SQL_TINYINT.value: ddbc_sql_const.SQL_C_TINYINT.value, - ddbc_sql_const.SQL_SMALLINT.value: ddbc_sql_const.SQL_C_SHORT.value, - ddbc_sql_const.SQL_INTEGER.value: ddbc_sql_const.SQL_C_LONG.value, - ddbc_sql_const.SQL_BIGINT.value: ddbc_sql_const.SQL_C_SBIGINT.value, - ddbc_sql_const.SQL_REAL.value: ddbc_sql_const.SQL_C_FLOAT.value, - ddbc_sql_const.SQL_FLOAT.value: ddbc_sql_const.SQL_C_DOUBLE.value, - ddbc_sql_const.SQL_DOUBLE.value: ddbc_sql_const.SQL_C_DOUBLE.value, - ddbc_sql_const.SQL_BINARY.value: ddbc_sql_const.SQL_C_BINARY.value, - ddbc_sql_const.SQL_VARBINARY.value: ddbc_sql_const.SQL_C_BINARY.value, - ddbc_sql_const.SQL_LONGVARBINARY.value: ddbc_sql_const.SQL_C_BINARY.value, - ddbc_sql_const.SQL_SS_UDT.value: ddbc_sql_const.SQL_C_BINARY.value, - # ODBC 3.x date/time types (reported by ODBC 18 driver) - ddbc_sql_const.SQL_TYPE_DATE.value: ddbc_sql_const.SQL_C_TYPE_DATE.value, - ddbc_sql_const.SQL_TYPE_TIME.value: ddbc_sql_const.SQL_C_TYPE_TIME.value, - ddbc_sql_const.SQL_TYPE_TIMESTAMP.value: ddbc_sql_const.SQL_C_TYPE_TIMESTAMP.value, - ddbc_sql_const.SQL_SS_TIME2.value: ddbc_sql_const.SQL_C_TYPE_TIME.value, - ddbc_sql_const.SQL_DATETIMEOFFSET.value: ddbc_sql_const.SQL_C_SS_TIMESTAMPOFFSET.value, - # ODBC 2.x aliases (accepted by setinputsizes via SQLTypes) - ddbc_sql_const.SQL_DATE.value: ddbc_sql_const.SQL_C_TYPE_DATE.value, - ddbc_sql_const.SQL_TIME.value: ddbc_sql_const.SQL_C_TYPE_TIME.value, - ddbc_sql_const.SQL_TIMESTAMP.value: ddbc_sql_const.SQL_C_TYPE_TIMESTAMP.value, - # Other types - ddbc_sql_const.SQL_GUID.value: ddbc_sql_const.SQL_C_GUID.value, - ddbc_sql_const.SQL_SS_XML.value: ddbc_sql_const.SQL_C_WCHAR.value, - ddbc_sql_const.SQL_SS_VARIANT.value: ddbc_sql_const.SQL_C_BINARY.value, - } - return sql_to_c_type.get(sql_type, ddbc_sql_const.SQL_C_DEFAULT.value) + return self._get_sql_to_c_type_map().get(sql_type, ddbc_sql_const.SQL_C_DEFAULT.value) def _create_parameter_types_list( # pylint: disable=too-many-arguments,too-many-positional-arguments self, @@ -1366,14 +1393,20 @@ def _build_converter_map(self): """ Build a pre-computed converter map for output converters. Returns a list where each element is either a converter function or None. + An empty tuple means no converters apply; None is reserved for uncached + direct Row construction and its legacy connection lookup. + String fallback converters check each value's type because variant and + unknown SQL types can have str metadata but non-string values. This eliminates the need to look up converters for every row. """ + generation = self._connection._converters_generation if ( not self.description or not hasattr(self.connection, "_output_converters") or not self.connection._output_converters ): - return None + self._cached_converters_generation = generation + return () sql_type_codes = self._column_sql_types converter_map = [] @@ -1391,17 +1424,17 @@ def _build_converter_map(self): # (e.g. decimal.Decimal) - the pre-existing mssql-python key style. if converter is None: converter = self.connection.get_output_converter(desc[1]) - # 3) Legacy WVARCHAR fallback: only apply it when the column's mapped type - # is str/bytes, so a registered SQL_WVARCHAR converter is never used as an - # unconditional catch-all for INT/DECIMAL/DATE/etc. columns (GH #691). This - # mirrors the isinstance(value, (str, bytes)) gate in - # Row._apply_output_converters. + # 3) The WVARCHAR fallback must also check the value: SQL_VARIANT and + # unknown types map to str but can return non-string Python values. if converter is None and desc[1] in (str, bytes): converter = self.connection.get_output_converter(ddbc_sql_const.SQL_WVARCHAR.value) + if converter: + converter = _string_only_output_converter(converter) converter_map.append(converter) - return converter_map + self._cached_converters_generation = generation + return converter_map if any(converter is not None for converter in converter_map) else () def _compute_uuid_str_indices(self): """ @@ -1451,7 +1484,9 @@ def _get_column_and_converter_maps(self): # Fallback to legacy column name map if no cached map column_map = column_map or getattr(self, "_column_name_map", None) - # Get cached converter map + # Refresh once per settings change, not once per row. + if self._cached_converters_generation != self._connection._converters_generation: + self._cached_converter_map = self._build_converter_map() converter_map = getattr(self, "_cached_converter_map", None) # Snapshot canonical column names once per result set (identity-tracked against @@ -2785,8 +2820,10 @@ def fetchone(self) -> Union[None, Row]: """ self._check_closed() # Check if the cursor is closed - char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) - wchar_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_WCHAR.value) + if self._cached_decoding_generation != self._connection._decoding_generation: + self._refresh_decoding_cache() + char_enc = self._cached_char_encoding + wchar_enc = self._cached_wchar_encoding # Fetch raw data row_data = [] @@ -2795,19 +2832,20 @@ def fetchone(self) -> Union[None, Row]: ret = ddbc_bindings.DDBCSQLFetchOne( self.hstmt, row_data, - char_decoding.get("encoding", "utf-16le"), - wchar_decoding.get("encoding", "utf-16le"), - char_decoding.get("ctype", ddbc_sql_const.SQL_WCHAR.value), + char_enc, + wchar_enc, + self._cached_char_ctype, ) + check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret) with perf_phase("py::fetchone::diag_records"): + # The native bridge's final status can mask earlier fetch warnings. if self.hstmt: self.messages.extend(ddbc_bindings.DDBCSQLGetAllDiagRecords(self.hstmt)) if ret == ddbc_sql_const.SQL_NO_DATA.value: # No more data available if self._next_row_index == 0 and self.description is not None: - # This is an empty result set, set rowcount to 0 self.rowcount = 0 return None @@ -2823,6 +2861,10 @@ def fetchone(self) -> Union[None, Row]: # Get column and converter maps column_map, converter_map, column_map_lower = self._get_column_and_converter_maps() with perf_phase("py::fetchone::row_wrap"): + if not converter_map and not self._uuid_str_indices: + return Row._fast_create( + row_data, column_map, self, column_map_lower, self._cached_result_columns + ) return Row( row_data, column_map, @@ -2856,8 +2898,10 @@ def fetchmany(self, size: Optional[int] = None) -> List[Row]: if size <= 0: return [] - char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) - wchar_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_WCHAR.value) + if self._cached_decoding_generation != self._connection._decoding_generation: + self._refresh_decoding_cache() + char_enc = self._cached_char_encoding + wchar_enc = self._cached_wchar_encoding # Fetch raw data rows_data = [] @@ -2867,11 +2911,12 @@ def fetchmany(self, size: Optional[int] = None) -> List[Row]: self.hstmt, rows_data, size, - char_decoding.get("encoding", "utf-16le"), - wchar_decoding.get("encoding", "utf-16le"), - char_decoding.get("ctype", ddbc_sql_const.SQL_WCHAR.value), + char_enc, + wchar_enc, + self._cached_char_ctype, ) + check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret) with perf_phase("py::fetchmany::diag_records"): if self.hstmt: self.messages.extend(ddbc_bindings.DDBCSQLGetAllDiagRecords(self.hstmt)) @@ -2894,6 +2939,15 @@ def fetchmany(self, size: Optional[int] = None) -> List[Row]: # Convert raw data to Row objects uuid_idx = self._uuid_str_indices with perf_phase("py::fetchmany::row_wrap"): + if not converter_map and not uuid_idx: + return ddbc_bindings.construct_rows( + rows_data, + Row, + column_map, + self, + column_map_lower, + self._cached_result_columns, + ) return [ Row( row_data, @@ -2921,8 +2975,10 @@ def fetchall(self) -> List[Row]: if not self._has_result_set and self.description: self._reset_rownumber() - char_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_CHAR.value) - wchar_decoding = self._get_decoding_settings(ddbc_sql_const.SQL_WCHAR.value) + if self._cached_decoding_generation != self._connection._decoding_generation: + self._refresh_decoding_cache() + char_enc = self._cached_char_encoding + wchar_enc = self._cached_wchar_encoding # Fetch raw data rows_data = [] @@ -2931,9 +2987,9 @@ def fetchall(self) -> List[Row]: ret = ddbc_bindings.DDBCSQLFetchAll( self.hstmt, rows_data, - char_decoding.get("encoding", "utf-16le"), - wchar_decoding.get("encoding", "utf-16le"), - char_decoding.get("ctype", ddbc_sql_const.SQL_WCHAR.value), + char_enc, + wchar_enc, + self._cached_char_ctype, ) # Check for errors @@ -2960,6 +3016,15 @@ def fetchall(self) -> List[Row]: # Convert raw data to Row objects uuid_idx = self._uuid_str_indices with perf_phase("py::fetchall::row_wrap"): + if not converter_map and not uuid_idx: + return ddbc_bindings.construct_rows( + rows_data, + Row, + column_map, + self, + column_map_lower, + self._cached_result_columns, + ) return [ Row( row_data, diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp index 1f235f829..c67723678 100644 --- a/mssql_python/pybind/ddbc_bindings.cpp +++ b/mssql_python/pybind/ddbc_bindings.cpp @@ -12,6 +12,7 @@ #include "param_detect.hpp" #include "py_ref.hpp" #include "py_type_cache.hpp" +#include "row_factory.hpp" #include "utf_utils.h" #include "fetch_text.hpp" @@ -3618,8 +3619,9 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_INTEGER: { SQLINTEGER intValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_LONG, &intValue, 0, NULL); - if (SQL_SUCCEEDED(ret)) { + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_LONG, &intValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { row.append(static_cast(intValue)); } else { row.append(py::none()); @@ -3628,7 +3630,12 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_SMALLINT: { SQLSMALLINT smallIntValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_SHORT, &smallIntValue, 0, NULL); + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_SHORT, &smallIntValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(static_cast(smallIntValue)); } else { @@ -3641,7 +3648,12 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_REAL: { SQLREAL realValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_FLOAT, &realValue, 0, NULL); + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_FLOAT, &realValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(realValue); } else { @@ -3714,7 +3726,12 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p case SQL_DOUBLE: case SQL_FLOAT: { SQLDOUBLE doubleValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_DOUBLE, &doubleValue, 0, NULL); + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_DOUBLE, &doubleValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(doubleValue); } else { @@ -3727,7 +3744,12 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_BIGINT: { SQLBIGINT bigintValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_SBIGINT, &bigintValue, 0, NULL); + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_SBIGINT, &bigintValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(static_cast(bigintValue)); } else { @@ -3740,9 +3762,10 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_TYPE_DATE: { SQL_DATE_STRUCT dateValue; - ret = - SQLGetData_ptr(hStmt, i, SQL_C_TYPE_DATE, &dateValue, sizeof(dateValue), NULL); - if (SQL_SUCCEEDED(ret)) { + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_TYPE_DATE, &dateValue, sizeof(dateValue), + &indicator); + if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) { row.append(PyTypeCache::get_date_class_obj()(dateValue.year, dateValue.month, dateValue.day)); } else { @@ -3772,8 +3795,13 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p case SQL_TYPE_TIMESTAMP: case SQL_DATETIME: { SQL_TIMESTAMP_STRUCT timestampValue; + SQLLEN indicator = 0; ret = SQLGetData_ptr(hStmt, i, SQL_C_TYPE_TIMESTAMP, ×tampValue, - sizeof(timestampValue), NULL); + sizeof(timestampValue), &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(PyTypeCache::get_datetime_class_obj()( timestampValue.year, timestampValue.month, timestampValue.day, @@ -3877,7 +3905,12 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_TINYINT: { SQLCHAR tinyIntValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_TINYINT, &tinyIntValue, 0, NULL); + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_TINYINT, &tinyIntValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(static_cast(tinyIntValue)); } else { @@ -3890,7 +3923,12 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p } case SQL_BIT: { SQLCHAR bitValue; - ret = SQLGetData_ptr(hStmt, i, SQL_C_BIT, &bitValue, 0, NULL); + SQLLEN indicator = 0; + ret = SQLGetData_ptr(hStmt, i, SQL_C_BIT, &bitValue, 0, &indicator); + if (SQL_SUCCEEDED(ret) && indicator == SQL_NULL_DATA) { + row.append(py::none()); + break; + } if (SQL_SUCCEEDED(ret)) { row.append(static_cast(bitValue)); } else { @@ -6222,6 +6260,14 @@ PYBIND11_MODULE(ddbc_bindings, m) { // Add a version attribute m.attr("__version__") = "1.0.0"; + // Fast Row construction in C++ — replaces Python list comprehension + m.def("construct_rows", &RowFactory::construct_rows, + "Build Row objects in C++ for fetchall/fetchmany fast path", + py::arg("rows_data"), py::arg("row_class"), + py::arg("column_map"), py::arg("cursor"), + py::arg("column_map_lower") = py::none(), + py::arg("column_names") = py::none()); + // Expose logger bridge function to Python m.def("update_log_level", &mssql_python::logging::LoggerBridge::updateLevel, "Update the cached log level in C++ bridge"); diff --git a/mssql_python/pybind/row_factory.hpp b/mssql_python/pybind/row_factory.hpp new file mode 100644 index 000000000..85d271128 --- /dev/null +++ b/mssql_python/pybind/row_factory.hpp @@ -0,0 +1,57 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#pragma once + +#include "py_ref.hpp" + +namespace RowFactory { + +// Wrap fetched values without converter or UUID processing. +// Accepts Row and its subclasses, bypassing __init__. +inline py::list construct_rows(const py::list& rows_data, const py::object& row_class, + const py::object& column_map, const py::object& cursor_obj, + const py::object& column_map_lower, const py::object& column_names) { + if (!PyType_Check(row_class.ptr())) { + throw py::type_error("row_class must be a type"); + } + PyTypeObject* row_type = reinterpret_cast(row_class.ptr()); + const py::object row_base = py::module_::import("mssql_python.row").attr("Row"); + if (!PyType_Check(row_base.ptr()) || + !PyType_IsSubtype(row_type, reinterpret_cast(row_base.ptr()))) { + throw py::type_error("row_class must be Row or a Row subclass"); + } + Py_ssize_t n = PyList_GET_SIZE(rows_data.ptr()); + + // Keep Python-owned names local to this call and its interpreter. + py::str attr_values("_values"); + py::str attr_column_map("_column_map"); + py::str attr_cursor("_cursor"); + py::str attr_column_map_lower("_column_map_lower"); + py::str attr_column_names("_column_names"); + + py::list result(n); + + for (Py_ssize_t i = 0; i < n; ++i) { + py::object row = steal(row_type->tp_alloc(row_type, 0)); + if (!row) + throw py::error_already_set(); + + PyObject* row_data = PyList_GET_ITEM(rows_data.ptr(), i); + + if (PyObject_GenericSetAttr(row.ptr(), attr_values.ptr(), row_data) < 0 || + PyObject_GenericSetAttr(row.ptr(), attr_column_map.ptr(), column_map.ptr()) < 0 || + PyObject_GenericSetAttr(row.ptr(), attr_cursor.ptr(), cursor_obj.ptr()) < 0 || + PyObject_GenericSetAttr(row.ptr(), attr_column_map_lower.ptr(), + column_map_lower.ptr()) < 0 || + PyObject_GenericSetAttr(row.ptr(), attr_column_names.ptr(), column_names.ptr()) < 0) { + throw py::error_already_set(); + } + + PyList_SET_ITEM(result.ptr(), i, row.release().ptr()); + } + + return result; +} + +} // namespace RowFactory diff --git a/mssql_python/row.py b/mssql_python/row.py index 0e7716a0d..79427067f 100644 --- a/mssql_python/row.py +++ b/mssql_python/row.py @@ -17,6 +17,9 @@ class Row: A row of data from a cursor fetch operation. Provides both tuple-like indexing and attribute access to column values. + Rows support user-defined attributes, ``vars(row)``, and weak references. + Internal fields use slots; ``vars(row)`` contains user-defined attributes. + For dict-like access, use the read-only ``row._mapping`` view (a ``collections.abc.Mapping`` of column name -> value). Iterating the Row itself (for x in row) yields values, not keys — consistent with pyodbc.Row and @@ -38,6 +41,33 @@ class Row: print(value) """ + # Slot internal fields while preserving dynamic attributes and weak references. + __slots__ = ( + "_values", + "_column_map", + "_cursor", + "_column_map_lower", + "_column_names", + "__dict__", + "__weakref__", + ) + + @staticmethod + def _fast_create(values, column_map, cursor, column_map_lower=None, column_names=None): + """Construct a Row bypassing __init__ — for the common fast path. + + Used by fetchall/fetchmany when no output converters and no UUID + stringification are needed (the vast majority of queries). Skips + the entire if/elif/else chain and keyword argument overhead in __init__. + """ + r = Row.__new__(Row) + r._values = values + r._column_map = column_map + r._cursor = cursor + r._column_map_lower = column_map_lower + r._column_names = column_names + return r + def __init__( self, values, @@ -54,7 +84,8 @@ def __init__( values: List of values for this row column_map: Pre-built column name to index mapping (shared across rows) cursor: Optional cursor reference (for backward compatibility and lowercase access) - converter_map: Pre-computed converter map (shared across rows for performance) + converter_map: Pre-computed converter map (shared across rows for performance). + An empty sequence skips converters; None enables the connection fallback. uuid_str_indices: Tuple of column indices whose uuid.UUID values should be converted to str. Pre-computed once per result set when native_uuid=False. None means no conversion (native_uuid=True, the default). @@ -67,22 +98,19 @@ def __init__( cursor snapshot; ``_mapping_keys()`` then reconstructs names from ``column_map``. """ - # Apply output converters if available using pre-computed converter map if converter_map: self._values = self._apply_output_converters_optimized(values, converter_map) elif ( - cursor + converter_map is None + and cursor and hasattr(cursor.connection, "_output_converters") and cursor.connection._output_converters ): - # Fallback to original method for backward compatibility + # Support direct Row construction without a pre-computed converter map. self._values = self._apply_output_converters(values, cursor) else: self._values = values - # Convert UUID columns to str when native_uuid=False. - # uuid_str_indices is pre-computed once at execute() time, so this is - # O(num_uuid_columns) per row — zero cost when native_uuid=True (the default). if uuid_str_indices: self._stringify_uuids(uuid_str_indices) diff --git a/tests/test_fetch_settings_cache.py b/tests/test_fetch_settings_cache.py new file mode 100644 index 000000000..58a1423f3 --- /dev/null +++ b/tests/test_fetch_settings_cache.py @@ -0,0 +1,892 @@ +""" +Copyright (c) Microsoft Corporation. +Licensed under the MIT license. + +Regression and operation-count tests for connection settings cached by fetch APIs. +All integration queries are read-only and each test owns its connection. +""" + +import datetime +import subprocess +import sys +import uuid +import weakref +from unittest.mock import Mock, patch + +import pytest +import mssql_python +from mssql_python.constants import ConstantsDDBC +from mssql_python.row import Row + +FETCH_METHODS = ("fetchone", "fetchmany", "fetchall") +SQL_WVARCHAR = ConstantsDDBC.SQL_WVARCHAR.value +SQL_SS_VARIANT = ConstantsDDBC.SQL_SS_VARIANT.value +UUID_TEXT = "00112233-4455-6677-8899-AABBCCDDEEFF" +MIXED_SELECT = ( + "SELECT CAST('abc' AS VARCHAR(10)) AS narrow, " + "CAST(N'def' AS NVARCHAR(10)) AS wide, CAST(42 AS INT) AS number, " + f"CAST('{UUID_TEXT}' AS UNIQUEIDENTIFIER) AS id, " + "CAST(0x0102 AS VARBINARY(2)) AS binary_value, CAST(NULL AS NVARCHAR(10)) AS empty_value" +) +MIXED_CONVERTER_INPUTS = [b"a\x00b\x00c\x00", b"d\x00e\x00f\x00", b"\x01\x02"] +VARIANT_SELECT = ( + "SELECT value FROM (VALUES " + "(1, CAST(42 AS SQL_VARIANT)), " + "(2, CAST(CAST('20260102' AS DATE) AS SQL_VARIANT)), " + f"(3, CAST(CAST('{UUID_TEXT}' AS UNIQUEIDENTIFIER) AS SQL_VARIANT)), " + "(4, CAST(CAST(N'abc' AS NVARCHAR(10)) AS SQL_VARIANT)), " + "(5, CAST(CAST(0x0102 AS VARBINARY(2)) AS SQL_VARIANT)), " + "(6, CAST(NULL AS SQL_VARIANT))) AS v(n, value) ORDER BY n" +) +VARIANT_VALUES = [ + 42, + datetime.date(2026, 1, 2), + uuid.UUID(UUID_TEXT), + "abc", + b"\x01\x02", + None, +] + + +@pytest.fixture +def connection(conn_str): + try: + conn = mssql_python.connect(conn_str, timeout=5) + except mssql_python.Error as error: + pytest.fail( + f"Connection failed: {type(error).__name__}; connection details withheld", + pytrace=False, + ) + try: + yield conn + finally: + conn.close() + + +def fetch_rows(cursor, method): + if method == "fetchone": + row = cursor.fetchone() + return [] if row is None else [row] + if method == "fetchmany": + return cursor.fetchmany(10) + return cursor.fetchall() + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("lob", (False, True), ids=("bound", "lob")) +@pytest.mark.parametrize( + ("sql_type", "literal", "expected"), + ( + ("INT", "42", 42), + ("SMALLINT", "42", 42), + ("BIGINT", "42", 42), + ("TINYINT", "42", 42), + ("BIT", "1", True), + ("REAL", "1.25", 1.25), + ("FLOAT", "1.25", 1.25), + ("DATE", "'20260102'", datetime.date(2026, 1, 2)), + ("DATETIME", "'20260102'", datetime.datetime(2026, 1, 2)), + ("DATETIME2", "'20260102'", datetime.datetime(2026, 1, 2)), + ("SMALLDATETIME", "'20260102'", datetime.datetime(2026, 1, 2)), + ), +) +def test_fixed_width_null_fetch(connection, method, lob, sql_type, literal, expected): + prefix = "CAST(N'payload' AS NVARCHAR(MAX)), " if lob else "" + with connection.cursor() as cursor: + cursor.execute( + f"SELECT {prefix}CAST(CASE WHEN n = 2 THEN NULL ELSE {literal} END AS {sql_type}) " + "AS value FROM (VALUES (1), (2)) AS v(n) ORDER BY n" + ) + rows = fetch_rows(cursor, method) + if method == "fetchone": + rows.extend(fetch_rows(cursor, method)) + assert len(rows) == 2 + assert rows[0][-1] == expected + assert type(rows[0][-1]) is type(expected) + assert rows[1][-1] is None + assert not cursor.messages + assert fetch_rows(cursor, method) == [] + cursor.execute(f"SELECT CAST({literal} AS {sql_type})") + assert cursor.fetchval() == expected + assert not cursor.messages + + +@pytest.mark.parametrize("expression", ("NULL", "OBJECT_ID('tempdb..#missing_fetch_null_table')")) +def test_fetchval_null_expression(connection, expression): + with connection.cursor() as cursor: + cursor.execute(f"SELECT {expression}") + assert cursor.fetchval() is None + assert not cursor.messages + cursor.execute("SELECT 42") + assert cursor.fetchval() == 42 + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("lob", (False, True), ids=("bound", "lob")) +@pytest.mark.parametrize("wide", (False, True), ids=("char", "wchar")) +@pytest.mark.parametrize("when", ("before_execute", "after_execute", "between_fetches")) +def test_decoding_changes_on_existing_cursor(connection, method, lob, wide, when): + sqltype = mssql_python.SQL_WCHAR if wide else mssql_python.SQL_CHAR + if wide: + # Native WCHAR decoding is fixed in this revision; these are value parity controls. + connection.setdecoding(sqltype, encoding="utf-16be") + size = "MAX" if lob else "1" + expression = ( + f"CAST(NCHAR(233) AS NVARCHAR({size}))" if wide else f"CONVERT(VARCHAR({size}), 0xE9)" + ) + encoding = "utf-16le" if wide else "latin-1" + + with connection.cursor() as cursor: + if when == "before_execute": + connection.setdecoding(sqltype, encoding=encoding) + cursor.execute(f"SELECT {expression} AS txt FROM (VALUES (1), (2)) AS v(n)") + if when == "between_fetches": + assert cursor.fetchone() is not None + if when != "before_execute": + connection.setdecoding(sqltype, encoding=encoding) + rows = fetch_rows(cursor, method) + assert rows + assert all(row.txt == "\u00e9" for row in rows) + + +@pytest.mark.parametrize( + ("method", "bridge_name"), + ( + ("fetchone", "DDBCSQLFetchOne"), + ("fetchmany", "DDBCSQLFetchMany"), + ("fetchall", "DDBCSQLFetchAll"), + ), +) +def test_wchar_decoding_forwarded_to_live_fetch_bridge(connection, method, bridge_name): + bridge = getattr(mssql_python.ddbc_bindings, bridge_name) + with ( + patch.object(connection, "getdecoding", wraps=connection.getdecoding) as reads, + patch.object(mssql_python.ddbc_bindings, bridge_name, wraps=bridge) as fetch, + connection.cursor() as cursor, + ): + assert reads.call_count == 2 + previous_encoding = "utf-16le" + expected_reads = 2 + encodings = ("utf-16le", "utf-16be", "utf-16be", "utf-16le", "utf-16le") + for index, encoding in enumerate(encodings, 1): + cursor.execute("SELECT CAST(NCHAR(233) AS NVARCHAR(1)) AS txt") + if encoding != previous_encoding: + connection.setdecoding(mssql_python.SQL_WCHAR, encoding=encoding) + expected_reads += 2 + assert fetch_rows(cursor, method)[0].txt == "\u00e9" + assert fetch.call_count == index + assert fetch.call_args.args[-3:] == ("utf-16le", encoding, mssql_python.SQL_WCHAR) + assert reads.call_count == expected_reads + previous_encoding = encoding + assert [call.args[0] for call in reads.call_args_list] == [ + mssql_python.SQL_CHAR, + mssql_python.SQL_WCHAR, + ] * 3 + + +def test_decoding_cache_reuse_and_multiple_cursors(connection): + with patch.object(connection, "getdecoding", wraps=connection.getdecoding) as reads: + with connection.cursor() as first, connection.cursor() as second: + assert reads.call_count == 4 + for cursor in (first, second): + cursor.execute("SELECT n FROM (VALUES (1), (2), (3)) AS v(n) ORDER BY n") + assert cursor.fetchone()[0] == 1 + assert cursor.fetchmany(1)[0][0] == 2 + assert cursor.fetchall()[0][0] == 3 + assert reads.call_count == 4 + + connection.setdecoding(mssql_python.SQL_CHAR, encoding="latin-1") + for index, cursor in enumerate((first, second), 1): + cursor.execute( + "SELECT CONVERT(VARCHAR(1), 0xE9) AS txt FROM (VALUES (1), (2), (3)) AS v(n)" + ) + assert cursor.fetchone().txt == "\u00e9" + assert cursor.fetchmany(1)[0].txt == "\u00e9" + assert cursor.fetchall()[0].txt == "\u00e9" + assert reads.call_count == 4 + 2 * index + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("native_uuid", (True, False)) +@pytest.mark.parametrize("mutation", ("add", "replace", "remove", "clear")) +def test_converter_changes_after_execute(connection, method, native_uuid, mutation, monkeypatch): + monkeypatch.setattr(mssql_python, "native_uuid", native_uuid) + original = Mock(side_effect=lambda raw: "original:" + raw.decode("utf-16-le")) + replacement = Mock(side_effect=lambda raw: "converted:" + raw.decode("utf-16-le")) + if mutation != "add": + connection.add_output_converter(SQL_WVARCHAR, original) + + with connection.cursor() as cursor: + cursor.execute( + "SELECT CAST(N'abc' AS NVARCHAR(10)) AS txt, " + f"CAST('{UUID_TEXT}' AS UNIQUEIDENTIFIER) AS id, " + "CAST(NULL AS NVARCHAR(10)) AS empty_value" + ) + if mutation in ("add", "replace"): + connection.add_output_converter(SQL_WVARCHAR, replacement) + elif mutation == "remove": + connection.remove_output_converter(SQL_WVARCHAR) + else: + connection.clear_output_converters() + + row = fetch_rows(cursor, method)[0] + assert row.txt == ("converted:abc" if mutation in ("add", "replace") else "abc") + assert row.id == (uuid.UUID(UUID_TEXT) if native_uuid else UUID_TEXT) + assert row.empty_value is None + assert original.call_count == 0 + if mutation in ("add", "replace"): + replacement.assert_called_once_with(b"a\x00b\x00c\x00") + else: + replacement.assert_not_called() + + +def test_converter_cache_reuse_between_fetches(connection): + converter = Mock(side_effect=lambda raw: "converted:" + raw.decode("utf-16-le")) + replacement = Mock(side_effect=lambda raw: "new:" + raw.decode("utf-16-le")) + with connection.cursor() as cursor: + with ( + patch.object( + cursor, "_build_converter_map", wraps=cursor._build_converter_map + ) as builds, + patch.object( + connection, "get_output_converter", wraps=connection.get_output_converter + ) as lookups, + patch.object(Row, "_fast_create", wraps=Row._fast_create) as fast_one, + patch.object( + mssql_python.ddbc_bindings, + "construct_rows", + wraps=mssql_python.ddbc_bindings.construct_rows, + ) as fast_batch, + ): + cursor.execute( + "SELECT CAST(N'abc' AS NVARCHAR(10)) AS txt " + "FROM (VALUES (1), (2), (3), (4), (5), (6), (7), (8), (9)) AS v(n)" + ) + assert builds.call_count == 1 + assert lookups.call_count == 0 + assert cursor.fetchone().txt == "abc" + assert fast_one.call_count == 1 + + connection.add_output_converter(SQL_WVARCHAR, converter) + assert cursor.fetchone().txt == "converted:abc" + assert [row.txt for row in cursor.fetchmany(2)] == ["converted:abc"] * 2 + assert converter.call_count == 3 + assert builds.call_count == 2 + assert lookups.call_count == 1 + + connection.add_output_converter(SQL_WVARCHAR, replacement) + assert cursor.fetchone().txt == "new:abc" + assert builds.call_count == 3 + assert lookups.call_count == 2 + assert replacement.call_count == 1 + + connection.remove_output_converter(SQL_WVARCHAR) + assert cursor.fetchmany(1)[0].txt == "abc" + assert builds.call_count == 4 + assert lookups.call_count == 2 + assert fast_batch.call_count == 1 + + connection.add_output_converter(SQL_WVARCHAR, converter) + assert cursor.fetchone().txt == "converted:abc" + assert builds.call_count == 5 + assert lookups.call_count == 3 + connection.clear_output_converters() + assert [row.txt for row in cursor.fetchall()] == ["abc", "abc"] + assert builds.call_count == 6 + assert lookups.call_count == 3 + assert fast_batch.call_count == 2 + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("native_uuid", (True, False)) +@pytest.mark.parametrize("mutation", ("add", "replace", "remove", "clear")) +def test_late_converter_mixed_values(connection, method, native_uuid, mutation, monkeypatch): + monkeypatch.setattr(mssql_python, "native_uuid", native_uuid) + original = Mock(return_value="original") + replacement = Mock(return_value="converted") + if mutation != "add": + connection.add_output_converter(SQL_WVARCHAR, original) + + with connection.cursor() as cursor: + cursor.execute(MIXED_SELECT) + if mutation in ("add", "replace"): + connection.add_output_converter(SQL_WVARCHAR, replacement) + elif mutation == "remove": + connection.remove_output_converter(SQL_WVARCHAR) + else: + connection.clear_output_converters() + row = fetch_rows(cursor, method)[0] + converted = mutation in ("add", "replace") + assert list(row) == [ + "converted" if converted else "abc", + "converted" if converted else "def", + 42, + uuid.UUID(UUID_TEXT) if native_uuid else UUID_TEXT, + "converted" if converted else b"\x01\x02", + None, + ] + original.assert_not_called() + assert [call.args[0] for call in replacement.call_args_list] == ( + MIXED_CONVERTER_INPUTS if converted else [] + ) + + +@pytest.mark.parametrize("native_uuid", (True, False)) +def test_late_converter_mixed_values_between_fetches(connection, native_uuid, monkeypatch): + monkeypatch.setattr(mssql_python, "native_uuid", native_uuid) + converter = Mock(return_value="converted") + replacement = Mock(return_value="new") + with connection.cursor() as cursor: + cursor.execute(MIXED_SELECT + " FROM (VALUES (1), (2), (3), (4), (5)) AS v(n)") + assert cursor.fetchone().number == 42 + connection.add_output_converter(SQL_WVARCHAR, converter) + first = cursor.fetchone() + connection.add_output_converter(SQL_WVARCHAR, replacement) + second = cursor.fetchmany(1)[0] + connection.remove_output_converter(SQL_WVARCHAR) + third = cursor.fetchone() + connection.add_output_converter(SQL_WVARCHAR, converter) + connection.clear_output_converters() + fourth = cursor.fetchall()[0] + for row, text in ((first, "converted"), (second, "new"), (third, "abc"), (fourth, "abc")): + assert row.narrow == text + assert row.number == 42 + assert row.id == (uuid.UUID(UUID_TEXT) if native_uuid else UUID_TEXT) + assert row.empty_value is None + assert [call.args[0] for call in converter.call_args_list] == MIXED_CONVERTER_INPUTS + assert [call.args[0] for call in replacement.call_args_list] == MIXED_CONVERTER_INPUTS + + +def test_preconfigured_converter_keeps_existing_fallback_semantics(connection): + converter = Mock(return_value="converted") + connection.add_output_converter(SQL_WVARCHAR, converter) + with connection.cursor() as cursor: + cursor.execute(MIXED_SELECT) + assert list(cursor.fetchone()) == [ + "converted", + "converted", + 42, + uuid.UUID(UUID_TEXT), + "converted", + None, + ] + assert [call.args[0] for call in converter.call_args_list] == MIXED_CONVERTER_INPUTS + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("when", ("before_execute", "after_execute", "between_fetches")) +def test_variant_string_fallback_is_value_gated(connection, method, when): + converter = Mock(return_value="converted") + with ( + connection.cursor() as cursor, + patch.object(cursor, "_build_converter_map", wraps=cursor._build_converter_map) as builds, + patch.object( + connection, "get_output_converter", wraps=connection.get_output_converter + ) as lookups, + ): + if when == "before_execute": + connection.add_output_converter(SQL_WVARCHAR, converter) + cursor.execute(VARIANT_SELECT) + expected = VARIANT_VALUES[:3] + ["converted", "converted", None] + if when == "between_fetches": + assert cursor.fetchone()[0] == 42 + expected = expected[1:] + if when != "before_execute": + connection.add_output_converter(SQL_WVARCHAR, converter) + rows = [] + while batch := fetch_rows(cursor, method): + rows.extend(batch) + assert [row[0] for row in rows] == expected + assert [type(row[0]) for row in rows] == [type(value) for value in expected] + assert [call.args[0] for call in converter.call_args_list] == [ + b"a\x00b\x00c\x00", + b"\x01\x02", + ] + assert builds.call_count == (1 if when == "before_execute" else 2) + assert lookups.call_count == 3 + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("explicit_type", ("sql", "python")) +def test_late_variant_explicit_converter_precedence(connection, method, explicit_type): + fallback = Mock(return_value="fallback") + python_converter = Mock(return_value="python") + sql_converter = Mock(return_value="sql") + with connection.cursor() as cursor: + cursor.execute(VARIANT_SELECT) + connection.add_output_converter(SQL_WVARCHAR, fallback) + connection.add_output_converter(str, python_converter) + if explicit_type == "sql": + connection.add_output_converter(SQL_SS_VARIANT, sql_converter) + rows = [] + while batch := fetch_rows(cursor, method): + rows.extend(batch) + assert [row[0] for row in rows] == [explicit_type] * 5 + [None] + selected = sql_converter if explicit_type == "sql" else python_converter + assert [call.args[0] for call in selected.call_args_list] == ( + VARIANT_VALUES[:3] + [b"a\x00b\x00c\x00", b"\x01\x02"] + ) + fallback.assert_not_called() + if explicit_type == "sql": + python_converter.assert_not_called() + else: + sql_converter.assert_not_called() + + +def test_unknown_sql_type_string_fallback_is_value_gated(connection): + converter = Mock(return_value="converted") + connection.add_output_converter(SQL_WVARCHAR, converter) + with connection.cursor() as cursor: + cursor._initialize_description( + [ + { + "ColumnName": "value", + "DataType": 123456, + "ColumnSize": 100, + "DecimalDigits": 0, + "Nullable": ConstantsDDBC.SQL_NULLABLE.value, + } + ] + ) + assert cursor.description[0][1] is str + converter_map = cursor._build_converter_map() + rows = [Row([value], {"value": 0}, converter_map=converter_map) for value in VARIANT_VALUES] + assert [row[0] for row in rows] == VARIANT_VALUES[:3] + ["converted", "converted", None] + assert [call.args[0] for call in converter.call_args_list] == [ + b"a\x00b\x00c\x00", + b"\x01\x02", + ] + + +def test_converter_cache_multiple_cursors_and_result_shapes(connection): + converter = Mock(side_effect=lambda raw: "converted:" + raw.decode("utf-16-le")) + with connection.cursor() as first, connection.cursor() as second: + for cursor in (first, second): + cursor.execute("SELECT CAST(N'abc' AS NVARCHAR(10)) AS txt") + assert cursor.fetchone().txt == "abc" + connection.add_output_converter(SQL_WVARCHAR, converter) + for cursor in (first, second): + with patch.object( + cursor, "_build_converter_map", wraps=cursor._build_converter_map + ) as builds: + assert cursor.fetchall() == [] + builds.assert_called_once_with() + cursor.execute( + "SELECT CAST(NULL AS INT) AS empty_value, CAST(N'def' AS NVARCHAR(10)) AS txt; " + "SELECT CAST(N'ghi' AS NVARCHAR(10)) AS renamed" + ) + row = cursor.fetchone() + assert list(row) == [None, "converted:def"] + assert cursor.nextset() + assert cursor.fetchone().renamed == "converted:ghi" + assert converter.call_count == 4 + + +@pytest.mark.parametrize("stringify_uuid", (False, True)) +def test_direct_row_converter_fallback(connection, stringify_uuid): + converter = Mock(side_effect=lambda raw: "converted:" + raw.decode("utf-16-le")) + with connection.cursor() as cursor: + cursor.execute( + "SELECT CAST(N'abc' AS NVARCHAR(10)) AS txt, " + f"CAST('{UUID_TEXT}' AS UNIQUEIDENTIFIER) AS id" + ) + values = list(cursor.fetchone()) + connection.add_output_converter(SQL_WVARCHAR, converter) + row = Row( + values, + {"txt": 0, "id": 1}, + cursor=cursor, + uuid_str_indices=(1,) if stringify_uuid else None, + ) + assert row.txt == "converted:abc" + assert row.id == (UUID_TEXT if stringify_uuid else uuid.UUID(UUID_TEXT)) + converter.assert_called_once_with(b"a\x00b\x00c\x00") + + +def test_direct_row_without_converters_is_zero_copy(): + values = [1, "abc", None] + row = Row(values, {"number": 0, "txt": 1, "empty_value": 2}) + assert row._values is values + + +def test_decoding_cache_refresh_failure_is_retried(connection): + with connection.cursor() as cursor: + cursor.execute("SELECT CONVERT(VARCHAR(1), 0xE9) AS txt") + generation = cursor._cached_decoding_generation + connection.setdecoding(mssql_python.SQL_CHAR, encoding="latin-1") + read_settings = connection.getdecoding + + def fail_wchar(sqltype): + if sqltype == mssql_python.SQL_WCHAR: + raise RuntimeError("injected settings failure") + return read_settings(sqltype) + + with patch.object(connection, "getdecoding", side_effect=fail_wchar): + with pytest.raises(RuntimeError, match="injected settings failure"): + cursor.fetchone() + assert cursor._cached_decoding_generation == generation + assert cursor.fetchone().txt == "\u00e9" + + +def test_converter_cache_refresh_failure_is_retried(connection): + with connection.cursor() as cursor: + cursor.execute("SELECT CAST(N'abc' AS NVARCHAR(10)) AS txt FROM (VALUES (1), (2)) AS v(n)") + generation = cursor._cached_converters_generation + connection.add_output_converter(SQL_WVARCHAR, lambda raw: "converted") + with patch.object( + connection, "get_output_converter", side_effect=RuntimeError("injected map failure") + ): + with pytest.raises(RuntimeError, match="injected map failure"): + cursor.fetchone() + assert cursor._cached_converters_generation == generation + assert cursor.fetchone().txt == "converted" + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("lowercase", (False, True)) +@pytest.mark.parametrize("processing", ("none", "converter", "uuid")) +def test_fetch_preserves_column_name_maps(connection, method, lowercase, processing, monkeypatch): + monkeypatch.setattr(mssql_python, "lowercase", lowercase) + monkeypatch.setattr(mssql_python, "native_uuid", processing != "uuid") + if processing == "converter": + connection.add_output_converter(SQL_WVARCHAR, lambda raw: "converted") + with connection.cursor() as cursor: + cursor.execute( + "SELECT CAST(N'abc' AS NVARCHAR(10)) AS MixedName, " + f"CAST('{UUID_TEXT}' AS UNIQUEIDENTIFIER) AS MixedId " + "FROM (VALUES (1), (2)) AS v(n)" + ) + rows = fetch_rows(cursor, method) + if method == "fetchone": + rows.extend(fetch_rows(cursor, method)) + row = rows[0] + name = "mixedname" if lowercase else "MixedName" + expected = "converted" if processing == "converter" else "abc" + assert row[name] == getattr(row, name) == expected + if lowercase: + assert row["MIXEDNAME"] == row.MIXEDNAME == expected + else: + with pytest.raises(KeyError): + row["MIXEDNAME"] + with pytest.raises(AttributeError): + row.MIXEDNAME + assert row._column_map_lower is cursor._cached_column_map_lower + assert row[1] == (UUID_TEXT if processing == "uuid" else uuid.UUID(UUID_TEXT)) + id_name = "mixedid" if lowercase else "MixedId" + names = cursor._cached_result_columns + assert names == (name, id_name) + assert len(rows) == 2 + assert all(item._column_names is names for item in rows) + expected_mapping = {name: expected, id_name: row[1]} + assert all(dict(item._mapping) == expected_mapping for item in rows) + cursor.execute("SELECT 42 AS replacement") + replacement = fetch_rows(cursor, method)[0] + assert replacement._column_names is cursor._cached_result_columns + assert replacement._column_names is not names + assert dict(replacement._mapping) == {"replacement": 42} + assert all(dict(item._mapping) == expected_mapping for item in rows) + + +@pytest.mark.parametrize("native", (False, True)) +def test_fast_row_without_column_snapshot_mapping(native): + values = [1, "abc"] + column_map = {"number": 0, "text": 1} + if native: + row = mssql_python.ddbc_bindings.construct_rows([values], Row, column_map, None)[0] + else: + row = Row._fast_create(values, column_map, None) + assert row._values is values + assert row._column_names is None + assert dict(row._mapping) == {"number": 1, "text": "abc"} + + +def assert_row_instance_capabilities(row): + row.metadata = "initial" + assert vars(row)["metadata"] == "initial" + vars(row)["metadata"] = "updated" + assert row.metadata == "updated" + del row.metadata + assert not hasattr(row, "metadata") + assert row.number == row["number"] == row[0] == 42 + assert dict(row._mapping) == {"number": 42} + reference = weakref.ref(row) + assert reference() is row + return reference + + +@pytest.mark.parametrize("construction", ("direct", "python_fast", "native")) +def test_constructed_row_preserves_instance_capabilities(construction): + values = [42] + column_map = {"number": 0} + if construction == "direct": + row = Row(values, column_map) + elif construction == "python_fast": + row = Row._fast_create(values, column_map, None) + else: + row = mssql_python.ddbc_bindings.construct_rows([values], Row, column_map, None)[0] + assert row._values is values + reference = assert_row_instance_capabilities(row) + del row + assert reference() is None + + +@pytest.mark.parametrize("method", FETCH_METHODS) +def test_fetched_row_preserves_instance_capabilities(connection, method): + with connection.cursor() as cursor: + cursor.execute("SELECT 42 AS number") + row = fetch_rows(cursor, method)[0] + reference = assert_row_instance_capabilities(row) + del row + assert reference() is None + + +@pytest.mark.parametrize("size", (0, 1, 3)) +def test_construct_rows_repeated_calls_release_references(size): + values = [[index] for index in range(size)] + column_map = {"Number": 0} + column_map_lower = {"number": 0} + column_names = ("Number",) + cursor = object() + tracked = (values, column_map, column_map_lower, column_names, cursor, *values) + references = [sys.getrefcount(value) for value in tracked] + for _ in range(10): + rows = mssql_python.ddbc_bindings.construct_rows( + values, Row, column_map, cursor, column_map_lower, column_names + ) + assert len(rows) == size + assert all(row._values is values[index] for index, row in enumerate(rows)) + assert all(row._column_map is column_map for row in rows) + assert all(row._column_map_lower is column_map_lower for row in rows) + assert all(row._column_names is column_names for row in rows) + assert all(row._cursor is cursor for row in rows) + del rows + assert [sys.getrefcount(value) for value in tracked] == references + + +def test_construct_rows_releases_partial_batch_on_attribute_error(): + class FailingRow(Row): + __slots__ = () + + @property + def _column_names(self): + return None + + @_column_names.setter + def _column_names(self, names): + if self._values[0] == 2: + raise RuntimeError("injected slot assignment failure") + + values = [[1], [2]] + column_map = {"number": 0} + column_names = ("number",) + cursor = object() + tracked = (values, column_map, column_names, cursor, *values) + references = [sys.getrefcount(value) for value in tracked] + for _ in range(10): + with pytest.raises(RuntimeError, match="injected slot assignment failure"): + mssql_python.ddbc_bindings.construct_rows( + values, FailingRow, column_map, cursor, None, column_names + ) + assert [sys.getrefcount(value) for value in tracked] == references + + +def test_construct_rows_accepts_row_subclasses(): + class CustomRow(Row): + __slots__ = () + + values = [42] + row = mssql_python.ddbc_bindings.construct_rows([values], CustomRow, {"number": 0}, None)[0] + assert type(row) is CustomRow + assert row._values is values + reference = assert_row_instance_capabilities(row) + del row + assert reference() is None + + +@pytest.mark.parametrize( + ("method", "bridge_name"), + ( + ("fetchone", "DDBCSQLFetchOne"), + ("fetchmany", "DDBCSQLFetchMany"), + ("fetchall", "DDBCSQLFetchAll"), + ), +) +def test_char_decoding_ctype_refresh(connection, method, bridge_name): + bridge = getattr(mssql_python.ddbc_bindings, bridge_name) + with ( + patch.object(connection, "getdecoding", wraps=connection.getdecoding) as reads, + patch.object(mssql_python.ddbc_bindings, bridge_name, wraps=bridge) as fetch, + connection.cursor() as cursor, + ): + for encoding, ctype in ( + ("utf-16le", mssql_python.SQL_WCHAR), + ("latin-1", mssql_python.SQL_CHAR), + ("utf-16le", mssql_python.SQL_WCHAR), + ): + cursor.execute("SELECT CONVERT(VARCHAR(1), 0xE9) AS txt") + connection.setdecoding(mssql_python.SQL_CHAR, encoding=encoding, ctype=ctype) + assert fetch_rows(cursor, method)[0].txt == "\u00e9" + assert fetch.call_args.args[-3:] == (encoding, "utf-16le", ctype) + assert reads.call_count == 8 + + +@pytest.mark.parametrize("invalid_type", ("None", "object()", "42", "'Row'")) +@pytest.mark.parametrize("rows", ("[]", "[[1]]")) +def test_construct_rows_rejects_non_types_in_subprocess(invalid_type, rows): + code = f""" +import sys +if sys.platform == "win32": + import ctypes + ctypes.windll.kernel32.SetErrorMode(0x0001 | 0x0002) +from mssql_python import ddbc_bindings +try: + ddbc_bindings.construct_rows({rows}, {invalid_type}, {{}}, None) +except TypeError as error: + assert str(error) == "row_class must be a type", str(error) +else: + raise AssertionError("Expected TypeError") +""" + result = subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True, timeout=30 + ) + assert result.returncode == 0, (result.returncode, result.stdout, result.stderr) + + +@pytest.mark.parametrize( + "unrelated_type", + ("types.FunctionType", "types.CodeType", "types.SimpleNamespace", "object", "int", "dict"), +) +@pytest.mark.parametrize("rows", ("[]", "[[1]]")) +def test_construct_rows_rejects_unrelated_types_in_subprocess(unrelated_type, rows): + code = f""" +import sys +import types +if sys.platform == "win32": + import ctypes + ctypes.windll.kernel32.SetErrorMode(0x0001 | 0x0002) +from mssql_python import ddbc_bindings +try: + ddbc_bindings.construct_rows({rows}, {unrelated_type}, {{}}, None) +except TypeError as error: + assert str(error) == "row_class must be Row or a Row subclass", str(error) +else: + raise AssertionError("Expected TypeError") +""" + result = subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True, timeout=30 + ) + assert result.returncode == 0, (result.returncode, result.stdout, result.stderr) + + +@pytest.mark.parametrize("method", FETCH_METHODS) +@pytest.mark.parametrize("stringify_uuid", (False, True)) +def test_unrelated_converter_preserves_zero_copy_fast_path( + connection, method, stringify_uuid, monkeypatch +): + monkeypatch.setattr(mssql_python, "native_uuid", not stringify_uuid) + converter = Mock(return_value="unexpected") + connection.add_output_converter(ConstantsDDBC.SQL_INTEGER.value, converter) + with connection.cursor() as cursor: + cursor.execute( + "SELECT CAST(N'abc' AS NVARCHAR(10)) AS txt, " + f"CAST('{UUID_TEXT}' AS UNIQUEIDENTIFIER) AS id" + ) + with ( + patch.object(Row, "_fast_create", wraps=Row._fast_create) as fast_one, + patch.object( + mssql_python.ddbc_bindings, + "construct_rows", + wraps=mssql_python.ddbc_bindings.construct_rows, + ) as fast_batch, + patch.object( + Row, "_apply_output_converters", side_effect=AssertionError("per-row lookup") + ), + patch.object( + Row, + "_apply_output_converters_optimized", + side_effect=AssertionError("unnecessary row copy"), + ), + patch.object( + connection, "get_output_converter", wraps=connection.get_output_converter + ) as lookups, + ): + row = fetch_rows(cursor, method)[0] + assert row.txt == "abc" + assert row.id == (UUID_TEXT if stringify_uuid else uuid.UUID(UUID_TEXT)) + assert lookups.call_count == 0 + assert fast_one.call_count == int(not stringify_uuid and method == "fetchone") + assert fast_batch.call_count == int(not stringify_uuid and method != "fetchone") + converter.assert_not_called() + + +@pytest.mark.parametrize( + ("method", "bridge_name"), + ( + ("fetchone", "DDBCSQLFetchOne"), + ("fetchmany", "DDBCSQLFetchMany"), + ("fetchall", "DDBCSQLFetchAll"), + ), +) +@pytest.mark.parametrize( + "status", + ( + ConstantsDDBC.SQL_SUCCESS.value, + ConstantsDDBC.SQL_SUCCESS_WITH_INFO.value, + ConstantsDDBC.SQL_NO_DATA.value, + ), +) +def test_fetch_drains_diagnostics_independent_of_final_status( + connection, method, bridge_name, status +): + bridge = getattr(mssql_python.ddbc_bindings, bridge_name) + warning = ("01000", 0, "injected fetch warning") + with connection.cursor() as cursor: + cursor.execute("SELECT CAST(N'abc' AS NVARCHAR(MAX)) AS txt") + + def fetch_with_final_status(*args): + bridge(*args) + return status + + with ( + patch.object( + mssql_python.ddbc_bindings, bridge_name, side_effect=fetch_with_final_status + ), + patch.object( + mssql_python.ddbc_bindings, "DDBCSQLGetAllDiagRecords", return_value=[warning] + ) as diagnostics, + ): + fetch_rows(cursor, method) + diagnostics.assert_called_once_with(cursor.hstmt) + assert warning in cursor.messages + + +@pytest.mark.parametrize( + ("method", "bridge_name"), + ( + ("fetchone", "DDBCSQLFetchOne"), + ("fetchmany", "DDBCSQLFetchMany"), + ("fetchall", "DDBCSQLFetchAll"), + ), +) +def test_fetch_error_is_raised_before_wrapping_rows(connection, method, bridge_name): + from types import SimpleNamespace + + with connection.cursor() as cursor: + cursor.execute("SELECT 1 AS number") + position = cursor._next_row_index + with ( + patch.object( + mssql_python.ddbc_bindings, + bridge_name, + return_value=ConstantsDDBC.SQL_ERROR.value, + ), + patch.object( + mssql_python.ddbc_bindings, + "DDBCSQLCheckError", + return_value=SimpleNamespace(sqlState="HY000", ddbcErrorMsg="injected fetch error"), + ), + ): + with pytest.raises(mssql_python.DatabaseError, match="injected fetch error"): + fetch_rows(cursor, method) + assert cursor._next_row_index == position + assert tuple(fetch_rows(cursor, method)[0]) == (1,)