Skip to content
Draft
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
8 changes: 6 additions & 2 deletions mssql_python/pybind/ddbc_bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6306,8 +6306,12 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row,
}
CaptureFetchDiagnostics(hStmt, ret, messages);
if (SQL_SUCCEEDED(ret)) {
// Retrieve column count
SQLSMALLINT colCount = SQLNumResultCols_wrap(StatementHandle, messages);
const auto snapshot = StatementHandle->resultMetadata.snapshot();
SQLSMALLINT colCount = snapshot.fullColumnCount;
if (colCount < 0) {
colCount = SQLNumResultCols_wrap(StatementHandle, messages);
StatementHandle->resultMetadata.publishFullColumnCount(snapshot.generation, colCount);
}
ret = SQLGetData_wrap(StatementHandle, colCount, row, charEncoding, wcharEncoding,
charCtype, messages);
if (!SQL_SUCCEEDED(ret)) {
Expand Down
13 changes: 12 additions & 1 deletion mssql_python/pybind/result_metadata.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,11 +30,12 @@ class ResultMetadataCache {
struct Snapshot {
uint64_t generation;
std::shared_ptr<const ResultMetadata> metadata;
SQLSMALLINT fullColumnCount;
};

Snapshot snapshot() const {
std::lock_guard<std::mutex> lock(mutex_);
return {generation_, metadata_};
return {generation_, metadata_, fullColumnCount_};
}

void publish(uint64_t generation, std::shared_ptr<const ResultMetadata> metadata) {
Expand All @@ -44,16 +45,26 @@ class ResultMetadataCache {
}
}

void publishFullColumnCount(uint64_t generation, SQLSMALLINT columnCount) {
std::lock_guard<std::mutex> lock(mutex_);
if (generation == generation_ && columnCount >= 0) {
fullColumnCount_ = columnCount;
}
}

void clear() {
std::lock_guard<std::mutex> lock(mutex_);
++generation_;
metadata_.reset();
fullColumnCount_ = -1;
}

private:
mutable std::mutex mutex_;
uint64_t generation_ = 0;
std::shared_ptr<const ResultMetadata> metadata_;
// Prefix GetData metadata does not establish the full result cardinality.
SQLSMALLINT fullColumnCount_ = -1;
};

class ResultMetadataFailureGuard {
Expand Down
250 changes: 250 additions & 0 deletions tests/test_fetch_settings_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -1111,3 +1111,253 @@ def test_fetch_error_is_raised_before_wrapping_rows(connection, method, bridge_n
fetch_rows(cursor, method)
assert cursor._next_row_index == position
assert tuple(fetch_rows(cursor, method)[0]) == (1,)


@pytest.mark.skipif(
sys.platform == "win32",
reason="Windows does not export the native driver function-pointer globals",
)
@pytest.mark.parametrize(
"mode",
(
"warm",
"prefix",
"shapes",
"nextset",
"fetch-error",
"count-error",
"decode-error",
"generation",
"mixed",
),
)
def test_native_fetchone_full_column_count_cache(conn_str, mode):
if not conn_str:
pytest.skip("DB_CONNECTION_STRING is required")
code = (
"import runpy, sys; "
"runpy.run_path(sys.argv[1])['_check_native_fetchone_full_column_count_cache']"
"(sys.argv[2], sys.argv[3])"
)
result = subprocess.run(
[
sys.executable,
"-c",
code,
str(Path(__file__).resolve()),
mode,
str(Path(mssql_python.ddbc_bindings.module.__file__).resolve()),
],
capture_output=True,
text=True,
timeout=60,
)
assert result.returncode == 0, (result.returncode, result.stdout, result.stderr)


def _check_native_fetchone_full_column_count_cache(mode, expected_native):
import ctypes
import os

native = Path(mssql_python.ddbc_bindings.module.__file__).resolve()
assert native == Path(expected_native)
library = ctypes.CDLL(str(native))
count_pointer = ctypes.c_void_p.in_dll(library, "SQLNumResultCols_ptr")
fetch_pointer = ctypes.c_void_p.in_dll(library, "SQLFetch_ptr")
count_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p, ctypes.POINTER(ctypes.c_short))
fetch_type = ctypes.CFUNCTYPE(ctypes.c_short, ctypes.c_void_p)
success = ConstantsDDBC.SQL_SUCCESS.value
error = ConstantsDDBC.SQL_ERROR.value
ddbc = mssql_python.ddbc_bindings
count_calls, callback_errors = [], []
fail_fetch = fail_count = invalidate_count = False

@count_type
def count_columns(handle, count):
nonlocal fail_count, invalidate_count
try:
count_calls.append(handle)
if fail_count:
fail_count = False
return error
result = original_count(handle, count)
if invalidate_count:
invalidate_count = False
status = ddbc.DDBCSQLSetStmtAttr(
cursor.hstmt, ConstantsDDBC.SQL_ATTR_QUERY_TIMEOUT.value, 0
)
if status != success:
callback_errors.append("Failed to invalidate the metadata generation")
return error
return result
except BaseException as failure:
callback_errors.append(type(failure).__name__)
return error

@fetch_type
def fetch_row(handle):
nonlocal fail_fetch
try:
if fail_fetch:
fail_fetch = False
return error
return original_fetch(handle)
except BaseException as failure:
callback_errors.append(type(failure).__name__)
return error

try:
connection = mssql_python.connect(os.environ["DB_CONNECTION_STRING"], timeout=5)
except mssql_python.Error as failure:
raise AssertionError(
f"Connection failed: {type(failure).__name__}; connection details withheld"
) from None
with connection, connection.cursor() as cursor:
query = (
"SELECT n AS number, CAST(N'text' AS NVARCHAR(10)) AS txt "
"FROM (VALUES (1), (2), (3), (4), (5), (6)) AS v(n) ORDER BY n"
)
cursor.execute(query)
saved_count, saved_fetch = count_pointer.value, fetch_pointer.value
assert saved_count and saved_fetch
original_count, original_fetch = count_type(saved_count), fetch_type(saved_fetch)
try:
count_pointer.value = ctypes.cast(count_columns, ctypes.c_void_p).value
fetch_pointer.value = ctypes.cast(fetch_row, ctypes.c_void_p).value
if mode == "prefix":
assert ddbc.DDBCSQLFetch(cursor.hstmt) == success
prefix = []
assert (
ddbc.DDBCSQLGetData(
cursor.hstmt,
1,
prefix,
"utf-16le",
"utf-16le",
ConstantsDDBC.SQL_C_WCHAR.value,
)
== success
)
assert prefix == [1]
assert count_calls == []
assert tuple(cursor.fetchone()) == (2, "text")
assert len(count_calls) == 1
assert tuple(cursor.fetchone()) == (3, "text")
assert len(count_calls) == 1
elif mode == "shapes":
for sql, expected in (
("SELECT 7 AS number", (7,)),
("SELECT CAST(N'new' AS NVARCHAR(10)) AS txt", ("new",)),
("SELECT CAST(NULL AS INT) AS empty_value, 9 AS number", (None, 9)),
):
cursor.execute(sql + " FROM (VALUES (1), (2)) AS v(n)")
count_calls.clear()
for _ in range(2):
row = tuple(cursor.fetchone())
assert row == expected
assert tuple(map(type, row)) == tuple(map(type, expected))
assert len(count_calls) == 1
assert cursor.fetchone() is None
assert len(count_calls) == 1
elif mode == "nextset":
cursor.execute(
"SELECT 1 AS number FROM (VALUES (1), (2)) AS v(n); "
"SELECT CAST(N'changed' AS NVARCHAR(10)) AS txt "
"FROM (VALUES (1), (2)) AS v(n)"
)
count_calls.clear()
assert [cursor.fetchone()[0] for _ in range(2)] == [1, 1]
assert len(count_calls) == 1
assert cursor.nextset()
count_calls.clear()
assert [cursor.fetchone()[0] for _ in range(2)] == ["changed", "changed"]
assert len(count_calls) == 1
elif mode == "fetch-error":
assert tuple(cursor.fetchone()) == (1, "text")
assert len(count_calls) == 1
fail_fetch = True
row = []
assert ddbc.DDBCSQLFetchOne(cursor.hstmt, row) == error
assert row == []
assert len(count_calls) == 1
assert tuple(cursor.fetchone()) == (2, "text")
assert len(count_calls) == 2
assert tuple(cursor.fetchone()) == (3, "text")
assert len(count_calls) == 2
elif mode == "count-error":
fail_count = True
with pytest.raises(mssql_python.DatabaseError):
ddbc.DDBCSQLFetchOne(cursor.hstmt, [])
assert len(count_calls) == 1
assert tuple(cursor.fetchone()) == (2, "text")
assert len(count_calls) == 2
assert tuple(cursor.fetchone()) == (3, "text")
assert len(count_calls) == 2
elif mode == "decode-error":
cursor.execute(
"SELECT CASE WHEN n = 2 THEN CAST(0x00D8 AS NVARCHAR(10)) "
"ELSE CAST(N'ok' AS NVARCHAR(10)) END AS txt "
"FROM (VALUES (1), (2), (3), (4)) AS v(n) ORDER BY n"
)
count_calls.clear()
assert cursor.fetchone()[0] == "ok"
with pytest.raises(UnicodeDecodeError):
cursor.fetchone()
assert len(count_calls) == 1
assert [cursor.fetchone()[0] for _ in range(2)] == ["ok", "ok"]
assert len(count_calls) == 2
elif mode == "generation":
invalidate_count = True
assert tuple(cursor.fetchone()) == (1, "text")
assert len(count_calls) == 1
assert tuple(cursor.fetchone()) == (2, "text")
assert len(count_calls) == 2
assert tuple(cursor.fetchone()) == (3, "text")
assert len(count_calls) == 2
elif mode == "mixed":
assert [tuple(row) for row in cursor.fetchmany(2)] == [(1, "text"), (2, "text")]
count_calls.clear()
assert tuple(cursor.fetchone()) == (3, "text")
assert len(count_calls) == 1
assert tuple(cursor.fetchmany(1)[0]) == (4, "text")
count_calls.clear()
assert tuple(cursor.fetchone()) == (5, "text")
assert count_calls == []
assert [tuple(row) for row in cursor.fetchall()] == [(6, "text")]
count_calls.clear()
assert cursor.fetchone() is None
assert count_calls == []
cursor.execute(query)
try:
other_connection = mssql_python.connect(
os.environ["DB_CONNECTION_STRING"], timeout=5
)
except mssql_python.Error as failure:
raise AssertionError(
f"Connection failed: {type(failure).__name__}; connection details withheld"
) from None
with other_connection, other_connection.cursor() as other:
other.execute("SELECT CAST(N'other' AS NVARCHAR(10))")
count_calls.clear()
assert cursor.fetchone()[0] == 1
assert other.fetchone()[0] == "other"
assert cursor.fetchone()[0] == 2
assert len(count_calls) == 2
else:
assert mode == "warm"
assert [tuple(cursor.fetchone()) for _ in range(6)] == [
(n, "text") for n in range(1, 7)
]
assert len(count_calls) == 1
assert cursor.fetchone() is None
assert len(count_calls) == 1
# Direct callers still perform a real query on every invocation.
assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2
assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2
assert len(count_calls) == 3
assert not callback_errors, callback_errors
assert not cursor.messages
finally:
count_pointer.value, fetch_pointer.value = saved_count, saved_fetch
cursor.execute("SELECT 42")
assert cursor.fetchone()[0] == 42
Loading