diff --git a/eng/pipelines/pr-validation-pipeline.yml b/eng/pipelines/pr-validation-pipeline.yml
index 2107a5bfe..d6cdcad81 100644
--- a/eng/pipelines/pr-validation-pipeline.yml
+++ b/eng/pipelines/pr-validation-pipeline.yml
@@ -1009,7 +1009,7 @@ jobs:
ACCEPT_EULA=Y apt-get install -y --no-install-recommends git unixodbc unixodbc-dev libodbc2 libodbcinst2 odbcinst msodbcsql18
git config --global --add safe.directory /workspace
odbcinst -q -d -n "ODBC Driver 18 for SQL Server"
- python -m eng.profiler_benchmarks.controller --reuse-candidate --leg "$(profilerLeg)" --output profiler-results
+ python -m eng.profiler_benchmarks.controller --reuse-candidate --ci-report --leg "$(profilerLeg)" --output profiler-results
'
else
echo "Skipping performance benchmarks on $(distroName) (only runs on Ubuntu with local SQL Server)"
diff --git a/eng/profiler_benchmarks/controller.py b/eng/profiler_benchmarks/controller.py
index 143164e21..db5577cc1 100644
--- a/eng/profiler_benchmarks/controller.py
+++ b/eng/profiler_benchmarks/controller.py
@@ -1,8 +1,10 @@
"""Build and measure base/candidate in isolated directories on the same CI agent."""
import argparse
+import ast
import contextlib
import faulthandler
+import hashlib
import importlib.util
import io
import json
@@ -11,13 +13,26 @@
import platform
import re
import signal
+import shutil
import subprocess
import sys
import tarfile
import tempfile
import time
-
-from .report import LEGS
+import textwrap
+
+from .report import (
+ LEGS,
+ MODES,
+ validate,
+ validate_ci_mode,
+ validate_ci_header,
+ ci_mode_reports,
+ validate_row_route,
+ expected_constructors,
+ validate_python_sources,
+ validate_samples,
+)
from . import workloads
ROOT = Path(__file__).resolve().parents[2]
@@ -28,18 +43,33 @@
LOCAL_BENCHMARK_TIMEOUT = 105 * 60
WORKER_TIMEOUT = 6 * 60
WINDOWS = os.name == "nt"
+FETCH_WORKER_TIMEOUT = 30
+CI_FINISH_RESERVE = 180
+
+
+class ProcessCleanupError(RuntimeError):
+ """Further work is unsafe until the previous process tree is reaped."""
-def git(*args):
- return subprocess.check_output(["git", "-C", str(ROOT), *args], text=True).strip()
+def git(*args, timeout=None):
+ return subprocess.check_output(
+ ["git", "-C", str(ROOT), *args], text=True, timeout=timeout
+ ).strip()
-def resolve_revisions(base, candidate):
- candidate = git("rev-parse", "--verify", "--end-of-options", f"{candidate}^{{commit}}")
+def resolve_revisions(base, candidate, timeout=None):
+ options = {"timeout": timeout} if timeout is not None else {}
+ candidate = git(
+ "rev-parse", "--verify", "--end-of-options", f"{candidate}^{{commit}}", **options
+ )
# ADO validates refs/pull/N/merge. Its first parent is the exact target snapshot,
# not whichever main build happened to finish most recently.
base = git(
- "rev-parse", "--verify", "--end-of-options", f"{base or candidate + '^1'}^{{commit}}"
+ "rev-parse",
+ "--verify",
+ "--end-of-options",
+ f"{base or candidate + '^1'}^{{commit}}",
+ **options,
)
return base, candidate
@@ -90,16 +120,55 @@ def terminate_process_tree(process):
process.wait(timeout=5)
-def build(path, log, timeout=900):
- env = dict(os.environ, ENABLE_PROFILING="1")
+def build(path, log, timeout=900, profiling=True):
+ # CMake >= 3.22 also initializes archived, older build scripts from this variable.
+ env = dict(os.environ, ENABLE_PROFILING="1" if profiling else "0", CMAKE_BUILD_TYPE="Release")
# build scripts find Python via PATH; keep the controller's interpreter.
env["PATH"] = str(Path(sys.executable).parent) + os.pathsep + env["PATH"]
command = ["cmd", "/c", "build.bat"] if os.name == "nt" else ["bash", "build.sh"]
+ run_process(command, log, timeout, cwd=path / "mssql_python/pybind", env=env)
+ configuration = release_build_configuration(path)
+ with log.open("a", encoding="utf-8") as output:
+ output.write("Verified Release configuration: " + json.dumps(configuration) + "\n")
+
+
+def release_build_configuration(source_root):
+ cache = source_root / "mssql_python/pybind/build/CMakeCache.txt"
+ if not cache.is_file():
+ tag = f"py{sys.version_info.major}{sys.version_info.minor}"
+ candidates = list(cache.parent.glob(f"*/{tag}/CMakeCache.txt"))
+ if len(candidates) > 1:
+ raise ValueError("Ambiguous native build configuration for this Python version")
+ if candidates:
+ cache = candidates[0]
+ entries = dict(
+ re.findall(r"^(CMAKE_[A-Z_]+):[^=\n]+=(.*)$", cache.read_text(encoding="utf-8"), re.M)
+ )
+ configurations = entries.get("CMAKE_CONFIGURATION_TYPES", "")
+ if (
+ "Release" not in configurations.split(";")
+ if configurations
+ else entries.get("CMAKE_BUILD_TYPE") != "Release"
+ ):
+ raise ValueError(
+ "Benchmarks require a Release native build (CMake >= 3.22 for old scripts)"
+ )
+ return {
+ key: entries[key]
+ for key in (
+ "CMAKE_GENERATOR",
+ "CMAKE_CXX_FLAGS",
+ "CMAKE_CXX_FLAGS_RELEASE",
+ "CMAKE_CONFIGURATION_TYPES" if configurations else "CMAKE_BUILD_TYPE",
+ )
+ }
+
+
+def run_process(command, log, timeout, **options):
with log.open("w", encoding="utf-8") as output:
process = subprocess.Popen(
command,
- cwd=path / "mssql_python/pybind",
- env=env,
+ **options,
stdout=output,
stderr=subprocess.STDOUT,
start_new_session=not WINDOWS,
@@ -108,13 +177,17 @@ def build(path, log, timeout=900):
try:
returncode = process.wait(timeout=timeout)
except subprocess.TimeoutExpired:
- terminate_process_tree(process)
+ try:
+ terminate_process_tree(process)
+ except (OSError, RuntimeError, subprocess.TimeoutExpired) as error:
+ raise ProcessCleanupError("Could not reap the timed-out process tree") from error
raise
if returncode:
raise subprocess.CalledProcessError(returncode, command)
def check_build(source_root, profiling):
+ print("Verified Release configuration: " + json.dumps(release_build_configuration(source_root)))
sys.path.insert(0, str(source_root))
import mssql_python
import mssql_python_odbc
@@ -145,10 +218,285 @@ def load_suite():
return core, workloads
+class _FetchContext:
+ """Controller-only recording policy; the interactive Profiler stays unchanged."""
+
+ def __init__(self, mode, native, python):
+ if mode not in ("latency", "route") or (native is not None) != (mode == "route"):
+ raise ValueError("Invalid fetch recording configuration")
+ self.native, self.python = native, python
+
+ def disable(self):
+ self.python.disable()
+ self.python.disable_timeline()
+ if self.native is not None:
+ self.native.disable()
+ self.native.disable_timeline()
+
+ def enable(self):
+ self.disable()
+ self.python.reset()
+ if self.native is not None:
+ self.native.reset()
+ self.native.enable()
+
+ def collect(self):
+ if self.python.is_enabled():
+ raise RuntimeError("Python phases became enabled during a fetch measurement")
+ self.disable()
+ py = self.python.get_stats()
+ if py:
+ raise RuntimeError("Python phase samples contaminate the fetch route")
+ return self.native.get_stats() if self.native is not None else {}, py
+
+
+def verify_reused_source(source_root, revision):
+ if source_root.resolve() == ROOT.resolve():
+ if git("rev-parse", "HEAD", timeout=5) != revision:
+ raise ValueError("Reused checkout revision changed")
+ git(
+ "diff",
+ "--exit-code",
+ "--quiet",
+ revision,
+ "--",
+ "mssql_python",
+ "mssql_python_odbc",
+ timeout=5,
+ )
+
+
+# Closed, reviewed dispatch bodies: main666, f539, many-only, and plain completion.
+# The options snapshot changes marshalling, not the plain Row-completion policy.
+# Method/comment/docstring edits require review and an explicit fingerprint update.
+# These reviewed method-source SHA256 digests are not credentials.
+_LEGACY_PYTHON_ROUTE = (
+ "52681c15f92e85c01e0b43a9bd763872051fdc956701046534a5beb2467e6c3c" # DevSkim: ignore DS173237
+)
+_LEGACY_FUSED_ROUTE = (
+ "d4c6671b89ace027c22585e761c4f4c3243e297be2162213d9bbb38d5bed03c6" # DevSkim: ignore DS173237
+)
+_MANY_ONLY_ROUTE = (
+ "544d5b0ccbf9df1f0b6ec72c79ecbfa49a3517cab0437a5520940b8941966759" # DevSkim: ignore DS173237
+)
+_PLAIN_COMPLETION_ROUTE = (
+ "776ebc4532b69fe2492317ef6af9d4b2bb08e8a0f112613a7ce2600500bc37d8" # DevSkim: ignore DS173237
+)
+_FETCH_OPTIONS_ROUTE = (
+ "6c8545f38956380e22b8cc9376e215e032d45ec68d1a29c669201f09477ec507" # DevSkim: ignore DS173237
+)
+
+
+def python_source_identity(source_root, revision):
+ """Read a selected root's descriptive policy, never infer it from measurements."""
+ if not SHA.fullmatch(revision):
+ raise ValueError("Invalid Python source revision")
+ source = (source_root / "mssql_python" / "cursor.py").read_text(encoding="utf-8")
+ try:
+ tree = ast.parse(source)
+ except SyntaxError as error:
+ raise ValueError("Invalid Python route source syntax") from error
+ name = "_DEFAULT_NATIVE_ROW_ROUTE"
+ mentions = [node for node in ast.walk(tree) if isinstance(node, ast.Name) and node.id == name]
+ declarations = [
+ node
+ for node in tree.body
+ if isinstance(node, ast.Assign)
+ and len(node.targets) == 1
+ and isinstance(node.targets[0], ast.Name)
+ and node.targets[0].id == name
+ ]
+ if mentions and (len(mentions) != 1 or len(declarations) != 1):
+ raise ValueError("Malformed or duplicate Python route declaration")
+ classes = [
+ node for node in tree.body if isinstance(node, ast.ClassDef) and node.name == "Cursor"
+ ]
+ if len(classes) != 1 or classes[0].decorator_list:
+ raise ValueError("Missing, duplicate or decorated Cursor class")
+ methods = []
+ for method in ("fetchone", "fetchmany", "fetchval"):
+ nodes = [
+ node
+ for node in classes[0].body
+ if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == method
+ ]
+ if len(nodes) != 1:
+ raise ValueError("Missing or duplicate fetch method")
+ if nodes[0].decorator_list or isinstance(nodes[0], ast.AsyncFunctionDef):
+ raise ValueError("Decorated or asynchronous fetch method is not a reviewed route")
+ methods.append(textwrap.dedent(ast.get_source_segment(source, nodes[0])))
+ # Check class-namespace bindings, not locals in methods or nested scopes.
+ route_names = {"fetchone", "fetchmany", "fetchval"}
+ pending = list(classes[0].body)
+ while pending:
+ node = pending.pop()
+ if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
+ if node.name in route_names and (
+ isinstance(node, ast.ClassDef) or node not in classes[0].body
+ ):
+ raise ValueError("Rebound Python fetch method")
+ pending.extend(node.decorator_list)
+ if isinstance(node, ast.ClassDef):
+ pending.extend(node.bases + node.keywords)
+ else:
+ pending.append(node.args)
+ if node.returns is not None:
+ pending.append(node.returns)
+ continue
+ if isinstance(node, ast.Lambda):
+ pending.append(node.args)
+ continue
+ if isinstance(node, (ast.ListComp, ast.SetComp, ast.DictComp, ast.GeneratorExp)):
+ continue
+ if (
+ isinstance(node, ast.Name)
+ and isinstance(node.ctx, (ast.Store, ast.Del))
+ and node.id in route_names
+ or isinstance(node, ast.alias)
+ and (node.asname or node.name.split(".")[0]) in route_names
+ or isinstance(node, (ast.ExceptHandler, ast.MatchAs, ast.MatchStar))
+ and node.name in route_names
+ or isinstance(node, ast.MatchMapping)
+ and node.rest in route_names
+ ):
+ raise ValueError("Rebound Python fetch method")
+ pending.extend(ast.iter_child_nodes(node))
+ fingerprint = hashlib.sha256(
+ json.dumps(methods, ensure_ascii=True, separators=(",", ":")).encode()
+ ).hexdigest()
+ if fingerprint in (_LEGACY_PYTHON_ROUTE, _PLAIN_COMPLETION_ROUTE, _FETCH_OPTIONS_ROUTE):
+ values = (False, False, False)
+ elif fingerprint == _LEGACY_FUSED_ROUTE:
+ values = (True, True, True)
+ elif fingerprint == _MANY_ONLY_ROUTE:
+ values = (False, True, False)
+ else:
+ raise ValueError("Unknown Python fetch dispatch fingerprint")
+ expected = dict(version=1, methods=dict(zip(("fetchone", "fetchmany", "fetchval"), values)))
+ if declarations:
+ expression = declarations[0].value
+ # literal_eval alone silently accepts duplicate dictionary keys.
+ for node in ast.walk(expression):
+ if isinstance(node, ast.Dict):
+ keys = [ast.literal_eval(key) for key in node.keys]
+ if any(type(key) is not str for key in keys) or len(set(keys)) != len(keys):
+ raise ValueError("Duplicate or invalid route declaration key")
+ route = ast.literal_eval(expression)
+ validate_row_route(route)
+ if route != expected:
+ raise ValueError("Python source-policy drift")
+ elif fingerprint in (_MANY_ONLY_ROUTE, _PLAIN_COMPLETION_ROUTE, _FETCH_OPTIONS_ROUTE):
+ raise ValueError("Missing Python route declaration")
+ return dict(
+ source_commit=revision,
+ python_cursor_sha256=hashlib.sha256(source.encode()).hexdigest(),
+ row_route=expected,
+ )
+
+
+def native_identity(source_root, revision, profiling):
+ from mssql_python import ddbc_bindings
+ import mssql_python.cursor as cursor_module
+
+ if (
+ Path(cursor_module.__file__).resolve()
+ != (source_root / "mssql_python" / "cursor.py").resolve()
+ ):
+ raise ValueError("Python cursor is outside the selected checkout")
+ python_identity = python_source_identity(source_root, revision)
+ guarded = hasattr(ddbc_bindings, "DDBCSQLFetchRow")
+ validate_row_route(python_identity["row_route"], guarded)
+ native_file = Path(ddbc_bindings.module.__file__).resolve()
+ if not native_file.is_relative_to(source_root.resolve()):
+ raise RuntimeError("Native binary is outside the selected checkout")
+ return dict(
+ **python_identity,
+ native_file=str(native_file),
+ native_sha256=hashlib.sha256(native_file.read_bytes()).hexdigest(),
+ native_profiling=profiling,
+ guarded_row=guarded,
+ )
+
+
+def fetch_worker(args, mode):
+ if mode not in ("latency", "route") or not SHA.fullmatch(args.revision or ""):
+ raise ValueError("Fetch measurement requires a mode and exact revision")
+ verify_reused_source(args.source_root, args.revision)
+ check_build(args.source_root, profiling=mode == "route")
+ import mssql_python
+ from mssql_python import ddbc_bindings, perf_timer
+
+ provenance = native_identity(args.source_root, args.revision, mode == "route")
+ cases = workloads.single_row_registry()
+ chosen = args.scenarios if args.scenarios is not None else list(cases)
+ if not chosen or set(chosen) - set(cases):
+ raise ValueError("Unknown or empty single-row workload selection")
+ ctx = _FetchContext(mode, ddbc_bindings.profiling if mode == "route" else None, perf_timer)
+ output = {}
+
+ def checkpoint(active=None, environment=None):
+ args.output.write_text(
+ json.dumps(
+ dict(
+ status="running" if environment is None else "complete",
+ active_scenario=active,
+ mode=mode,
+ provenance=provenance,
+ scenarios=output,
+ environment=environment,
+ ),
+ allow_nan=False,
+ ),
+ encoding="utf-8",
+ )
+
+ checkpoint()
+ try:
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as conn:
+ for name in chosen:
+ checkpoint(name)
+ result = cases[name](conn, ctx)
+ if mode == "route":
+ timer = (
+ "ddbc::FetchMany_wrap"
+ if name.endswith("fetchmany")
+ else "ddbc::FetchOne_wrap"
+ )
+ if result["cpp"].get(timer, {}).get("calls") != workloads.SINGLE_ROW_COUNT + 1:
+ raise RuntimeError("Native fetch call count does not match the workload")
+ constructors = (
+ result["cpp"].get("ddbc::FetchRow::construct_row", {}).get("calls", 0)
+ )
+ if constructors != expected_constructors(provenance, name.split("_")[1]):
+ raise RuntimeError("Native Row construction route was not established")
+ output[name] = {key: result[key] for key in ("wall_ms", "cpp", "py")}
+ output[name]["work"] = result["detail"]
+ checkpoint()
+ with conn.cursor() as cursor:
+ cursor.execute("SELECT CAST(SERVERPROPERTY('ProductVersion') AS VARCHAR(80))")
+ sql_version = cursor.fetchone()[0]
+ environment = dict(
+ os=platform.system(),
+ architecture=platform.machine().lower(),
+ python=platform.python_version(),
+ sql_version=sql_version,
+ )
+ finally:
+ ctx.disable()
+ checkpoint(environment=environment)
+
+
def worker(args):
+ mode = getattr(args, "mode", "diagnostic")
+ if mode != "diagnostic":
+ return fetch_worker(args, mode)
# Import the chosen driver FIRST, then the SAME workload/controller for both
# revisions. Never mix two native extensions into one interpreter.
+ revision = getattr(args, "revision", None)
+ if revision:
+ verify_reused_source(args.source_root, revision)
check_build(args.source_root, profiling=True)
+ provenance = native_identity(args.source_root, revision, True) if revision else None
core, workloads = load_suite()
cases = workloads.registry()
@@ -192,13 +540,21 @@ def worker(args):
python=platform.python_version(),
sql_version=sql_version,
)
- args.output.write_text(
- json.dumps(dict(environment=environment, scenarios=output), allow_nan=False),
- encoding="utf-8",
- )
-
-
-def measure(path, output, scenarios, timeout=WORKER_TIMEOUT):
+ result = dict(environment=environment, scenarios=output)
+ if provenance is not None:
+ result.update(status="complete", mode="diagnostic", provenance=provenance)
+ args.output.write_text(json.dumps(result, allow_nan=False), encoding="utf-8")
+
+
+def measure(
+ path,
+ output,
+ scenarios,
+ timeout=WORKER_TIMEOUT,
+ mode="diagnostic",
+ revision=None,
+ isolated=False,
+):
command = [
sys.executable,
"-u",
@@ -210,11 +566,18 @@ def measure(path, output, scenarios, timeout=WORKER_TIMEOUT):
"--output",
str(output),
]
+ if mode != "diagnostic" or revision is not None:
+ command += ["--mode", mode, "--revision", revision]
if scenarios:
command += ["--scenarios", *scenarios]
output.unlink(missing_ok=True)
- with output.with_suffix(".log").open("w", encoding="utf-8") as log:
- subprocess.run(command, stdout=log, stderr=subprocess.STDOUT, timeout=timeout, check=True)
+ if isolated:
+ run_process(command, output.with_suffix(".log"), timeout)
+ else:
+ with output.with_suffix(".log").open("w", encoding="utf-8") as log:
+ subprocess.run(
+ command, stdout=log, stderr=subprocess.STDOUT, timeout=timeout, check=True
+ )
return json.loads(output.read_text(encoding="utf-8"))
@@ -225,7 +588,224 @@ def remaining(deadline, limit):
return min(seconds, limit)
+def write_report(path, report):
+ temporary = path.with_suffix(".tmp")
+ temporary.write_text(json.dumps(report, allow_nan=False), encoding="utf-8")
+ temporary.replace(path)
+
+
+def run_ci_report(args):
+ if not args.reuse_candidate or args.mode != "diagnostic" or args.scenarios is not None:
+ raise ValueError(
+ "--ci-report requires --reuse-candidate and the complete default workload sets"
+ )
+ if args.samples != 5 or args.warmups != 1:
+ raise ValueError("--ci-report requires five pairs and one warmup pair")
+ finish_deadline = time.monotonic() + BENCHMARK_TIMEOUT
+ deadline = finish_deadline - CI_FINISH_RESERVE
+ base, candidate = resolve_revisions(args.base, args.candidate, timeout=remaining(deadline, 30))
+ if candidate != git("rev-parse", "HEAD", timeout=remaining(deadline, 30)):
+ raise ValueError("--ci-report must reuse checkout HEAD")
+ head = os.environ.get("SYSTEM_PULLREQUEST_SOURCECOMMITID", candidate)
+ if not SHA.fullmatch(head):
+ raise ValueError("Invalid PR head identity")
+ args.output.mkdir(parents=True, exist_ok=True)
+ report_path = args.output / "report.json"
+ common = dict(
+ status="incomplete",
+ leg=args.leg,
+ base_commit=base,
+ source_commit=candidate,
+ head_commit=head,
+ build_id=int(os.environ.get("BUILD_BUILDID", "0")),
+ samples=5,
+ warmups=1,
+ )
+ modes = {
+ mode: dict(
+ common,
+ schema_version=1 if mode == "diagnostic" else 2,
+ pairs=[],
+ unavailable_reason="Not started: waiting for earlier modes",
+ )
+ for mode in MODES
+ }
+ for mode in ("latency", "route"):
+ modes[mode]["mode"] = mode
+ bundle = modes["diagnostic"]
+ bundle.update(
+ measurement_bundle_version=1,
+ fetch_measurements={mode: modes[mode] for mode in ("latency", "route")},
+ )
+ validate_ci_header(bundle)
+ write_report(report_path, bundle)
+ on_identities = {}
+ failures = (
+ subprocess.CalledProcessError,
+ subprocess.TimeoutExpired,
+ TimeoutError,
+ ValueError,
+ KeyError,
+ TypeError,
+ OSError,
+ )
+ directory = tempfile.mkdtemp(prefix="profiler-ci-bundle-")
+ safe_to_clean = True
+ try:
+ roots = {name: Path(directory) / name for name in ("base-off", "candidate-off", "base-on")}
+ on_ready = False
+ for mode in ("latency", "route", "diagnostic"):
+ report = modes[mode]
+ stage = "admission"
+ try:
+ remaining(deadline, 1)
+ if mode == "diagnostic" and not on_ready:
+ raise ValueError("Shared ON build unavailable")
+ builds = (
+ (("base-off", base), ("candidate-off", candidate))
+ if mode == "latency"
+ else (("base-on", base),) if mode == "route" else ()
+ )
+ for name, revision in builds:
+ stage = "archive " + name
+ report["unavailable_reason"] = "Incomplete: " + stage
+ write_report(report_path, bundle)
+ run_process(
+ [
+ sys.executable,
+ "-m",
+ "eng.profiler_benchmarks.controller",
+ "--archive-source",
+ revision,
+ "--source-root",
+ str(roots[name]),
+ ],
+ args.output / ("archive-" + name + ".log"),
+ remaining(deadline, 60),
+ )
+ stage = "build " + name
+ report["unavailable_reason"] = "Incomplete: " + stage
+ write_report(report_path, bundle)
+ build(
+ roots[name],
+ args.output / ("build-" + name + ".log"),
+ remaining(deadline, 900),
+ profiling=mode != "latency",
+ )
+ if mode == "route":
+ on_ready = True
+ paths = {
+ "base": roots["base-off" if mode == "latency" else "base-on"],
+ "candidate": roots["candidate-off"] if mode == "latency" else ROOT,
+ }
+ sources = {
+ side: python_source_identity(paths[side], base if side == "base" else candidate)
+ for side in ("base", "candidate")
+ }
+ for previous in modes.values():
+ if "python_sources" in previous and previous["python_sources"] != sources:
+ raise ValueError("Python source identity changed across modes")
+ report["python_sources"] = sources
+ validate_python_sources(report)
+ identity = None
+ for sample in range(6):
+ pair = {}
+ for side in (
+ ("base", "candidate") if sample % 2 == 0 else ("candidate", "base")
+ ):
+ stage = f"{mode} pair {sample} {side}"
+ report["unavailable_reason"] = "Incomplete: " + stage
+ write_report(report_path, bundle)
+ pair[side] = measure(
+ paths[side],
+ args.output / f"{mode}-{side}-{sample}.json",
+ None,
+ remaining(
+ deadline,
+ WORKER_TIMEOUT if mode == "diagnostic" else FETCH_WORKER_TIMEOUT,
+ ),
+ mode=mode,
+ revision=base if side == "base" else candidate,
+ isolated=True,
+ )
+ probe = dict(report, pairs=[*report["pairs"], pair])
+ validate_ci_mode(probe, mode)
+ observed = {
+ side: (pair[side]["provenance"], pair[side]["environment"]) for side in pair
+ }
+ if identity is not None and observed != identity:
+ raise ValueError("Worker identity changed after warmup")
+ identity = observed
+ if mode != "latency":
+ for side in pair:
+ if side in on_identities and on_identities[side] != observed[side]:
+ raise ValueError("Shared ON worker identity mismatch")
+ on_identities[side] = observed[side]
+ if sample:
+ report["pairs"].append(pair)
+ write_report(report_path, bundle)
+ remaining(deadline, 1)
+ report["status"] = "complete"
+ validate_ci_mode(report, mode)
+ report.pop("unavailable_reason", None)
+ except ProcessCleanupError:
+ safe_to_clean = False
+ bundle["cleanup_required"] = directory
+ for pending in modes.values():
+ if pending["status"] != "complete":
+ pending["unavailable_reason"] = (
+ "Not completed: prior process cleanup failed"
+ )
+ report["status"] = "incomplete"
+ report["unavailable_reason"] = (
+ "Process cleanup failed; further modes were not started"
+ )
+ write_report(report_path, bundle)
+ raise
+ except failures as error:
+ report["status"] = "incomplete"
+ detail = (
+ str(error)[:140]
+ if isinstance(error, ValueError)
+ else "see raw worker/build evidence"
+ )
+ report["unavailable_reason"] = f"{stage}: {type(error).__name__}: {detail}"
+ print(report["unavailable_reason"], file=sys.stderr, flush=True)
+ finally:
+ write_report(report_path, bundle)
+ finally:
+ if safe_to_clean:
+ try:
+ shutil.rmtree(directory)
+ except OSError as error:
+ bundle["status"] = "incomplete"
+ bundle["cleanup_required"] = directory
+ bundle["unavailable_reason"] = (
+ "Build-directory cleanup failed; retained evidence requires cleanup"
+ )
+ write_report(report_path, bundle)
+ raise ProcessCleanupError("CI build-directory cleanup failed") from error
+ _, errors = ci_mode_reports(bundle)
+ for mode, reason in errors.items():
+ modes[mode]["status"] = "incomplete"
+ modes[mode]["unavailable_reason"] = reason
+ write_report(report_path, bundle)
+ if time.monotonic() > finish_deadline:
+ bundle["status"] = "incomplete"
+ bundle["unavailable_reason"] = "Aggregate finish deadline exceeded during finalization"
+ write_report(report_path, bundle)
+ if any(report["status"] != "complete" for report in modes.values()):
+ raise RuntimeError("CI performance report is incomplete; see per-mode reasons and raw logs")
+
+
def run(args):
+ mode = getattr(args, "mode", "diagnostic")
+ if mode not in MODES:
+ raise ValueError("Unknown measurement mode")
+ if mode == "latency" and args.reuse_candidate:
+ raise ValueError(
+ "Latency requires fresh native-OFF builds; cannot reuse the CI profiling build"
+ )
base, candidate = resolve_revisions(args.base, args.candidate)
args.output.mkdir(parents=True, exist_ok=True)
report_path = args.output / "report.json"
@@ -234,7 +814,7 @@ def run(args):
if not SHA.fullmatch(head):
head = candidate
report = dict(
- schema_version=1,
+ schema_version=1 if mode == "diagnostic" else 2,
status="incomplete",
leg=args.leg,
base_commit=base,
@@ -245,6 +825,8 @@ def run(args):
warmups=args.warmups,
pairs=[],
)
+ if mode != "diagnostic":
+ report["mode"] = mode
report_path.write_text(json.dumps(report), encoding="utf-8")
timeout = BENCHMARK_TIMEOUT if args.reuse_candidate else LOCAL_BENCHMARK_TIMEOUT
deadline = time.monotonic() + timeout
@@ -270,8 +852,22 @@ def run(args):
)
continue
checkout(revision, paths[side])
- print(f"Building profiling {side}: {revision}", flush=True)
- build(paths[side], args.output / f"build-{side}.log", remaining(deadline, 900))
+ print(f"Building {mode} {side}: {revision}", flush=True)
+ build_options = {"profiling": mode != "latency"} if mode != "diagnostic" else {}
+ build(
+ paths[side],
+ args.output / f"build-{side}.log",
+ remaining(deadline, 900),
+ **build_options,
+ )
+ if mode != "diagnostic":
+ report["python_sources"] = {
+ side: python_source_identity(paths[side], base if side == "base" else candidate)
+ for side in ("base", "candidate")
+ }
+ validate_python_sources(report)
+ report["unavailable_reason"] = "Collecting selected workload pairs"
+ identity = None
for sample in range(args.warmups + args.samples):
pair = {}
order = ("base", "candidate") if sample % 2 == 0 else ("candidate", "base")
@@ -282,14 +878,35 @@ def run(args):
args.output / f"{side}-{sample}.json",
args.scenarios,
remaining(deadline, WORKER_TIMEOUT),
+ **(
+ {"mode": mode, "revision": base if side == "base" else candidate}
+ if mode != "diagnostic"
+ else {}
+ ),
)
if pair["base"]["environment"] != pair["candidate"]["environment"]:
raise RuntimeError("Base and candidate environments differ")
+ if mode != "diagnostic":
+ probe = dict(
+ report, pairs=[*report["pairs"], pair], unavailable_reason="Collecting pairs"
+ )
+ if args.scenarios is None:
+ validate_ci_mode(probe, mode)
+ else:
+ validate_samples(probe, selected_cases=args.scenarios)
+ observed = {
+ side: (pair[side]["provenance"], pair[side]["environment"]) for side in pair
+ }
+ if identity is not None and identity != observed:
+ raise ValueError("Worker identity changed after warmup")
+ identity = observed
if sample >= args.warmups:
report["pairs"].append(pair)
report_path.write_text(json.dumps(report, allow_nan=False), encoding="utf-8")
if args.scenarios is None:
report["status"] = "complete"
+ if mode != "diagnostic":
+ validate(report)
report_path.write_text(json.dumps(report, allow_nan=False), encoding="utf-8")
print(f"Paired profiler report: {report_path}", flush=True)
@@ -303,6 +920,19 @@ def main():
parser.add_argument("--samples", type=int, default=5)
parser.add_argument("--warmups", type=int, default=1)
parser.add_argument("--scenarios", nargs="+", help="Local subset; CI runs the full registry")
+ parser.add_argument(
+ "--mode",
+ choices=MODES,
+ default="diagnostic",
+ help="diagnostic: both recorders; latency: native OFF/Python OFF; route: native ON/Python OFF",
+ )
+ parser.add_argument("--revision", help=argparse.SUPPRESS)
+ parser.add_argument(
+ "--ci-report",
+ action="store_true",
+ help="Bounded latency-first CI report with route and legacy diagnostics",
+ )
+ parser.add_argument("--archive-source", help=argparse.SUPPRESS)
parser.add_argument("--worker", action="store_true", help=argparse.SUPPRESS)
parser.add_argument("--source-root", type=Path, help=argparse.SUPPRESS)
parser.add_argument(
@@ -316,11 +946,19 @@ def main():
help="Verify native compile configuration and recording OFF, then exit",
)
args = parser.parse_args()
- if args.check_build:
+ if args.ci_report and (args.worker or args.archive_source or args.check_build):
+ parser.error("--ci-report cannot be combined with worker, archive or build-check modes")
+ if args.archive_source:
+ if not SHA.fullmatch(args.archive_source) or args.source_root is None:
+ parser.error("--archive-source requires an exact revision and --source-root")
+ checkout(args.archive_source, args.source_root)
+ elif args.check_build:
check_build(ROOT, profiling=args.check_build == "on")
elif args.worker:
if args.source_root is None or args.output is None:
parser.error("--worker requires --source-root and --output")
+ if args.mode != "diagnostic" and not SHA.fullmatch(args.revision or ""):
+ parser.error("Fetch measurement workers require an exact --revision")
# Dumps contain stack locations, not locals or connection strings. The
# parent still kills/reaps the worker at its deadline if it cannot finish.
faulthandler.enable()
@@ -337,7 +975,10 @@ def main():
or not 1 <= args.warmups <= 3
):
parser.error("Choose a leg, 3-15 measured pairs and 1-3 warmup pairs")
- run(args)
+ if args.ci_report:
+ run_ci_report(args)
+ else:
+ run(args)
if __name__ == "__main__":
diff --git a/eng/profiler_benchmarks/report.py b/eng/profiler_benchmarks/report.py
index 07f176f18..749756dd5 100644
--- a/eng/profiler_benchmarks/report.py
+++ b/eng/profiler_benchmarks/report.py
@@ -41,6 +41,34 @@
"scalar_fetchval": "10,000 scalar values / fetchval() (debug disabled)",
}
CASES = tuple(TASK_NAMES)
+MODES = ("diagnostic", "latency", "route")
+FETCH_CASES = tuple(
+ f"{shape}_{method}"
+ for shape in ("numeric", "mixed")
+ for method in ("fetchone", "fetchmany", "fetchval")
+)
+TASK_NAMES.update(
+ {
+ name: f"1,000 {name.split('_')[0]} rows / "
+ + ("fetchmany(1)" if name.endswith("fetchmany") else name.split("_")[1] + "()")
+ for name in FETCH_CASES
+ }
+)
+
+
+def measurement_mode(report):
+ version = report.get("schema_version")
+ if version == 1 and "mode" not in report:
+ return "diagnostic"
+ if version == 2 and report.get("mode") in ("latency", "route"):
+ return report["mode"]
+ raise ValueError("Unsupported measurement schema or mode")
+
+
+def cases_for(report):
+ return CASES if measurement_mode(report) == "diagnostic" else FETCH_CASES
+
+
MAX_BYTES = 8 * 1024 * 1024
MAX_COMMENT_CHARS = 60000
MAX_DIAGNOSTIC_ROWS = 20
@@ -114,8 +142,9 @@ def validate(report, build_id=None, head=None, source=None, base=None):
def _validate(report, build_id=None, head=None, source=None, base=None):
- if not isinstance(report, dict) or report.get("schema_version") != 1:
+ if not isinstance(report, dict):
raise ValueError("Unsupported report schema")
+ mode = measurement_mode(report)
if report.get("leg") not in LEGS or report.get("status") not in ("complete", "incomplete"):
raise ValueError("Invalid report status or leg")
if type(report.get("build_id")) is not int or report["build_id"] < 0:
@@ -139,12 +168,25 @@ def _validate(report, build_id=None, head=None, source=None, base=None):
pairs = report.get("pairs")
if not isinstance(pairs, list) or len(pairs) > samples:
raise ValueError("Invalid sample pairs")
+ validate_python_sources(report)
if report["status"] == "incomplete":
- return report
+ if "python_sources" in report:
+ text(report.get("unavailable_reason"), limit=240)
+ return validate_samples(report)
if len(pairs) != samples:
raise ValueError("Incomplete sample pairs")
+ return validate_samples(report)
+
+
+def validate_samples(report, selected_cases=None):
+ mode = measurement_mode(report)
+ pairs = report["pairs"]
+ cases = cases_for(report) if selected_cases is None else selected_cases
+ if not cases or set(cases) - set(cases_for(report)):
+ raise ValueError("Invalid selected scenarios")
environment = None
work = {}
+ provenance = {}
for pair in pairs:
if not isinstance(pair, dict) or set(pair) != {"base", "candidate"}:
raise ValueError("Invalid paired sample")
@@ -152,6 +194,13 @@ def _validate(report, build_id=None, head=None, source=None, base=None):
sample = pair[side]
if not isinstance(sample, dict):
raise ValueError("Invalid sample")
+ if mode != "diagnostic" or "python_sources" in report or "provenance" in sample:
+ if sample.get("mode") != mode or sample.get("status") != "complete":
+ raise ValueError("Incomplete or mismatched measurement mode")
+ native_identity = validate_native_identity(sample, report, side, mode)
+ if side in provenance and provenance[side] != native_identity:
+ raise ValueError("Native identity changed between samples")
+ provenance[side] = native_identity
env = sample["environment"]
if not isinstance(env, dict) or set(env) != {
"os",
@@ -173,7 +222,7 @@ def _validate(report, build_id=None, head=None, source=None, base=None):
scenarios = sample["scenarios"]
if not isinstance(scenarios, dict):
raise ValueError("Invalid scenarios object")
- if set(scenarios) != set(CASES):
+ if set(scenarios) != set(cases):
raise ValueError("Scenario set incomplete or changed")
for name, scenario in scenarios.items():
if not isinstance(scenario, dict):
@@ -185,12 +234,34 @@ def _validate(report, build_id=None, head=None, source=None, base=None):
if name in work and work[name] != identity:
raise ValueError(f"Workload changed for {name}")
work[name] = identity
+ if mode != "diagnostic":
+ if scenario["py"] or (mode == "latency" and scenario["cpp"]):
+ raise ValueError("Recording contaminated the measurement mode")
+ shape, method = name.split("_")
+ if scenario["work"] != f"Rows: 1000; shape: {shape}; API: {method}; EOF: 1":
+ raise ValueError("Single-row workload identity mismatch")
+ if mode == "route":
+ timer = (
+ "ddbc::FetchMany_wrap"
+ if method == "fetchmany"
+ else "ddbc::FetchOne_wrap"
+ )
+ stats = scenario["cpp"]
+ if not isinstance(stats, dict) or not isinstance(stats.get(timer), dict):
+ raise ValueError("Missing native fetch route evidence")
+ if stats[timer].get("calls") != 1001:
+ raise ValueError("Native fetch route count mismatch")
+ constructor = stats.get("ddbc::FetchRow::construct_row", {})
+ if not isinstance(constructor, dict) or constructor.get("calls", 0) != (
+ expected_constructors(native_identity, method)
+ ):
+ raise ValueError("Native constructor route count mismatch")
for layer in ("cpp", "py"):
stats = scenario[layer]
if (
not isinstance(stats, dict)
or len(stats) > 300
- or (layer == "cpp" and not stats)
+ or (layer == "cpp" and mode != "latency" and not stats)
):
raise ValueError("Missing or oversized profiling data")
for label, counter in stats.items():
@@ -209,6 +280,237 @@ def _validate(report, build_id=None, head=None, source=None, base=None):
return report
+def validate_row_route(route, guarded=None):
+ if (
+ type(route) is not dict
+ or set(route) != {"version", "methods"}
+ or type(route["version"]) is not int
+ or route["version"] != 1
+ ):
+ raise ValueError("Unsupported Python row route version or shape")
+ methods = route["methods"]
+ if (
+ type(methods) is not dict
+ or set(methods) != {"fetchone", "fetchmany", "fetchval"}
+ or any(type(value) is not bool for value in methods.values())
+ ):
+ raise ValueError("Invalid Python row route methods")
+ values = tuple(methods[name] for name in ("fetchone", "fetchmany", "fetchval"))
+ if values not in ((False, False, False), (True, True, True), (False, True, False)):
+ raise ValueError("Unsupported Python row route policy")
+ if guarded is not None and (type(guarded) is not bool or (any(values) and not guarded)):
+ raise ValueError("Native binding contradicts Python row route")
+ return methods
+
+
+def expected_constructors(identity, method):
+ if "row_route" not in identity:
+ return 1000 if identity["guarded_row"] else 0
+ return 1000 if validate_row_route(identity["row_route"], identity["guarded_row"])[method] else 0
+
+
+def route_description(report, method):
+ descriptions = []
+ for side in ("base", "candidate"):
+ identity = report["pairs"][0][side]["provenance"]
+ descriptions.append(
+ f"{side}: binding={identity['guarded_row']}, default {method} native Row="
+ f"{bool(expected_constructors(identity, method))}"
+ )
+ return "; ".join(descriptions) + "."
+
+
+def validate_python_sources(report):
+ if "python_sources" not in report:
+ return
+ sources = report["python_sources"]
+ if type(sources) is not dict or set(sources) != {"base", "candidate"}:
+ raise ValueError("Missing Python source anchors")
+ for side, source in sources.items():
+ if type(source) is not dict or set(source) != {
+ "source_commit",
+ "python_cursor_sha256",
+ "row_route",
+ }:
+ raise ValueError("Invalid Python source anchor")
+ if source["source_commit"] != report["base_commit" if side == "base" else "source_commit"]:
+ raise ValueError("Python source revision mismatch")
+ if not isinstance(source["python_cursor_sha256"], str) or not re.fullmatch(
+ r"[0-9a-f]{64}", source["python_cursor_sha256"]
+ ):
+ raise ValueError("Invalid Python cursor digest")
+ validate_row_route(source["row_route"])
+
+
+def validate_native_identity(sample, report, side, mode):
+ identity = sample.get("provenance")
+ legacy = {"source_commit", "native_file", "native_sha256", "native_profiling", "guarded_row"}
+ modern = legacy | {"python_cursor_sha256", "row_route"}
+ if not isinstance(identity, dict) or set(identity) not in (legacy, modern):
+ raise ValueError("Missing native measurement identity")
+ if (set(identity) == modern) != ("python_sources" in report):
+ raise ValueError("Missing Python source anchors or worker route fields")
+ if set(identity) == modern:
+ validate_python_sources(report)
+ validate_row_route(identity["row_route"], identity["guarded_row"])
+ anchor = {
+ key: identity[key] for key in ("source_commit", "python_cursor_sha256", "row_route")
+ }
+ if anchor != report["python_sources"][side]:
+ raise ValueError("Python source-policy drift")
+ if identity["source_commit"] != report["base_commit" if side == "base" else "source_commit"]:
+ raise ValueError("Worker revision mismatch")
+ text(identity["native_file"], limit=4096)
+ if not isinstance(identity["native_sha256"], str) or not re.fullmatch(
+ r"[0-9a-f]{64}", identity["native_sha256"]
+ ):
+ raise ValueError("Invalid native binary digest")
+ if (
+ identity["native_profiling"] is not (mode != "latency")
+ or type(identity["guarded_row"]) is not bool
+ ):
+ raise ValueError("Invalid native measurement configuration")
+ return identity
+
+
+def validate_ci_mode(report, mode):
+ if measurement_mode(report) != mode:
+ raise ValueError("CI measurement mode mismatch")
+ validate(report)
+ if report["status"] == "incomplete":
+ text(report["unavailable_reason"], limit=240)
+ validate_samples(report)
+ identities = {}
+ for pair in report["pairs"]:
+ for side, sample in pair.items():
+ if sample.get("status") != "complete" or sample.get("mode") != mode:
+ raise ValueError("Incomplete or mismatched worker mode")
+ identity = validate_native_identity(sample, report, side, mode)
+ if side in identities and identities[side] != identity:
+ raise ValueError("Native identity changed between samples")
+ identities[side] = identity
+ return report
+
+
+def has_ci_bundle(report):
+ return isinstance(report, dict) and (
+ "measurement_bundle_version" in report or "fetch_measurements" in report
+ )
+
+
+def validate_ci_header(report, build_id=None, head=None, source=None, base=None):
+ if (
+ type(report.get("measurement_bundle_version")) is not int
+ or report["measurement_bundle_version"] != 1
+ ):
+ raise ValueError("Unsupported CI measurement bundle version")
+ if measurement_mode(report) != "diagnostic":
+ raise ValueError("CI bundle root must be the diagnostic report")
+ header = dict(report, status="incomplete", pairs=[])
+ header.pop("python_sources", None)
+ validate(header, build_id, head, source, base)
+ if report["samples"] != 5 or report["warmups"] != 1:
+ raise ValueError("CI bundle requires five pairs and one warmup")
+ children = report.get("fetch_measurements")
+ if not isinstance(children, dict) or set(children) - {"latency", "route"}:
+ raise ValueError("Invalid CI measurement children")
+ return report
+
+
+def ci_mode_reports(report):
+ validate_ci_header(report)
+ valid, errors = {}, {}
+ for mode in ("latency", "route", "diagnostic"):
+ item = report if mode == "diagnostic" else report["fetch_measurements"].get(mode)
+ try:
+ if not isinstance(item, dict):
+ raise ValueError("Missing mode")
+ for key in (
+ "build_id",
+ "head_commit",
+ "source_commit",
+ "base_commit",
+ "leg",
+ "samples",
+ "warmups",
+ ):
+ if item.get(key) != report[key]:
+ raise ValueError("Mode provenance mismatch: " + key)
+ validate_ci_mode(item, mode)
+ valid[mode] = item
+ if item["status"] != "complete":
+ errors[mode] = item["unavailable_reason"]
+ except (KeyError, TypeError, ValueError) as error:
+ errors[mode] = "Invalid mode data: " + str(error)[:180]
+ on = [valid.get(mode) for mode in ("route", "diagnostic")]
+ if all(item is not None and item["pairs"] for item in on):
+ if any(
+ on[0]["pairs"][0][side][key] != on[1]["pairs"][0][side][key]
+ for side in ("base", "candidate")
+ for key in ("provenance", "environment")
+ ):
+ for mode in ("route", "diagnostic"):
+ valid.pop(mode)
+ errors[mode] = "Shared ON binary/environment identity mismatch"
+ anchored = [(mode, item) for mode, item in valid.items() if "python_sources" in item]
+ if anchored:
+ reference = next((item for mode, item in anchored if mode == "latency"), anchored[0][1])
+ for mode, item in list(valid.items()):
+ if item.get("python_sources") != reference["python_sources"] and (
+ item["pairs"] or "python_sources" in item
+ ):
+ valid.pop(mode)
+ errors[mode] = "Python source identity differs across modes"
+ latency = valid.get("latency")
+ if latency is not None and latency["pairs"]:
+ for mode in ("route", "diagnostic"):
+ item = valid.get(mode)
+ if (
+ item is not None
+ and item["pairs"]
+ and item["pairs"][0]["base"]["environment"]
+ != latency["pairs"][0]["base"]["environment"]
+ ):
+ valid.pop(mode)
+ errors[mode] = "Environment differs from the latency comparison"
+ return valid, errors
+
+
+def render_ci_reports(reports, head, build_id, issues=()):
+ diagnostics = []
+ diagnostic_issues = list(issues)
+ seen = set()
+ for report in reports:
+ leg = report["leg"]
+ if leg in seen:
+ raise ValueError("Duplicate performance report leg")
+ seen.add(leg)
+ if has_ci_bundle(report):
+ valid, errors = ci_mode_reports(report)
+ else:
+ validate(report)
+ valid = {measurement_mode(report): report}
+ errors = {}
+ diagnostic = valid.get("diagnostic")
+ if diagnostic is not None and diagnostic["status"] == "complete":
+ diagnostics.append(diagnostic)
+ else:
+ reason = errors.get("diagnostic", "missing or incomplete diagnostic measurement")
+ diagnostic_issues.append(leg + " (" + reason + ")")
+ body = render(diagnostics, head, build_id, diagnostic_issues)
+ artifact_note = "Raw samples and logs are attached to the ADO run as `profiler-*` artifacts."
+ body = body.replace(
+ artifact_note,
+ "This headline uses the original 22-task profiling-enabled diagnostics; separate "
+ "OFF/OFF latency and ON/OFF route measurements, when available, are retained in the raw artifacts "
+ "and are not headline inputs.\n\n" + artifact_note,
+ 1,
+ )
+ if len(body) > MAX_COMMENT_CHARS:
+ raise ValueError("CI performance comment exceeds its bounded size")
+ return body
+
+
def assess(evidence, artifact_urls, load_artifact, issues=()):
issues = list(issues)
try:
@@ -238,15 +540,13 @@ def assess(evidence, artifact_urls, load_artifact, issues=()):
# Match coverage's trust boundary: select the exact PR-head build and treat
# its bounded artifacts as data without requiring an identical producer tree.
reports = []
+ ci_requested = False
for leg, url in artifact_urls.items():
try:
- report = validate(
- artifact_report(load_artifact(url)),
- build_id,
- evidence.head,
- source,
- evidence.base,
- )
+ report = artifact_report(load_artifact(url))
+ ci_requested = ci_requested or has_ci_bundle(report)
+ validator = validate_ci_header if has_ci_bundle(report) else validate
+ validator(report, build_id, evidence.head, source, evidence.base)
if report["leg"] != leg:
raise ValueError("Artifact leg mismatch")
reports.append(report)
@@ -261,7 +561,8 @@ def assess(evidence, artifact_urls, load_artifact, issues=()):
issues.append(leg + " (invalid artifact)")
try:
- return render(reports, evidence.head, build_id, issues)
+ renderer = render_ci_reports if ci_requested else render
+ return renderer(reports, evidence.head, build_id, issues)
except ValueError:
return unavailable("Performance report rendering failed.")
@@ -269,7 +570,7 @@ def assess(evidence, artifact_urls, load_artifact, issues=()):
def comparisons(report):
"""Do not add inclusive phase totals together or treat them as wall-clock time."""
output = []
- for name in CASES:
+ for name in cases_for(report):
base = [pair["base"]["scenarios"][name] for pair in report["pairs"]]
candidate = [pair["candidate"]["scenarios"][name] for pair in report["pairs"]]
ratios = [new["wall_ms"] / old["wall_ms"] for old, new in zip(base, candidate)]
@@ -319,6 +620,9 @@ def comparisons(report):
base_ms=old,
candidate_ms=new,
change_pct=(ratio - 1) * 100,
+ ratio=ratio,
+ ratio_min=min(ratios),
+ ratio_max=max(ratios),
status=status,
phases=phases,
counts=sorted(changed_counts)[:3],
@@ -349,8 +653,12 @@ def issue_reason(leg, issues):
return global_issues[0] if global_issues else "incomplete benchmark"
-def render(reports, head, build_id, issues=()):
+def render(reports, head, build_id, issues=(), default_mode="diagnostic"):
url = f"https://dev.azure.com/sqlclientdrivers/public/_build/results?buildId={build_id}"
+ modes = {measurement_mode(r) for r in reports}
+ if len(modes) > 1:
+ raise ValueError("Cannot combine different measurement modes")
+ mode = next(iter(modes), default_mode)
by_leg = {r["leg"]: r for r in reports}
if len(by_leg) != len(reports):
raise ValueError("Duplicate performance report leg")
@@ -425,6 +733,12 @@ def render(reports, head, build_id, issues=()):
)
verdict = "✅ No regression detected"
+ if mode == "route" and completed:
+ verdict = "Native route attribution (instrumented)"
+ opening = (
+ "Instrumented paired timings below describe the verified native route, not production latency. "
+ + opening
+ )
improvement_tasks = len({row["name"] for _, row in improvements})
regression_tasks = len({row["name"] for _, row in regressions})
lines = [
@@ -442,13 +756,25 @@ def render(reports, head, build_id, issues=()):
f"{len(completed)}/{len(LEGS)} ENVIRONMENTS",
"",
]
+ if mode != "diagnostic":
+ lines += [
+ "**Measurement:** "
+ + (
+ "native instrumentation OFF / Python phases OFF; controlled fetch-loop latency."
+ if mode == "latency"
+ else "native instrumentation ON / Python phases OFF; route attribution, not production latency."
+ ),
+ "",
+ ]
if noisy:
noisy_tasks = len({row["name"] for _, row in noisy})
lines += [
f"{noisy_tasks} INCONSISTENT SLOWDOWN" f"{'S' if noisy_tasks != 1 else ''}",
"",
]
- affected_tasks = [name for name in CASES if any(row["name"] == name for _, row in highlighted)]
+ affected_tasks = [
+ name for name in TASK_NAMES if any(row["name"] == name for _, row in highlighted)
+ ]
if highlighted and len(affected_tasks) <= MAX_FINGERPRINT_TASKS:
affected_legs = [leg for leg in LEGS if any(item_leg == leg for item_leg, _ in highlighted)]
by_signal = {(leg, row["name"]): row for leg, row in highlighted}
@@ -476,7 +802,9 @@ def render(reports, head, build_id, issues=()):
lines.append("")
if regressions:
lines.append(
- "The largest recorded phase increases for these tasks are shown below. "
+ "Paired fetch-loop timings are shown below. Phase attribution is unavailable in latency mode."
+ if mode == "latency"
+ else "The largest recorded phase increases for these tasks are shown below. "
"Phase timings are supporting evidence, not root-cause proof."
)
if noisy:
@@ -537,7 +865,11 @@ def render(reports, head, build_id, issues=()):
diagnostics += 1
phases = "; ".join(f"{escape(label)} {delta:+.3f} ms" for delta, label in row["phases"])
counts = "; ".join(escape(label) for label in row["counts"])
- detail = phases or "no measured phase delta"
+ detail = phases or (
+ "phase attribution unavailable (recording OFF)"
+ if mode == "latency"
+ else "no measured phase delta"
+ )
if counts:
detail += f". Call changes: {counts}"
lines.append(f"**{TASK_NAMES[row['name']]}:** {detail}.")
@@ -564,8 +896,12 @@ def render(reports, head, build_id, issues=()):
lines += [
"",
f"### {environment_name(leg)}",
- "| Database task | Before | After | Paired change | Result |",
- "|---|---:|---:|---:|---|",
+ (
+ "| Database task | Before | After | Paired change | Result |"
+ if mode == "diagnostic"
+ else "| Database task | Before | After | Paired change | Ratio median [min, max] | Result |"
+ ),
+ "|---|---:|---:|---:|---|" if mode == "diagnostic" else "|---|---:|---:|---:|---:|---|",
]
for row in rows:
result = {
@@ -576,7 +912,13 @@ def render(reports, head, build_id, issues=()):
}[row["status"]]
lines.append(
f"| {TASK_NAMES[row['name']]} | {row['base_ms']:.3f} ms | "
- f"{row['candidate_ms']:.3f} ms | {row['change_pct']:+.1f}% | {result} |"
+ f"{row['candidate_ms']:.3f} ms | {row['change_pct']:+.1f}% | "
+ + (
+ f"{row['ratio']:.3f} [{row['ratio_min']:.3f}, {row['ratio_max']:.3f}] | "
+ if mode != "diagnostic"
+ else ""
+ )
+ + f"{result} |"
)
lines += [
"",
@@ -587,7 +929,11 @@ def render(reports, head, build_id, issues=()):
"",
]
lines += [
- f"[ADO build {build_id}]({url})",
+ (
+ f"[ADO build {build_id}]({url})"
+ if build_id or mode == "diagnostic"
+ else "Local paired comparison (no ADO build)"
+ ),
"",
f"PR head: `{head}`",
]
@@ -595,7 +941,7 @@ def render(reports, head, build_id, issues=()):
first = next(iter(completed.values()))[0]
lines += [
f"Base: `{first['base_commit']}`",
- f"Measured merge: `{first['source_commit']}`",
+ f"Measured {'merge' if mode == 'diagnostic' else 'source'}: `{first['source_commit']}`",
"",
]
for leg, (report, _) in completed.items():
@@ -605,6 +951,22 @@ def render(reports, head, build_id, issues=()):
f"{escape(env['architecture'])}, SQL {escape(env['sql_version'])}; "
f"{report['samples']} paired comparisons and {report['warmups']} warmup."
)
+ if mode != "diagnostic":
+ for side in ("base", "candidate"):
+ identity = report["pairs"][0][side]["provenance"]
+ lines.append(
+ f"- {environment_name(leg)} {side} native SHA256: `{identity['native_sha256']}`; "
+ f"guarded Row entry available: {identity['guarded_row']}. "
+ + (
+ "Default native Row routes: "
+ + ", ".join(
+ f"{method}={enabled}"
+ for method, enabled in identity["row_route"]["methods"].items()
+ )
+ if "row_route" in identity
+ else "Historical all-or-none route contract."
+ )
+ )
lines += [
"",
"A consistent change requires more than 20% median paired movement, at least "
@@ -619,9 +981,15 @@ def render(reports, head, build_id, issues=()):
lines += ["", "Unavailable or rejected data: " + ", ".join(escape(x) for x in issues)]
lines += [
"",
- "Both revisions use profiling-enabled builds on the same agent and database, "
- "with alternating order and discarded warmups. Results are diagnostic and do "
- "not represent production-wheel latency.",
+ (
+ "Both revisions use native-instrumentation-OFF builds with Python phases OFF on "
+ "the same agent and database, alternating order and discarded warmups. This is "
+ "controlled fetch-loop latency, not a customer-production or pyodbc comparison."
+ if mode == "latency"
+ else "Both revisions use profiling-enabled builds on the same agent and database, "
+ "with alternating order and discarded warmups. Results are diagnostic and do "
+ "not represent production-wheel latency."
+ ),
"",
"Raw samples and logs are attached to the ADO run as `profiler-*` artifacts.",
"",
@@ -648,17 +1016,20 @@ def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("reports", nargs="+", type=Path)
args = parser.parse_args()
- reports = [validate(json.loads(path.read_text(encoding="utf-8"))) for path in args.reports]
+ reports = [json.loads(path.read_text(encoding="utf-8")) for path in args.reports]
+ for report in reports:
+ (validate_ci_header if has_ci_bundle(report) else validate)(report)
first = reports[0]
for report in reports[1:]:
- validate(
+ (validate_ci_header if has_ci_bundle(report) else validate)(
report,
build_id=first["build_id"],
head=first["head_commit"],
source=first["source_commit"],
base=first["base_commit"],
)
- print(render(reports, first["head_commit"], first["build_id"]))
+ renderer = render_ci_reports if any(has_ci_bundle(report) for report in reports) else render
+ print(renderer(reports, first["head_commit"], first["build_id"]))
if __name__ == "__main__":
diff --git a/eng/profiler_benchmarks/workloads.py b/eng/profiler_benchmarks/workloads.py
index 9059ab6c9..bb1096bdd 100644
--- a/eng/profiler_benchmarks/workloads.py
+++ b/eng/profiler_benchmarks/workloads.py
@@ -209,3 +209,69 @@ def registry():
result["lob_varchar_256k_fetchall"] = (lob_fetch, False)
result["scalar_fetchval"] = (scalar_fetchval, True)
return result
+
+
+SINGLE_ROW_COUNT = 1000
+
+
+def single_row_fetch(conn, ctx, method, shape):
+ """Time one API through EOF; validate every column outside the timed window."""
+ from mssql_python.logging import logger
+
+ if method not in ("fetchone", "fetchmany", "fetchval") or shape not in ("numeric", "mixed"):
+ raise ValueError("Unknown single-row workload")
+ if logger.is_debug_enabled:
+ raise ValueError("Single-row measurements require debug logging OFF")
+ second = "n + 10" if shape == "numeric" else "CAST(N'text' AS NVARCHAR(10))"
+ query = (
+ "WITH digits(n) AS (SELECT n FROM (VALUES (0),(1),(2),(3),(4),(5),(6),(7),(8),(9)) d(n)), "
+ "numbers(n) AS (SELECT a.n + 10*b.n + 100*c.n FROM digits a CROSS JOIN digits b "
+ "CROSS JOIN digits c) "
+ f"SELECT n AS number, {second} AS value FROM numbers ORDER BY n"
+ )
+ with conn.cursor() as cursor:
+ cursor.execute(query)
+ try:
+ ctx.enable()
+ start = time.perf_counter()
+ if method == "fetchmany":
+ values = [cursor.fetchmany(1) for _ in range(SINGLE_ROW_COUNT)]
+ eof = cursor.fetchmany(1)
+ elif method == "fetchval":
+ values = [cursor.fetchval() for _ in range(SINGLE_ROW_COUNT)]
+ eof = cursor.fetchval()
+ else:
+ values = [cursor.fetchone() for _ in range(SINGLE_ROW_COUNT)]
+ eof = cursor.fetchone()
+ wall_ms = (time.perf_counter() - start) * 1000
+ cpp, py = ctx.collect()
+ if method == "fetchmany":
+ assert eof == [] and all(len(batch) == 1 for batch in values)
+ values = [batch[0] for batch in values]
+ else:
+ assert eof is None, "Single-row fetch did not reach EOF"
+ for n, value in enumerate(values):
+ if method == "fetchval":
+ assert type(value) is int and value == n
+ else:
+ expected = (n, n + 10 if shape == "numeric" else "text")
+ assert tuple(value) == expected
+ assert tuple(map(type, value)) == tuple(map(type, expected))
+ assert not cursor.messages, "Clean single-row fetch produced diagnostics"
+ return dict(
+ title=f"{shape} / {method}",
+ wall_ms=wall_ms,
+ cpp=cpp,
+ py=py,
+ detail=f"Rows: {SINGLE_ROW_COUNT}; shape: {shape}; API: {method}; EOF: 1",
+ )
+ finally:
+ ctx.disable()
+
+
+def single_row_registry():
+ return {
+ f"{shape}_{method}": partial(single_row_fetch, method=method, shape=shape)
+ for shape in ("numeric", "mixed")
+ for method in ("fetchone", "fetchmany", "fetchval")
+ }
diff --git a/mssql_python/cursor.py b/mssql_python/cursor.py
index 79436ab2e..d87953590 100644
--- a/mssql_python/cursor.py
+++ b/mssql_python/cursor.py
@@ -29,7 +29,7 @@
DatabaseError,
)
from mssql_python.row import Row
-from mssql_python.perf_timer import perf_phase
+from mssql_python.perf_timer import perf_phase, perf_start, perf_stop
from mssql_python import get_settings
from mssql_python.parameter_helper import (
detect_and_convert_parameters,
@@ -43,6 +43,37 @@
else:
pyarrow = None
+_DEFAULT_ROW_TYPE = Row
+_DEFAULT_FAST_ROW_CREATE = Row._fast_create
+_DEFAULT_FAST_ROW_CODE = _DEFAULT_FAST_ROW_CREATE.__code__
+_DEFAULT_FAST_ROW_DESCRIPTOR = vars(Row)["_fast_create"]
+# Describes the default routes; reporting reads this outside the fetch hot path.
+_DEFAULT_NATIVE_ROW_ROUTE = {
+ "version": 1,
+ "methods": {"fetchone": False, "fetchmany": False, "fetchval": False},
+}
+_DEFAULT_NATIVE_FETCH_ONE = ddbc_bindings.DDBCSQLFetchOne
+_DEFAULT_NATIVE_FETCH_MANY = ddbc_bindings.DDBCSQLFetchMany
+_NATIVE_ROW_PLAN = object()
+
+
+def _native_row_eligible(row_type):
+ # Inspect raw class entries: getattr would execute user descriptors before the factory.
+ if row_type is not _DEFAULT_ROW_TYPE or type(row_type) is not type:
+ return False
+ namespace = type.__getattribute__(row_type, "__dict__")
+ bases = type.__getattribute__(row_type, "__bases__")
+ return (
+ len(bases) == 1
+ and bases[0] is object
+ and "__new__" not in namespace
+ and "__setattr__" not in namespace
+ and namespace.get("_fast_create") is _DEFAULT_FAST_ROW_DESCRIPTOR
+ and _DEFAULT_FAST_ROW_CREATE.__code__ is _DEFAULT_FAST_ROW_CODE
+ and _DEFAULT_FAST_ROW_CREATE.__globals__.get("Row") is row_type
+ )
+
+
# Constants for string handling
MAX_INLINE_CHAR: int = (
4000 # NVARCHAR/VARCHAR inline limit; this triggers NVARCHAR(MAX)/VARCHAR(MAX) + DAE
@@ -646,8 +677,19 @@ def _refresh_decoding_cache(self):
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_fetch_options = None
self._cached_decoding_generation = generation
+ def _create_fetch_options(self):
+ generation = self._cached_decoding_generation
+ options = ddbc_bindings._FetchOptions(
+ self._cached_char_encoding, self._cached_wchar_encoding, self._cached_char_ctype
+ )
+ # Allocation callbacks can refresh decoding; never install an older snapshot.
+ if self._cached_decoding_generation == generation:
+ self._cached_fetch_options = options
+ return options
+
def _get_decoding_settings(self, sql_type):
"""
Get decoding settings for a specific SQL type.
@@ -2855,52 +2897,86 @@ def fetchone(self) -> Union[None, Row]:
# Fetch raw data
row_data = []
try:
- with perf_phase("py::fetchone::cpp_call"):
- ret = ddbc_bindings.DDBCSQLFetchOne(
- self.hstmt,
- row_data,
- char_enc,
- wchar_enc,
- self._cached_char_ctype,
- self.messages,
- )
+ started = perf_start()
+ try:
+ fetch = ddbc_bindings.DDBCSQLFetchOne
+ if fetch is _DEFAULT_NATIVE_FETCH_ONE:
+ options = self._cached_fetch_options
+ if options is None:
+ options = self._create_fetch_options()
+ ret = ddbc_bindings._fetchone_with_options(
+ self.hstmt, row_data, options, self.messages
+ )
+ else:
+ ret = fetch(
+ self.hstmt,
+ row_data,
+ char_enc,
+ wchar_enc,
+ self._cached_char_ctype,
+ self.messages,
+ )
+ finally:
+ if started:
+ perf_stop("py::fetchone::cpp_call", started)
- check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret)
+ return self._finish_fetchone(ret, row_data)
+ except Exception:
+ # On error, don't increment rownumber - rethrow the error
+ raise
- 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:
- self.rowcount = 0
- return None
+ def _finish_fetchone(self, ret, row_data, native=False):
+ check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret)
- # Update internal position after successful fetch
- if self._skip_increment_for_next_fetch:
- self._skip_increment_for_next_fetch = False
- self._next_row_index += 1
- else:
- self._increment_rownumber()
+ 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:
+ self.rowcount = 0
+ return None
- self.rowcount = self._next_row_index
+ # Update internal position after successful fetch
+ if self._skip_increment_for_next_fetch:
+ self._skip_increment_for_next_fetch = False
+ self._next_row_index += 1
+ else:
+ self._increment_rownumber()
- # 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
+ self.rowcount = self._next_row_index
+
+ # Get column and converter maps
+ column_map, converter_map, column_map_lower = self._get_column_and_converter_maps()
+ started = perf_start()
+ try:
+ if not converter_map and not self._uuid_str_indices:
+ if native and not started and _native_row_eligible(Row):
+ row_type = Row
+ factory = row_type._fast_create
+ column_names = self._cached_result_columns
+ return (
+ _NATIVE_ROW_PLAN,
+ row_type,
+ column_map,
+ self,
+ column_map_lower,
+ column_names,
+ factory,
+ _DEFAULT_FAST_ROW_CODE,
)
- return Row(
- row_data,
- column_map,
- cursor=self,
- converter_map=converter_map,
- uuid_str_indices=self._uuid_str_indices,
- column_map_lower=column_map_lower,
- column_names=self._cached_result_columns,
+ return Row._fast_create(
+ row_data, column_map, self, column_map_lower, self._cached_result_columns
)
- except Exception:
- # On error, don't increment rownumber - rethrow the error
- raise
+ return Row(
+ row_data,
+ column_map,
+ cursor=self,
+ converter_map=converter_map,
+ uuid_str_indices=self._uuid_str_indices,
+ column_map_lower=column_map_lower,
+ column_names=self._cached_result_columns,
+ )
+ finally:
+ if started:
+ perf_stop("py::fetchone::row_wrap", started)
def fetchmany(self, size: Optional[int] = None) -> List[Row]:
"""
@@ -2930,62 +3006,112 @@ def fetchmany(self, size: Optional[int] = None) -> List[Row]:
# Fetch raw data
rows_data = []
try:
- with perf_phase("py::fetchmany::cpp_call"):
- ret = ddbc_bindings.DDBCSQLFetchMany(
- self.hstmt,
- rows_data,
- size,
- char_enc,
- wchar_enc,
- self._cached_char_ctype,
- self.messages,
- )
-
- check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret)
-
- # Update rownumber for the number of rows actually fetched
- if rows_data and self._has_result_set:
- # advance counters by number of rows actually returned
- self._next_row_index += len(rows_data)
- self._rownumber = self._next_row_index - 1
-
- # Centralize rowcount assignment after fetch
- if len(rows_data) == 0 and self._next_row_index == 0:
- self.rowcount = 0
- else:
- self.rowcount = self._next_row_index
-
- # Get column and converter maps
- column_map, converter_map, column_map_lower = self._get_column_and_converter_maps()
-
- # 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,
+ started = perf_start()
+ try:
+ fetch = ddbc_bindings.DDBCSQLFetchMany
+ if fetch is _DEFAULT_NATIVE_FETCH_MANY:
+ options = self._cached_fetch_options
+ if options is None:
+ options = self._create_fetch_options()
+ ret = ddbc_bindings._fetchmany_with_options(
+ self.hstmt, rows_data, size, options, self.messages
)
- return [
- Row(
- row_data,
- column_map,
- cursor=self,
- converter_map=converter_map,
- uuid_str_indices=uuid_idx,
- column_map_lower=column_map_lower,
- column_names=self._cached_result_columns,
+ else:
+ ret = fetch(
+ self.hstmt,
+ rows_data,
+ size,
+ char_enc,
+ wchar_enc,
+ self._cached_char_ctype,
+ self.messages,
)
- for row_data in rows_data
- ]
+ finally:
+ if started:
+ perf_stop("py::fetchmany::cpp_call", started)
+
+ return self._finish_fetchmany(ret, rows_data, size=size)
except Exception:
# On error, don't increment rownumber - rethrow the error
raise
+ def _finish_fetchmany(self, ret, rows_data, native=False, size=1):
+ check_error(ddbc_sql_const.SQL_HANDLE_STMT.value, self.hstmt, ret)
+
+ # Update rownumber for the number of rows actually fetched
+ if rows_data and self._has_result_set:
+ # advance counters by number of rows actually returned
+ self._next_row_index += len(rows_data)
+ self._rownumber = self._next_row_index - 1
+
+ # Centralize rowcount assignment after fetch
+ if len(rows_data) == 0 and self._next_row_index == 0:
+ self.rowcount = 0
+ else:
+ self.rowcount = self._next_row_index
+
+ # Get column and converter maps
+ column_map, converter_map, column_map_lower = self._get_column_and_converter_maps()
+
+ # Convert raw data to Row objects
+ uuid_idx = self._uuid_str_indices
+ started = perf_start()
+ try:
+ if not converter_map and not uuid_idx:
+ if (
+ type(size) is int
+ and size == 1
+ and len(rows_data) == 1
+ and Row is _DEFAULT_ROW_TYPE
+ and Row._fast_create is _DEFAULT_FAST_ROW_CREATE
+ ):
+ if native and not started and _native_row_eligible(Row):
+ row_type = Row
+ factory = row_type._fast_create
+ column_names = self._cached_result_columns
+ return (
+ _NATIVE_ROW_PLAN,
+ row_type,
+ column_map,
+ self,
+ column_map_lower,
+ column_names,
+ factory,
+ _DEFAULT_FAST_ROW_CODE,
+ )
+ return [
+ Row._fast_create(
+ rows_data[0],
+ column_map,
+ self,
+ column_map_lower,
+ self._cached_result_columns,
+ )
+ ]
+ return ddbc_bindings.construct_rows(
+ rows_data,
+ Row,
+ column_map,
+ self,
+ column_map_lower,
+ self._cached_result_columns,
+ )
+ return [
+ Row(
+ row_data,
+ column_map,
+ cursor=self,
+ converter_map=converter_map,
+ uuid_str_indices=uuid_idx,
+ column_map_lower=column_map_lower,
+ column_names=self._cached_result_columns,
+ )
+ for row_data in rows_data
+ ]
+ finally:
+ if started:
+ perf_stop("py::fetchmany::row_wrap", started)
+
def fetchall(self) -> List[Row]:
"""
Fetch all (remaining) rows of a query result.
diff --git a/mssql_python/pybind/build.sh b/mssql_python/pybind/build.sh
index 05bc69568..d12b5dd0c 100755
--- a/mssql_python/pybind/build.sh
+++ b/mssql_python/pybind/build.sh
@@ -121,10 +121,12 @@ if [[ "$COVERAGE_MODE" == "true" && "$OS" == "Linux" ]]; then
else
if [[ "$OS" == "macOS" ]]; then
echo "[ACTION] Configuring for macOS (default build)"
- cmake -DMACOS_STRING_FIX=ON $PROFILING_FLAG "${SOURCE_DIR}"
+ cmake -DMACOS_STRING_FIX=ON -DCMAKE_BUILD_TYPE="${CMAKE_BUILD_TYPE:-Release}" \
+ $PROFILING_FLAG "${SOURCE_DIR}"
else
echo "[ACTION] Configuring for Linux with architecture: $DETECTED_ARCH"
- cmake -DARCHITECTURE="$DETECTED_ARCH" $PROFILING_FLAG "${SOURCE_DIR}"
+ cmake -DARCHITECTURE="$DETECTED_ARCH" -DCMAKE_BUILD_TYPE="${CMAKE_BUILD_TYPE:-Release}" \
+ $PROFILING_FLAG "${SOURCE_DIR}"
fi
fi
diff --git a/mssql_python/pybind/ddbc_bindings.cpp b/mssql_python/pybind/ddbc_bindings.cpp
index 4e4523955..0575b435e 100644
--- a/mssql_python/pybind/ddbc_bindings.cpp
+++ b/mssql_python/pybind/ddbc_bindings.cpp
@@ -23,6 +23,7 @@
#include // For std::memcpy
#include
#include
+#include
#include // std::forward
#include // CPython datetime API (PyDateTime_IMPORT, PyDateTime_GET_*, etc.)
@@ -3111,6 +3112,24 @@ SQLSMALLINT SQLNumResultCols_wrap(SqlHandlePtr statementHandle, py::handle messa
namespace {
+SQLSMALLINT CachedResultColumnCount(const SqlHandlePtr& handle, py::handle messages) {
+ const auto snapshot = handle->resultMetadata.snapshot();
+ if (snapshot.fullColumnCount >= 0) {
+ return snapshot.fullColumnCount;
+ }
+ const auto count = SQLNumResultCols_wrap(handle, messages);
+ handle->resultMetadata.publishFullColumnCount(snapshot.generation, count);
+ return count;
+}
+
+// Adopt a new cell reference even if appending to the caller's list fails.
+void AppendFetchedCell(py::list& row, PyObject* value) {
+ auto owned = steal(value);
+ if (!owned || PyList_Append(row.ptr(), owned.ptr()) < 0) {
+ throw py::error_already_set();
+ }
+}
+
py::dict GetFetchColumnMetadata(const py::list& columns, size_t index) {
return columns[index].cast();
}
@@ -3262,12 +3281,18 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT
while (true) {
++loopCount;
- std::vector chunk(DAE_CHUNK_SIZE, 0);
+ const size_t offset = buffer.size();
+ if (buffer.max_size() - offset < DAE_CHUNK_SIZE) {
+ throw std::length_error("LOB data exceeds maximum buffer size");
+ }
+ // Keep the existing read windows and zero-filled terminator room.
+ buffer.resize(offset + DAE_CHUNK_SIZE, 0);
+ char* chunk = buffer.data() + offset;
SQLLEN actualRead = 0;
{
// Release the GIL during blocking SQLGetData LOB streaming
py::gil_scoped_release release;
- ret = SQLGetData_ptr(hStmt, colIndex, cType, chunk.data(), DAE_CHUNK_SIZE, &actualRead);
+ ret = SQLGetData_ptr(hStmt, colIndex, cType, chunk, DAE_CHUNK_SIZE, &actualRead);
}
CaptureFetchDiagnostics(hStmt, ret, messages, true);
@@ -3310,11 +3335,13 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT
// Wide characters
size_t wcharSize = sizeof(SQLWCHAR);
if (bytesRead >= wcharSize && (bytesRead % wcharSize == 0)) {
- size_t wcharCount = bytesRead / wcharSize;
- std::vector alignedBuf(wcharCount);
- std::memcpy(alignedBuf.data(), chunk.data(), bytesRead);
- while (wcharCount > 0 && alignedBuf[wcharCount - 1] == 0) {
- --wcharCount;
+ while (bytesRead >= wcharSize) {
+ // The byte destination need not be aligned for SQLWCHAR.
+ const char* lastChar = chunk + bytesRead - wcharSize;
+ if (std::any_of(lastChar, lastChar + wcharSize,
+ [](char byte) { return byte != '\0'; })) {
+ break;
+ }
bytesRead -= wcharSize;
}
if (bytesRead < DAE_CHUNK_SIZE) {
@@ -3325,8 +3352,8 @@ py::object FetchLobColumnData(SQLHSTMT hStmt, SQLUSMALLINT colIndex, SQLSMALLINT
}
}
}
+ buffer.resize(offset + bytesRead);
if (bytesRead > 0) {
- buffer.insert(buffer.end(), chunk.begin(), chunk.begin() + bytesRead);
LOG("FetchLobColumnData: Appended %zu bytes at loop %d", bytesRead, loopCount);
}
if (ret == SQL_SUCCESS) {
@@ -3696,22 +3723,36 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
// Use Python's codec system to decode bytes.
const std::string decodeEncoding =
GetEffectiveCharDecoding(effectiveCharEnc);
- py::bytes raw_bytes(reinterpret_cast(dataBuffer.data()),
- static_cast(dataLen));
+ py::object decoded;
try {
- py::object decoded =
- raw_bytes.attr("decode")(decodeEncoding, "strict");
- row.append(decoded);
+ // bytes.decode rejects embedded NULs; the C codec API does not.
+ if (decodeEncoding.find('\0') != std::string::npos) {
+ PyErr_SetString(PyExc_ValueError, "embedded null character");
+ throw py::error_already_set();
+ }
+ decoded = steal(PyUnicode_Decode(
+ reinterpret_cast(dataBuffer.data()),
+ static_cast(dataLen),
+ decodeEncoding.c_str(), "strict"));
+ if (!decoded) throw py::error_already_set();
+ if (PyList_Append(row.ptr(), decoded.ptr()) < 0)
+ throw py::error_already_set();
LOG("SQLGetData: CHAR column %d decoded with '%s', %zu bytes "
"-> %zu chars",
i, decodeEncoding.c_str(), (size_t)dataLen,
py::len(decoded));
} catch (const py::error_already_set& e) {
+ if (e.matches(PyExc_MemoryError)) throw;
LOG_ERROR(
"SQLGetData: Failed to decode CHAR column %d with '%s': %s",
i, decodeEncoding.c_str(), e.what());
- // Return raw bytes as fallback
- row.append(raw_bytes);
+ // Preserve the existing codec-error bytes fallback.
+ decoded = steal(PyBytes_FromStringAndSize(
+ reinterpret_cast(dataBuffer.data()),
+ static_cast(dataLen)));
+ if (!decoded) throw py::error_already_set();
+ if (PyList_Append(row.ptr(), decoded.ptr()) < 0)
+ throw py::error_already_set();
}
} else {
// Buffer too small, fallback to streaming
@@ -3774,9 +3815,19 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
} else {
uint64_t fetchBufferSize =
(columnSize + 1) * sizeof(SQLWCHAR); // +1 for null terminator
- std::vector dataBuffer(columnSize + 1);
+ const size_t bufferChars = columnSize + 1;
+ SQLWCHAR inlineBuffer[64];
+ std::vector heapBuffer;
+ SQLWCHAR* dataBuffer;
+ if (bufferChars <= std::size(inlineBuffer)) {
+ std::fill_n(inlineBuffer, bufferChars, SQLWCHAR{});
+ dataBuffer = inlineBuffer;
+ } else {
+ heapBuffer.resize(bufferChars);
+ dataBuffer = heapBuffer.data();
+ }
SQLLEN dataLen;
- ret = SQLGetData_ptr(hStmt, i, SQL_C_WCHAR, dataBuffer.data(), fetchBufferSize,
+ ret = SQLGetData_ptr(hStmt, i, SQL_C_WCHAR, dataBuffer, fetchBufferSize,
&dataLen);
CaptureFetchDiagnostics(
hStmt, ret, messages,
@@ -3786,14 +3837,14 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
if (SQL_SUCCEEDED(ret)) {
if (dataLen > 0) {
uint64_t numCharsInData = dataLen / sizeof(SQLWCHAR);
- if (numCharsInData < dataBuffer.size()) {
+ if (numCharsInData < bufferChars) {
// Construct with explicit length: SQLGetData reports the
// exact number of characters via dataLen, so do not rely on
// null termination. This preserves embedded NULs and avoids
// any risk of reading past the valid range if the driver
// omits the terminator.
row.append(FetchText::from_utf16_native(
- reinterpret_cast(dataBuffer.data()),
+ reinterpret_cast(dataBuffer),
static_cast(numCharsInData * sizeof(SQLWCHAR))));
LOG("SQLGetData: Appended NVARCHAR string "
"length=%lu for column %d",
@@ -3847,7 +3898,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
ret = SQLGetData_ptr(hStmt, i, SQL_C_LONG, &intValue, 0, &indicator);
CaptureFetchDiagnostics(hStmt, ret, messages);
if (SQL_SUCCEEDED(ret) && indicator != SQL_NULL_DATA) {
- row.append(static_cast(intValue));
+ AppendFetchedCell(row, PyLong_FromLong(intValue));
} else {
row.append(py::none());
}
@@ -3863,7 +3914,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
break;
}
if (SQL_SUCCEEDED(ret)) {
- row.append(static_cast(smallIntValue));
+ AppendFetchedCell(row, PyLong_FromLong(smallIntValue));
} else {
LOG("SQLGetData: Error retrieving SQL_SMALLINT for column "
"%d - SQLRETURN=%d",
@@ -3882,7 +3933,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
break;
}
if (SQL_SUCCEEDED(ret)) {
- row.append(realValue);
+ AppendFetchedCell(row, PyFloat_FromDouble(realValue));
} else {
LOG("SQLGetData: Error retrieving SQL_REAL for column %d - "
"SQLRETURN=%d",
@@ -3962,7 +4013,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
break;
}
if (SQL_SUCCEEDED(ret)) {
- row.append(doubleValue);
+ AppendFetchedCell(row, PyFloat_FromDouble(doubleValue));
} else {
LOG("SQLGetData: Error retrieving SQL_DOUBLE/FLOAT for "
"column %d - SQLRETURN=%d",
@@ -3981,7 +4032,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
break;
}
if (SQL_SUCCEEDED(ret)) {
- row.append(static_cast(bigintValue));
+ AppendFetchedCell(row, PyLong_FromLongLong(bigintValue));
} else {
LOG("SQLGetData: Error retrieving SQL_BIGINT for column %d "
"- SQLRETURN=%d",
@@ -4152,7 +4203,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
break;
}
if (SQL_SUCCEEDED(ret)) {
- row.append(static_cast(tinyIntValue));
+ AppendFetchedCell(row, PyLong_FromLong(tinyIntValue));
} else {
LOG("SQLGetData: Error retrieving SQL_TINYINT for column "
"%d - SQLRETURN=%d",
@@ -4171,7 +4222,7 @@ SQLRETURN SQLGetData_wrap(SqlHandlePtr StatementHandle, SQLUSMALLINT colCount, p
break;
}
if (SQL_SUCCEEDED(ret)) {
- row.append(static_cast(bitValue));
+ AppendFetchedCell(row, PyBool_FromLong(bitValue != 0));
} else {
LOG("SQLGetData: Error retrieving SQL_BIT for column %d - "
"SQLRETURN=%d",
@@ -4250,6 +4301,7 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri
// Unbind any columns from previous fetch operations to avoid memory
// corruption
+ StatementHandle->unboundGeneration.reset();
SQLFreeStmt_ptr(StatementHandle->get(), SQL_UNBIND);
// Perform scroll operation
@@ -4276,10 +4328,13 @@ SQLRETURN SQLFetchScroll_wrap(SqlHandlePtr StatementHandle, SQLSMALLINT FetchOri
// For column in the result set, binds a buffer to retrieve column data
// TODO: Move to anonymous namespace, since it is not used outside this file
template
-SQLRETURN SQLBindColums(SQLHSTMT hStmt, ColumnBuffers& buffers, const Metadata& columnNames,
+SQLRETURN SQLBindColums(SqlHandlePtr handle, ColumnBuffers& buffers, const Metadata& columnNames,
SQLUSMALLINT numCols, int fetchSize, int charCtype = SQL_C_WCHAR,
py::handle messages = {}) {
PERF_TIMER("SQLBindColums");
+ // Invalidate before the first bind, including partial/failed binding.
+ handle->unboundGeneration.reset();
+ SQLHSTMT hStmt = handle->get();
SQLRETURN ret = SQL_SUCCESS;
const bool useWideChar = (charCtype == SQL_C_WCHAR);
// Bind columns based on their data types
@@ -4927,6 +4982,7 @@ struct FetchStateGuard {
ret = SQLSetStmtAttr_ptr(handle->get(), SQL_ATTR_ROWS_FETCHED_PTR, nullptr, 0);
break;
default:
+ handle->unboundGeneration.reset();
ret = SQLFreeStmt_ptr(handle->get(), SQL_UNBIND);
break;
}
@@ -4948,6 +5004,10 @@ struct FetchStateGuard {
}
};
+SQLRETURN FetchSingleRow(SqlHandlePtr handle, py::list& row, const std::string& charEncoding,
+ const std::string& wcharEncoding, int charCtype, py::handle messages,
+ SQLSMALLINT knownColumnCount = -1, bool scroll = false);
+
// FetchMany_wrap - Fetches multiple rows of data from the result set.
//
// @param StatementHandle: Handle to the statement from which data is to be
@@ -4976,8 +5036,10 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch
SQLRETURN ret = SQL_ERROR;
ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret);
SQLHSTMT hStmt = StatementHandle->get();
- // Retrieve column count
- SQLSMALLINT numCols = SQLNumResultCols_wrap(StatementHandle, messages);
+ // Keep count/name validation before advancing, including the size-one route.
+ SQLSMALLINT numCols = fetchSize == 1
+ ? CachedResultColumnCount(StatementHandle, messages)
+ : SQLNumResultCols_wrap(StatementHandle, messages);
// Retrieve column metadata
auto snapshot = StatementHandle->resultMetadata.snapshot();
@@ -5017,6 +5079,40 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch
ThrowStdException("Column metadata count does not match result column count");
}
+ // Only types with identical bound and SQLGetData conversions are eligible.
+ // Keep other types on the existing path, including their pre-fetch failures.
+ const bool singleNumericRow = fetchSize == 1 && numCols > 0 &&
+ std::all_of(columnNames.begin(), columnNames.end(), [](const auto& column) {
+ switch (column.dataType) {
+ case SQL_INTEGER:
+ case SQL_SMALLINT:
+ case SQL_BIGINT:
+ case SQL_TINYINT:
+ case SQL_BIT:
+ case SQL_REAL:
+ case SQL_DOUBLE:
+ case SQL_FLOAT:
+ return true;
+ default:
+ return false;
+ }
+ });
+ if (singleNumericRow) {
+ PERF_TIMER("FetchMany::single_numeric_row");
+ FetchStateGuard fetchStateGuard(StatementHandle, messages);
+ // No rows-fetched pointer is needed for a single unbound row.
+ fetchStateGuard.configure(nullptr, 1);
+ py::list row;
+ ret = FetchSingleRow(StatementHandle, row, charEncoding, wcharEncoding, charCtype,
+ messages, numCols, true);
+ CheckFetchError(StatementHandle, ret);
+ if (SQL_SUCCEEDED(ret)) {
+ rows.append(row);
+ }
+ fetchStateGuard.close();
+ return ret;
+ }
+
std::vector lobColumns;
for (SQLSMALLINT i = 0; i < numCols; i++) {
const auto& column = columnNames.at(i);
@@ -5059,7 +5155,8 @@ SQLRETURN FetchMany_wrap(SqlHandlePtr StatementHandle, py::list& rows, int fetch
FetchStateGuard fetchStateGuard(StatementHandle, messages);
// Bind columns
- ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, charCtype, messages);
+ ret = SQLBindColums(StatementHandle, buffers, columnNames, numCols, fetchSize, charCtype,
+ messages);
if (!SQL_SUCCEEDED(ret)) {
LOG("FetchMany_wrap: Error when binding columns - SQLRETURN=%d", ret);
return ret;
@@ -5380,7 +5477,8 @@ SQLRETURN FetchArrowBatch_wrap(SqlHandlePtr StatementHandle, py::list& capsules,
FetchStateGuard fetchStateGuard(StatementHandle, messages);
if (!hasLobColumns && fetchSize > 0) {
- ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, charCtype, messages);
+ ret = SQLBindColums(StatementHandle, buffers, columnNames, numCols, fetchSize, charCtype,
+ messages);
if (!SQL_SUCCEEDED(ret)) {
LOG("Error when binding columns");
return ret;
@@ -6238,7 +6336,8 @@ SQLRETURN FetchAll_wrap(SqlHandlePtr StatementHandle, py::list& rows,
FetchStateGuard fetchStateGuard(StatementHandle, messages);
// Bind columns
- ret = SQLBindColums(hStmt, buffers, columnNames, numCols, fetchSize, charCtype, messages);
+ ret = SQLBindColums(StatementHandle, buffers, columnNames, numCols, fetchSize, charCtype,
+ messages);
if (!SQL_SUCCEEDED(ret)) {
LOG("FetchAll_wrap: Error when binding columns - SQLRETURN=%d", ret);
return ret;
@@ -6285,28 +6384,42 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row,
// Issue #531: upgrade SQL_C_CHAR + utf-8 to SQL_C_WCHAR on Windows so the
// driver does lossless UTF-16 conversion instead of returning ACP bytes.
charCtype = EffectiveCharCtypeForFetch(charCtype, charEncoding);
+ return FetchSingleRow(StatementHandle, row, charEncoding, wcharEncoding, charCtype, messages);
+}
+
+SQLRETURN FetchSingleRow(SqlHandlePtr StatementHandle, py::list& row,
+ const std::string& charEncoding, const std::string& wcharEncoding,
+ int charCtype, py::handle messages, SQLSMALLINT knownColumnCount,
+ bool scroll) {
SQLRETURN ret = SQL_ERROR;
ResultMetadataFailureGuard metadataFailure(StatementHandle->resultMetadata, ret);
SQLHSTMT hStmt = StatementHandle->get();
- // Unbind any columns from previous fetch operations (e.g., fetchmany)
- // to avoid conflicts with SQLGetData. SQLGetData cannot be used on
- // columns that are already bound.
- ret = SQLFreeStmt_ptr(hStmt, SQL_UNBIND);
- CaptureFetchDiagnostics(hStmt, ret, messages);
- if (!SQL_SUCCEEDED(ret))
- return ret;
+ const auto generation = StatementHandle->resultMetadata.snapshot().generation;
+ if (StatementHandle->unboundGeneration != generation) {
+ StatementHandle->unboundGeneration.reset();
+ {
+ PERF_TIMER("FetchSingleRow::SQL_UNBIND");
+ ret = SQLFreeStmt_ptr(hStmt, SQL_UNBIND);
+ }
+ CaptureFetchDiagnostics(hStmt, ret, messages);
+ if (!SQL_SUCCEEDED(ret))
+ return ret;
+ StatementHandle->unboundGeneration = generation;
+ }
// Assume hStmt is already allocated and a query has been executed
{
// Release the GIL during the blocking ODBC fetch
py::gil_scoped_release release;
- ret = SQLFetch_ptr(hStmt);
+ ret = scroll ? SQLFetchScroll_ptr(hStmt, SQL_FETCH_NEXT, 0) : SQLFetch_ptr(hStmt);
}
CaptureFetchDiagnostics(hStmt, ret, messages);
if (SQL_SUCCEEDED(ret)) {
- // Retrieve column count
- SQLSMALLINT colCount = SQLNumResultCols_wrap(StatementHandle, messages);
+ SQLSMALLINT colCount = knownColumnCount;
+ if (colCount < 0) {
+ colCount = CachedResultColumnCount(StatementHandle, messages);
+ }
ret = SQLGetData_wrap(StatementHandle, colCount, row, charEncoding, wcharEncoding,
charCtype, messages);
if (!SQL_SUCCEEDED(ret)) {
@@ -6319,6 +6432,55 @@ SQLRETURN FetchOne_wrap(SqlHandlePtr StatementHandle, py::list& row,
return ret;
}
+// Python completion retains return validation, counters, maps and custom factory ordering.
+py::object FetchRow_wrap(SqlHandlePtr statement, py::list& data,
+ const std::string& charEncoding, const std::string& wcharEncoding,
+ int charCtype, py::handle messages, bool many,
+ const py::function& complete, const py::object& planToken,
+ const py::tuple& attributes, const py::str& rowGlobalName,
+ const py::str& newName, const py::str& setattrName) {
+ SQLRETURN ret = many
+ ? FetchMany_wrap(statement, data, 1, charEncoding, wcharEncoding, charCtype, messages)
+ : FetchOne_wrap(statement, data, charEncoding, wcharEncoding, charCtype, messages);
+ py::object result = complete(ret, data, true);
+ if (!PyTuple_CheckExact(result.ptr()) || PyTuple_GET_SIZE(result.ptr()) != 8 ||
+ PyTuple_GET_ITEM(result.ptr(), 0) != planToken.ptr()) {
+ return result;
+ }
+ if (many && PyList_GET_SIZE(data.ptr()) != 1) {
+ throw py::value_error("A single-row construction plan requires exactly one fetched row");
+ }
+ py::tuple plan = borrow(result.ptr());
+ py::object values = many ? borrow(PyList_GET_ITEM(data.ptr(), 0)) : borrow(data.ptr());
+ py::object factory = borrow(PyTuple_GET_ITEM(plan.ptr(), 6));
+ py::object row;
+ PyObject* factoryRowType = nullptr;
+ if (PyFunction_Check(factory.ptr())) {
+ factoryRowType = PyDict_GetItemWithError(PyFunction_GetGlobals(factory.ptr()),
+ rowGlobalName.ptr());
+ if (!factoryRowType && PyErr_Occurred()) {
+ throw py::error_already_set();
+ }
+ }
+ // Final argument evaluation or plan allocation can change the captured factory in place.
+ if (PyFunction_Check(factory.ptr()) &&
+ PyFunction_GetCode(factory.ptr()) == PyTuple_GET_ITEM(plan.ptr(), 7) &&
+ factoryRowType == PyTuple_GET_ITEM(plan.ptr(), 1) &&
+ RowFactory::has_default_row_allocation(factoryRowType, newName, setattrName)) {
+ PERF_TIMER("FetchRow::construct_row");
+ row = RowFactory::construct_row(values, plan[1].cast(), plan[2],
+ plan[3], plan[4], plan[5], attributes);
+ } else {
+ row = factory(values, plan[2], plan[3], plan[4], plan[5]);
+ }
+ if (!many) {
+ return row;
+ }
+ py::list rows(1);
+ PyList_SET_ITEM(rows.ptr(), 0, row.release().ptr());
+ return rows;
+}
+
// Wrap SQLMoreResults
SQLRETURN SQLMoreResults_wrap(SqlHandlePtr StatementHandle) {
PERF_TIMER("SQLMoreResults_wrap");
@@ -6514,6 +6676,20 @@ PYBIND11_MODULE(ddbc_bindings, m) {
py::arg("StatementHandle"), py::arg("colCount"), py::arg("row"), py::arg("charEncoding"),
py::arg("wcharEncoding"), py::arg("charCtype"), py::arg("messages") = py::none());
m.def("DDBCSQLMoreResults", &SQLMoreResults_wrap, "Check for more results in the result set");
+ py::class_(m, "_FetchOptions")
+ .def(py::init());
+ m.def("_fetchone_with_options",
+ [](SqlHandlePtr statement, py::list& row, const FetchOptions& options,
+ py::handle messages) {
+ return FetchOne_wrap(std::move(statement), row, options.charEncoding,
+ options.wcharEncoding, options.charCtype, messages);
+ });
+ m.def("_fetchmany_with_options",
+ [](SqlHandlePtr statement, py::list& rows, int size, const FetchOptions& options,
+ py::handle messages) {
+ return FetchMany_wrap(std::move(statement), rows, size, options.charEncoding,
+ options.wcharEncoding, options.charCtype, messages);
+ });
m.def("DDBCSQLFetchOne", &FetchOne_wrap, "Fetch one row from the result set",
py::arg("StatementHandle"), py::arg("row"), py::arg("charEncoding") = "utf-16le",
py::arg("wcharEncoding") = "utf-16le", py::arg("charCtype") = SQL_C_WCHAR,
@@ -6625,6 +6801,10 @@ PYBIND11_MODULE(ddbc_bindings, m) {
// Add a version attribute
m.attr("__version__") = "1.0.0";
+ m.def("_apply_output_converters", &RowFactory::apply_output_converters,
+ "Apply a cached converter map after native value materialization",
+ py::arg("values"), py::arg("converters"));
+
// 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",
@@ -6633,6 +6813,23 @@ PYBIND11_MODULE(ddbc_bindings, m) {
py::arg("column_map_lower") = py::none(),
py::arg("column_names") = py::none());
+ // Owned by binding defaults and released with the module; no static Python handles.
+ const py::tuple rowAttributes = py::make_tuple(
+ "_values", "_column_map", "_cursor", "_column_map_lower", "_column_names");
+ m.def("construct_row", &RowFactory::construct_row,
+ "Build one compatible Row with checked ownership",
+ py::arg("values"), py::arg("row_class"), py::arg("column_map"), py::arg("cursor"),
+ py::arg("column_map_lower") = py::none(), py::arg("column_names") = py::none(),
+ py::arg("attributes") = rowAttributes);
+ m.def("DDBCSQLFetchRow", &FetchRow_wrap,
+ "Fetch and construct one compatible Row with ordered Python completion",
+ py::arg("statement"), py::arg("data"), py::arg("char_encoding"),
+ py::arg("wchar_encoding"), py::arg("char_ctype"), py::arg("messages"),
+ py::arg("many"), py::arg("complete"), py::arg("plan_token"),
+ py::arg("attributes") = rowAttributes, py::arg("row_global_name") = py::str("Row"),
+ py::arg("new_name") = py::str("__new__"),
+ py::arg("setattr_name") = py::str("__setattr__"));
+
// 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/ddbc_bindings.h b/mssql_python/pybind/ddbc_bindings.h
index 32d9f8067..bd0c19c6a 100644
--- a/mssql_python/pybind/ddbc_bindings.h
+++ b/mssql_python/pybind/ddbc_bindings.h
@@ -7,6 +7,7 @@
#include
#include
#include
+#include
#include
#include
#include
@@ -331,6 +332,9 @@ class SqlHandle {
std::unordered_map describeCache;
void clearDescribeCache() { describeCache.clear(); }
ResultMetadataCache resultMetadata;
+ // Written only under the GIL. Metadata invalidation (including cancellation)
+ // changes the generation, so an old successful unbind cannot authorize reuse.
+ std::optional unboundGeneration;
private:
// The caller must release the GIL before waiting for native cleanup.
@@ -422,6 +426,15 @@ struct DateTimeOffset {
SQLSMALLINT timezone_minute; // Offset minutes from UTC
};
+struct FetchOptions {
+ const std::string charEncoding;
+ const std::string wcharEncoding;
+ const int charCtype;
+
+ FetchOptions(const std::string& charEncoding, const std::string& wcharEncoding, int charCtype)
+ : charEncoding(charEncoding), wcharEncoding(wcharEncoding), charCtype(charCtype) {}
+};
+
// Struct to hold data buffers and indicators for each column
struct ColumnBuffers {
std::vector> charBuffers;
diff --git a/mssql_python/pybind/result_metadata.hpp b/mssql_python/pybind/result_metadata.hpp
index 3a9912860..72ec0010b 100644
--- a/mssql_python/pybind/result_metadata.hpp
+++ b/mssql_python/pybind/result_metadata.hpp
@@ -30,11 +30,12 @@ class ResultMetadataCache {
struct Snapshot {
uint64_t generation;
std::shared_ptr metadata;
+ SQLSMALLINT fullColumnCount;
};
Snapshot snapshot() const {
std::lock_guard lock(mutex_);
- return {generation_, metadata_};
+ return {generation_, metadata_, fullColumnCount_};
}
void publish(uint64_t generation, std::shared_ptr metadata) {
@@ -44,16 +45,26 @@ class ResultMetadataCache {
}
}
+ void publishFullColumnCount(uint64_t generation, SQLSMALLINT columnCount) {
+ std::lock_guard lock(mutex_);
+ if (generation == generation_ && columnCount >= 0) {
+ fullColumnCount_ = columnCount;
+ }
+ }
+
void clear() {
std::lock_guard lock(mutex_);
++generation_;
metadata_.reset();
+ fullColumnCount_ = -1;
}
private:
mutable std::mutex mutex_;
uint64_t generation_ = 0;
std::shared_ptr metadata_;
+ // Prefix GetData metadata does not establish the full result cardinality.
+ SQLSMALLINT fullColumnCount_ = -1;
};
class ResultMetadataFailureGuard {
diff --git a/mssql_python/pybind/row_factory.hpp b/mssql_python/pybind/row_factory.hpp
index 85d271128..61c33b66c 100644
--- a/mssql_python/pybind/row_factory.hpp
+++ b/mssql_python/pybind/row_factory.hpp
@@ -5,8 +5,109 @@
#include "py_ref.hpp"
+#include
+
namespace RowFactory {
+// Only the post-fetch cached-map loop: no ODBC access, Row allocation, or UUID work.
+inline py::list apply_output_converters(const py::object& values, const py::object& converters) {
+ if (!PyList_CheckExact(values.ptr()) || !PyList_CheckExact(converters.ptr())) {
+ throw py::type_error("converter values and map must be exact lists");
+ }
+ py::list result =
+ steal(PyList_GetSlice(values.ptr(), 0, PyList_GET_SIZE(values.ptr())));
+ if (!result)
+ throw py::error_already_set();
+
+ // Keep the Python iterators: their retained tuples affect finalizer timing
+ // when a callback replaces itself or mutates the source lists.
+ py::object pairs = steal(PyObject_CallFunctionObjArgs(reinterpret_cast(&PyZip_Type),
+ values.ptr(), converters.ptr(), nullptr));
+ if (!pairs)
+ throw py::error_already_set();
+ py::object items =
+ steal(PyObject_CallOneArg(reinterpret_cast(&PyEnum_Type), pairs.ptr()));
+ if (!items)
+ throw py::error_already_set();
+ pairs = py::object();
+
+ // Retain the current inputs and last encoded value like the Python locals.
+ py::object value, converter, value_bytes;
+ for (Py_ssize_t i = 0;; ++i) {
+ py::object item = steal(PyIter_Next(items.ptr()));
+ if (!item) {
+ if (PyErr_Occurred())
+ throw py::error_already_set();
+ break;
+ }
+ PyObject* pair = PyTuple_GET_ITEM(item.ptr(), 1);
+ py::object next_value = borrow(PyTuple_GET_ITEM(pair, 0));
+ py::object next_converter = borrow(PyTuple_GET_ITEM(pair, 1));
+ value = std::move(next_value);
+ converter = std::move(next_converter);
+ item = py::object();
+ const int enabled = PyObject_IsTrue(converter.ptr());
+ if (enabled < 0)
+ throw py::error_already_set();
+ if (!enabled || value.is_none())
+ continue;
+
+ try {
+ const int is_string =
+ PyObject_IsInstance(value.ptr(), reinterpret_cast(&PyUnicode_Type));
+ if (is_string < 0)
+ throw py::error_already_set();
+ PyObject* argument = value.ptr();
+ if (is_string) {
+ // Match str.encode's codec lookup, including the spelling.
+ // Subclasses and __class__ proxies still dispatch encode dynamically.
+ py::object encoded =
+ steal(PyUnicode_CheckExact(value.ptr())
+ ? PyUnicode_AsEncodedString(value.ptr(), "utf-16-le", nullptr)
+ : PyObject_CallMethod(value.ptr(), "encode", "s", "utf-16-le"));
+ if (!encoded)
+ throw py::error_already_set();
+ value_bytes = std::move(encoded);
+ argument = value_bytes.ptr();
+ }
+ py::object converted = steal(PyObject_CallOneArg(converter.ptr(), argument));
+ if (!converted)
+ throw py::error_already_set();
+ // Checked assignment also preserves Python's caught IndexError if a
+ // callback grows the input lists beyond the initial result copy.
+ if (PyList_SetItem(result.ptr(), i, converted.release().ptr()) < 0)
+ throw py::error_already_set();
+ } catch (py::error_already_set& error) {
+ if (!error.matches(PyExc_Exception))
+ throw;
+ // Preserve the existing cached path's keep-original-on-Exception contract.
+ error.restore();
+ PyErr_Clear();
+ }
+ }
+ items = py::object();
+ return result;
+}
+
+inline void initialize_row(
+ const py::object& row, PyObject* row_data, const py::object& column_map,
+ const py::object& cursor_obj, const py::object& column_map_lower,
+ const py::object& column_names, const py::handle& attr_values, const py::handle& attr_column_map,
+ const py::handle& attr_cursor, const py::handle& attr_column_map_lower,
+ const py::handle& attr_column_names,
+ int (*set_attr)(PyObject*, PyObject*, PyObject*) = PyObject_GenericSetAttr) {
+ if (!row)
+ throw py::error_already_set();
+
+ if (set_attr(row.ptr(), attr_values.ptr(), row_data) < 0 ||
+ set_attr(row.ptr(), attr_column_map.ptr(), column_map.ptr()) < 0 ||
+ set_attr(row.ptr(), attr_cursor.ptr(), cursor_obj.ptr()) < 0 ||
+ set_attr(row.ptr(), attr_column_map_lower.ptr(), column_map_lower.ptr()) < 0 ||
+ set_attr(row.ptr(), attr_column_names.ptr(), column_names.ptr()) < 0) {
+ throw py::error_already_set();
+ }
+}
+
// 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,
@@ -36,22 +137,67 @@ inline py::list construct_rows(const py::list& rows_data, const py::object& row_
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);
+ initialize_row(row, row_data, column_map, cursor_obj, column_map_lower, column_names,
+ attr_values, attr_column_map, attr_cursor, attr_column_map_lower,
+ attr_column_names);
+ PyList_SET_ITEM(result.ptr(), i, row.release().ptr());
+ }
+
+ return result;
+}
- 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) {
+// Passive final eligibility check: do not execute newly installed class descriptors.
+inline bool has_default_row_allocation(PyObject* row_class, const py::str& new_name,
+ const py::str& setattr_name) {
+ if (!PyUnicode_CheckExact(new_name.ptr()) || !PyUnicode_CheckExact(setattr_name.ptr())) {
+ throw py::type_error("Row allocation guard names must be exact strings");
+ }
+ if (!row_class || Py_TYPE(row_class) != &PyType_Type) {
+ return false;
+ }
+ auto* row_type = reinterpret_cast(row_class);
+ if (!row_type->tp_bases || !PyTuple_CheckExact(row_type->tp_bases) ||
+ PyTuple_GET_SIZE(row_type->tp_bases) != 1 ||
+ PyTuple_GET_ITEM(row_type->tp_bases, 0) != reinterpret_cast(&PyBaseObject_Type) ||
+ !row_type->tp_dict) {
+ return false;
+ }
+ PyObject* names[] = {new_name.ptr(), setattr_name.ptr()};
+ for (PyObject* name : names) {
+ PyObject* member = PyDict_GetItemWithError(row_type->tp_dict, name);
+ if (member) {
+ return false;
+ }
+ if (PyErr_Occurred()) {
throw py::error_already_set();
}
-
- PyList_SET_ITEM(result.ptr(), i, row.release().ptr());
}
+ return true;
+}
- return result;
+// Attribute names are prepared once as binding defaults, not allocated for each row.
+inline py::object construct_row(const py::object& values, const py::type& row_class,
+ const py::object& column_map, const py::object& cursor_obj,
+ const py::object& column_map_lower, const py::object& column_names,
+ const py::tuple& attributes) {
+ if (attributes.size() != 5) {
+ throw py::value_error("Row construction requires five attribute names");
+ }
+ for (py::handle name : attributes) {
+ if (!PyUnicode_Check(name.ptr())) {
+ throw py::type_error("Row attribute names must be strings");
+ }
+ }
+ const py::object new_method = row_class.attr("__new__");
+ py::object row = steal(PyObject_CallOneArg(new_method.ptr(), row_class.ptr()));
+ initialize_row(row, values.ptr(), column_map, cursor_obj, column_map_lower, column_names,
+ py::handle(PyTuple_GET_ITEM(attributes.ptr(), 0)),
+ py::handle(PyTuple_GET_ITEM(attributes.ptr(), 1)),
+ py::handle(PyTuple_GET_ITEM(attributes.ptr(), 2)),
+ py::handle(PyTuple_GET_ITEM(attributes.ptr(), 3)),
+ py::handle(PyTuple_GET_ITEM(attributes.ptr(), 4)), PyObject_SetAttr);
+ return row;
}
} // namespace RowFactory
diff --git a/mssql_python/row.py b/mssql_python/row.py
index ccec53787..7eb766b0d 100644
--- a/mssql_python/row.py
+++ b/mssql_python/row.py
@@ -9,6 +9,7 @@
import uuid as _uuid
from collections.abc import Mapping
from typing import Any
+from mssql_python import ddbc_bindings
from mssql_python.logging import logger
@@ -198,6 +199,9 @@ def _apply_output_converters_optimized(self, values, converter_map):
"""
Apply output converters using pre-computed converter map for optimal performance.
+ Native materialization has already completed. Exact lists use a native
+ dispatch loop; both paths copy values and preserve dynamic string encoding.
+
Args:
values: Raw values from the database
converter_map: Pre-computed list of converters (one per column, None if no converter)
@@ -205,6 +209,11 @@ def _apply_output_converters_optimized(self, values, converter_map):
Returns:
List of converted values
"""
+ # Native fetches and the cursor cache supply exact lists. Keep arbitrary
+ # iterables on the Python path, including their iteration side effects.
+ if type(values) is list and type(converter_map) is list:
+ return ddbc_bindings._apply_output_converters(values, converter_map)
+
converted_values = list(values)
for i, (value, converter) in enumerate(zip(values, converter_map)):
diff --git a/profiler/README.md b/profiler/README.md
index 8c0662b14..a918c2a03 100644
--- a/profiler/README.md
+++ b/profiler/README.md
@@ -144,6 +144,113 @@ without re-entering a held counter lock. Python samples triggered recursively du
profiler bookkeeping are intentionally omitted; the cleanup itself still runs.
Ordinary nested workload timers and other threads are not suppressed.
+## Single-row fetch attribution
+
+`fetchone()` and the default iterator/`fetchval()` path share the native
+single-row helper. Full result column counts are cached per metadata generation,
+separately from prefix `SQLGetData` metadata. `fetchmany(1)` reuses that count for
+all result shapes, while still validating count and names before fetching. Cold
+`fetchmany(1)` can make two count calls (eager validation and `DescribeColumns`);
+subsequent size-one calls, including EOF, do not reacquire it within the same
+generation. Larger batches and direct `DDBCSQLNumResultCols` calls retain their
+uncached count behavior. `fetchval()` still calls `fetchone()` and constructs the full
+row, including converters for columns beyond the first.
+
+`ddbc::FetchSingleRow::SQL_UNBIND` counts attempted unbinds in that helper. A successful
+unbind can be reused within the same generation, but does not certify row-array
+attributes. Binding (including Arrow and partial binds), cleanup, and generation
+changes invalidate reuse.
+
+`ddbc::FetchMany::single_numeric_row` identifies the native `fetchmany(1)` route for
+all-numeric results (integer, bit, real, float/double). It retains eager count/name
+validation, row-array configuration and cleanup, and uses `SQLFetchScroll` plus
+per-column `SQLGetData`. Mixed INT/NVARCHAR, text, LOB, decimal, temporal, UUID and
+variant results retain their existing native routes. Numeric width can therefore
+change the tradeoff; do not generalize a narrow-row result.
+
+The Python one-row wrapping shortcut is independent of native eligibility. It
+requires a built-in `int` request of one, exactly one returned row, the canonical
+`Row` and factory, and no converter or UUID work. One-row tails of larger requests
+and substituted factories retain batch wrapping. The existing
+`py::fetchone::{cpp_call,row_wrap}` and `py::fetchmany::{cpp_call,row_wrap}` phases
+use paired start/stop calls, avoiding context-manager entry/exit when disabled.
+
+Row-wise integer, bit and floating-point values use checked CPython constructors
+and append operations. This removes generic scalar marshalling, not the scalar
+allocations themselves. The public `DDBCSQLGetData` destination still retains its
+identity, preexisting elements and completed cells on a later-column error; it
+is not replaced with a presized, partially initialized list.
+
+Bounded narrow-character GetData decodes the driver buffer directly through
+Python's codec API. Successful decoding avoids an intermediate bytes object,
+bound `decode` method and Python call arguments. Codec failures retain the logged
+bytes fallback, including codec names containing an embedded NUL rather than
+silently truncating them for the C API. Allocation failures now propagate instead
+of being mistaken for codec failures. The decoded value is still appended before
+debug logging; if a custom string's length raises during logging, the existing
+decoded cell and subsequent bytes fallback are retained. Wide text and LOB
+streaming are unchanged.
+
+Mixed INT/NVARCHAR is intentionally not routed through the numeric specialization:
+bound and GetData paths differ on malformed UTF-16, `SQL_NO_TOTAL`, truncation
+continuation and diagnostic timing. The unbind witness does not justify removing
+row-array configuration/cleanup. Row user attributes, weakrefs and substituted
+factories remain supported, as do dynamic `fetchone` overrides. No
+first-column-only `fetchval` route is introduced.
+
+The experimental shared `DDBCSQLFetchRow` entry performs the original native fetch,
+then calls the extracted Python completion body for return checking, row counters
+and map acquisition. An eligible completion returns a private construction plan;
+the same native call builds the Row and returns it (or a one-element list for
+`fetchmany(1)`). Unsupported/customized completion constructs the already-fetched
+values in Python. It never refetches. `fetchval()` still dispatches through
+`self.fetchone()` and all-column conversions remain in the completion body.
+
+Eligibility inspects raw class dictionaries and the original factory code without
+executing allocation/assignment descriptors or metaclass hooks. It is checked
+before fetching and again after map acquisition. The captured factory code and
+Row global are checked again in native code after final argument evaluation;
+late changes call that captured factory on the fetched values, without refetching.
+A final native raw-type/dictionary guard also rejects newly installed allocation or
+assignment hooks without invoking them. The captured Python factory then preserves
+`Row.__new__(Row)` lookup order, including a descriptor changing the second `Row`.
+In-place factory changes, custom
+allocation/assignment, converters and UUID policies retain their fallback behavior.
+The original `Row._fast_create` body is unchanged. Constructor attribute names are
+an immutable tuple owned by binding defaults, prepared once per module rather
+than five Python string allocations per row; no static owning Python handles or
+Row layout offsets are used. Native construction still uses `__new__` validation
+and normal checked attribute assignment, preserving descriptor callbacks.
+
+This is a shared fetch-and-construction call, **not** an all-native cursor:
+completion still re-enters Python, and eligible rows allocate a small plan tuple.
+The ordinary real-source frame trace removes one `_fast_create` frame but adds
+two `_native_row_eligible` frames and one `_finish_fetch*` frame: net **two more
+Python frames**. The outer Python-to-native call count is unchanged, and completion
+adds a native-to-Python crossing. This is not reduced Python dispatch or fewer
+crossings. These costs may outweigh native assignment work; no speedup or slowdown
+is established without measurement. Larger requests, integer subclasses and substituted low-level
+fetch bindings retain the original entry. Native numeric/non-numeric routing and
+cleanup are unchanged.
+
+When Python phase profiling is active, the original split entry is used so
+`cpp_call` and `row_wrap` retain their existing boundaries. With Python phases off,
+`ddbc::FetchRow::construct_row` attributes native construction when native profiling
+is available. Timing the split Python-profiled path is not evidence of fused-path
+latency; comparisons must identify which route actually ran.
+
+The isolated `test_single_row_fusion_native_counters_in_subprocess` cases enable
+native counters with Python phases disabled. For each API they require two native
+Row constructions for two rows, none at EOF, and none for converter fallback while
+all-column callbacks still run. These three cases skip on native profiling-OFF
+builds; the constructor-frame tests still run there. Split-route profiling alone
+cannot qualify fusion. These are future exact-source CI contracts, not local
+runtime measurements.
+
+Use profiling-enabled builds to check operation counts and normal uninstrumented
+Release builds for latency comparisons. These routes are optimization hypotheses,
+not a measured speedup; source-only checks do not establish native correctness.
+
## Adding a timer
To time a new spot in the code:
@@ -168,3 +275,166 @@ void MyFunction(...) {
`PERF_TIMER` compiles to nothing unless the build has `ENABLE_PROFILING`, so
adding timers costs nothing in released builds.
+
+
+### Explicit single-row comparison modes
+
+The paired controller retains its existing `diagnostic` default: native and Python
+recording are enabled. That instrumented report is not an uninstrumented latency
+measurement. Default fetching completes Rows in Python, even when the native
+Row-construction binding is available.
+
+Comparisons require Release builds on both sides. The controller sets
+`CMAKE_BUILD_TYPE=Release` for archived sources as well as the candidate; old Unix
+build scripts require CMake 3.22 or newer to consume that environment setting.
+Build checks reject a missing or non-Release configuration and retain the generator
+and C++ flags from `CMakeCache.txt` in build/worker logs. Unix non-coverage builds
+also explicitly configure Release by default. Historical measurements without
+verified Release configuration are not evidence of shipped Release latency.
+
+Two opt-in modes use the same six read-only workloads: `numeric_fetchone`,
+`numeric_fetchmany`, `numeric_fetchval`, and their `mixed_` equivalents. Each fetches
+1,000 ordered rows plus EOF; numeric rows contain two INT columns, mixed rows contain
+INT and NVARCHAR. `fetchmany` always requests the built-in integer `1`. `fetchval`
+uses its public API, not a first-column-only native shortcut.
+
+- `--mode latency` builds both revisions with native instrumentation OFF and keeps
+ Python phases OFF. It rejects `--reuse-candidate` rather than silently timing the
+ profiling-enabled CI extension. Setup/execute and validation are outside the
+ timed window; the API loop, result retention and EOF call are inside it.
+- `--mode route` uses native instrumentation ON with Python phases OFF. Each case
+ must record 1,001 native fetch calls. The version-1 default route contract requires
+ zero native Row constructions for all three APIs in the current implementation.
+ Here `1` is a built-in integer for both numeric and
+ mixed shapes; this is not the numeric-column-only native fetching shortcut.
+ Legacy f539 requires 1,000 constructions for all three APIs; the intermediate
+ many-only route requires 1,000 for `fetchmany(1)`; main666 and the current route
+ require none. This is route attribution, not production latency.
+ Existing native regression tests separately cover converter fallback, all-column
+ callbacks, repeated EOF, and customization failures.
+
+For example, after obtaining approval for the required builds and database work:
+
+```console
+python -m eng.profiler_benchmarks.controller --mode latency --base --candidate --leg Linux-SQL2022 --output
+python -m eng.profiler_benchmarks.controller --mode route --base --candidate --leg Linux-SQL2022 --output
+```
+
+Run these separately, not concurrently. No extra CI arm or automatic execution is
+introduced. Each invocation retains the controller's existing bounded sample,
+worker and overall time limits. Use distinct output directories; subset runs remain
+incomplete and cannot produce a full report. The standalone interactive Profiler's
+recording and timeline defaults are unchanged.
+
+These modes emit schema version 2 with explicit mode, worker revision, actual native
+binary path/SHA256 and compile-capability identity. The validator rejects mixed modes,
+changed binaries, missing/failed samples, unexpected recording, and wrong route counts.
+Schema version 1 remains the existing 22-workload diagnostic report. Failed runs stay
+incomplete/unavailable; missing timings are never replaced with zero. Raw samples
+retain every paired observation. Tables show signed subthreshold changes and median
+candidate/base ratios with their observed min/max range, not confidence intervals.
+The classification policy remains **more than 20% median paired change, at least
+1 ms between median times, and at least 80% of pairs beyond the relative threshold**.
+A subthreshold change is not a proven win or proof of no effect. Neither mode adds a
+pyodbc comparison or changes production fetch behavior.
+
+
+### Latency-first PR CI report
+
+The existing Ubuntu/SQL Server 2022 and Ubuntu/SQL Server 2025 PR legs invoke
+`--reuse-candidate --ci-report`. Windows profiling remains disabled. The aggregate
+runs OFF/OFF latency first, ON/OFF native-route attribution second, and the original
+22-workload ON/ON diagnostics last. The six numeric/mixed workloads in each fetch
+mode and all legacy workloads retain **five measured pairs and one warmup pair**.
+Each mode alternates base/candidate order independently, using separate workers.
+
+The aggregate builds isolated base-OFF, candidate-OFF and base-ON directories and
+reuses the already tested candidate-ON build. Route and diagnostic workers must
+attest the same ON binary paths, SHA256 digests, source revisions and environments.
+Each reused-checkout worker also checks HEAD and tracked driver/provider source
+cleanliness before importing the driver. The reused build is never rebuilt or
+toggled in place. Per leg this permits at most
+three performance-step builds, 36 measurement workers and 408 workload executions
+(including warmups), rather than one build, 12 workers and 264 executions. Both
+active legs together add at most four builds, 48 workers and 288 executions. No new
+matrix arm, pyodbc comparison or database setup is introduced.
+
+The controller retains a **90-minute aggregate budget**, including a three-minute
+finish reserve; the pipeline step remains 100 minutes and its job 160 minutes.
+Archives/preflight, builds and workers receive bounded timeouts clipped to the
+remaining work budget. Fetch workers receive at most 30 seconds; diagnostic workers
+retain their six-minute cap. All workers completed within the allowance in ADO180876
+on the two observed agents. That observation is not a future-host guarantee or
+qualification of this successor.
+The proposal already required 132 minutes if its build/worker caps and planning
+overhead were all consumed. Metadata operations also consume the shared deadline;
+they never extend it. Completing every mode is **not guaranteed**. Unused early time flows to later modes. Exhaustion or
+failure leaves that mode unavailable with a stage/reason and its raw checkpoints;
+there are no retries, trimmed workload sets, reduced pair counts or zero substitutes.
+Unsafe process cleanup aborts subsequent work and reports the retained temporary
+root rather than removing files beneath an unreaped worker.
+
+There is still exactly one `report.json` per existing artifact. Its root remains the
+schema-1 diagnostic report and its `status` describes diagnostics only. The additive
+`measurement_bundle_version: 1` extension contains `fetch_measurements.latency` and
+`fetch_measurements.route` (schema 2). Each mode has an independent complete/incomplete
+status. Root and child checkpoints are written atomically; an incomplete aggregate
+exits nonzero, while the pipeline's existing failure-tolerant artifact publication
+retains independently valid modes. The finish deadline is checked after final
+validation and the atomic write. An overrun gets one corrective incomplete checkpoint
+and a nonzero exit, preserving completed latency/route measurements without a recheck
+loop. Raw build logs and mode-prefixed worker JSON/log files stay in the same artifact,
+without nested files named `report.json`.
+
+The updated collector validates shared PR/build/base/merge identity, then each mode
+independently. The PR comment restores the original 22-task profiling-enabled
+**diagnostic** headline, section ordering and timing tables; it is not a shipped-latency
+verdict. Separate OFF/OFF latency and ON/OFF route measurements, when available, remain
+in the raw artifacts and are not headline inputs. Missing, malformed or incomplete
+diagnostics are explicitly unavailable even when the latency or route measurements
+complete. The standalone latency and route reports remain available with their existing
+labels. The updated reader accepts historical schema-1/schema-2
+reports as well as the additive route contract. Old strict five-key provenance readers
+reject new fetch workers; deploy producer and trusted consumer changes together rather
+than stripping fields. Interactive Profiler defaults remain unchanged.
+Fork reporting continues to use trusted base code. Existing artifact, download,
+comment-size and publisher time limits are not expanded.
+
+This successor is source-only until approved exact-head CI runs it. Source/fake-boundary
+controls and prior-head CI do not establish native correctness, completion within
+these allowances, or a speedup.
+
+### Default route identity
+
+`guarded_row` records native binding capability, not which Python API uses it.
+New workers carry exactly seven provenance fields: `source_commit`, `native_file`,
+`native_sha256`, `native_profiling`, `guarded_row`, `python_cursor_sha256`, and
+`row_route`. The latter is `{version: 1, methods: {fetchone: false, fetchmany: false,
+fetchval: false}}` for the current implementation. `guarded_row` may still be true:
+the low-level binding remains available but is not used by these default routes.
+The private literal in `cursor.py` is descriptive;
+it is never consulted by fetching. `fetchval` still calls the dynamically resolved
+`self.fetchone()`. The private `_finish_fetchone` completion now observes `native=False`
+for default fetchone/fetchval instead of the previous eligible `native=True` route.
+
+The controller independently reads each selected source root before measurement and
+records `python_sources.base` and `.candidate` anchors (revision, cursor digest, route).
+Workers verify their imported cursor origin and match these anchors. A closed recognizer
+accepts only reviewed main666, f539 and successor method bodies. Its constants hash the
+LF-normalized, dedented source segments of fetchone/fetchmany/fetchval, in that order,
+as a compact ASCII-escaped JSON list. Decorated methods, duplicate methods or literal
+keys, and non-definition class-scope bindings of the three route names are rejected.
+The binding check does not inspect method locals or nested scopes; it is a narrow
+recognizer guard, not analysis of arbitrary dynamic Python namespace mutations. Method-body, comment or docstring edits require explicit fingerprint
+review/update; this maintenance coupling is deliberate, not a semantic or security proof.
+The whole cursor digest still binds all other source bytes. Neither side labels nor
+observed counts determine policy. Future bases with either reviewed route work identically.
+
+Warmups and retained/partial pairs must preserve source, policy, native and environment
+identity. OFF and ON modes share Python anchors, but have distinct native binaries;
+route and diagnostic workers share the exact ON identity even for available partial data.
+Unknown versions, missing new fields/anchors, source-policy drift or contradictory native
+capability make evidence unavailable. Historical schema-2 workers without either new field
+retain the strict all-or-none `guarded_row` expectation; historical artifacts are not rewritten.
+An absent constructor timer is acceptable only for an expected-zero method with positive
+1,001-call native fetch proof. It never supplies a missing timing or successful measurement.
diff --git a/tests/test_017_fetch_bounded_text.py b/tests/test_017_fetch_bounded_text.py
index e300485d7..76d6bb862 100644
--- a/tests/test_017_fetch_bounded_text.py
+++ b/tests/test_017_fetch_bounded_text.py
@@ -1,17 +1,20 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.
-"""Bounded text payload fidelity, including row-wise routing beside a MAX column.
+"""Bounded text payload fidelity and LOB storage regression coverage.
-The MAX value here is only a routing control. Actual MAX text BOM/NUL fidelity
-belongs to the separate LOB decoder and is not covered by this regression.
+The bounded cases use MAX only as a routing control. The LOB cases preserve the
+existing streaming/decoder behavior; they do not redefine MAX text BOM/NUL fidelity.
"""
+import os
+import subprocess
import sys
+import textwrap
import pytest
-from mssql_python import SQL_CHAR, SQL_WCHAR
+from mssql_python import SQL_CHAR, SQL_WCHAR, ddbc_bindings as ddbc
@pytest.fixture(scope="module")
@@ -193,3 +196,380 @@ def test_bounded_nvarchar_batch_malformed_fallback(db_connection, raw, method):
assert _fetch_rows(cursor, method) == [(expected,)]
cursor.execute("SELECT CAST(N'recovered' AS nvarchar(64))")
assert cursor.fetchone()[0] == "recovered"
+
+
+@pytest.mark.parametrize(
+ "width", [10, 63, 64, 4000], ids=["small", "inline-edge", "heap-edge", "large"]
+)
+@pytest.mark.parametrize("method", ["fetchone", "fetchval"])
+def test_bounded_nvarchar_single_row_buffer_boundary(db_connection, width, method):
+ expected = [None, "", "A\0B", "\ufeff\ufffe", "x" * (width - 2) + "\U0001f642"]
+ values = ", ".join(
+ f"({index}, CAST("
+ + ("NULL" if value is None else f"0x{value.encode('utf-16le').hex()}")
+ + f" AS nvarchar({width})))"
+ for index, value in enumerate(expected)
+ )
+ with db_connection.cursor() as cursor:
+ cursor.execute(f"SELECT payload FROM (VALUES {values}) AS v(n, payload) ORDER BY n")
+ for value in expected:
+ row = getattr(cursor, method)()
+ actual = row[0] if method == "fetchone" else row
+ assert actual == value
+ assert actual is None or type(actual) is str
+ assert getattr(cursor, method)() is None
+ assert getattr(cursor, method)() is None
+ cursor.execute("SELECT CAST(N'reused' AS nvarchar(10))")
+ assert cursor.fetchone()[0] == "reused"
+
+
+@pytest.mark.parametrize("width", [63, 64], ids=["inline-edge", "heap-edge"])
+@pytest.mark.parametrize("method", ["fetchone", "fetchval"])
+@pytest.mark.parametrize("raw", ["00D8", "00DC"], ids=["unpaired-high", "unpaired-low"])
+def test_bounded_nvarchar_buffer_boundary_strict_error(db_connection, width, method, raw):
+ with db_connection.cursor() as cursor:
+ cursor.execute(f"SELECT CAST(0x{raw} AS nvarchar({width}))")
+ with pytest.raises(UnicodeDecodeError):
+ getattr(cursor, method)()
+ cursor.execute("SELECT CAST(N'recovered' AS nvarchar(10))")
+ row = getattr(cursor, method)()
+ assert (row[0] if method == "fetchone" else row) == "recovered"
+
+
+@pytest.mark.skipif(
+ sys.platform == "win32",
+ reason="Windows does not export the native driver function-pointer globals",
+)
+@pytest.mark.parametrize("width", [63, 64], ids=["inline-edge", "heap-edge"])
+def test_bounded_nvarchar_getdata_buffer_contract(conn_str, width):
+ import os
+ import subprocess
+ import textwrap
+
+ script = textwrap.dedent("""
+ import ctypes as c
+ import os
+ import sys
+ import mssql_python
+ from mssql_python import ddbc_bindings as ddbc
+
+ width = int(sys.argv[1])
+ pointer, short, length = c.c_void_p, c.c_short, c.c_ssize_t
+ get_type = c.CFUNCTYPE(short, pointer, c.c_ushort, short, pointer, length, pointer)
+ library = c.CDLL(ddbc.module.__file__)
+ slot = pointer.in_dll(library, "SQLGetData_ptr")
+ calls, errors = [], []
+
+ @get_type
+ def getdata(handle, column, ctype, buffer, capacity, indicator):
+ try:
+ assert column == 1 and ctype == -8 # SQL_C_WCHAR
+ assert capacity == (width + 1) * 2
+ assert c.string_at(buffer, capacity) == bytes(capacity)
+ calls.append(capacity)
+ return original(handle, column, ctype, buffer, capacity, indicator)
+ except BaseException as error:
+ errors.append(repr(error))
+ return -1
+
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as connection:
+ with connection.cursor() as cursor:
+ cursor.execute(f"SELECT CAST(REPLICATE(N'x', {width}) AS nvarchar({width}))")
+ saved = slot.value
+ assert saved
+ original = get_type(saved)
+ slot.value = c.cast(getdata, pointer).value
+ try:
+ assert cursor.fetchone()[0] == "x" * width
+ assert cursor.fetchone() is None
+ finally:
+ slot.value = saved
+ assert calls == [(width + 1) * 2] and not errors, (calls, errors)
+ cursor.execute("SELECT CAST(N'recovered' AS nvarchar(10))")
+ assert cursor.fetchval() == "recovered"
+ """)
+ result = subprocess.run(
+ [sys.executable, "-c", script, str(width)],
+ env={**os.environ, "DB_CONNECTION_STRING": conn_str},
+ capture_output=True,
+ text=True,
+ timeout=45,
+ check=False,
+ )
+ assert result.returncode == 0, result.stdout + result.stderr
+
+
+def _run_fetch_script(conn_str, script, *arguments, timeout=45):
+ result = subprocess.run(
+ [sys.executable, "-c", script, *arguments],
+ env={**os.environ, "DB_CONNECTION_STRING": conn_str},
+ capture_output=True,
+ text=True,
+ timeout=timeout,
+ check=False,
+ )
+ assert result.returncode == 0, result.stdout + result.stderr
+
+
+@pytest.mark.parametrize("size", (1, 2))
+def test_getdata_decode_error_keeps_preexisting_and_completed_cells(cursor, size):
+ cursor.execute("SELECT 7, CAST(0x00D8 AS NVARCHAR(1))")
+ assert ddbc.DDBCSQLFetch(cursor.hstmt) == 0
+ marker = object()
+ row = [marker]
+ if size == 1:
+ assert ddbc.DDBCSQLGetData(cursor.hstmt, size, row, "utf-16le", "utf-16le", -8) == 0
+ else:
+ with pytest.raises(UnicodeDecodeError):
+ ddbc.DDBCSQLGetData(cursor.hstmt, size, row, "utf-16le", "utf-16le", -8)
+ assert row == [marker, 7]
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+
+
+@pytest.mark.skipif(
+ sys.platform == "win32",
+ reason="Windows does not export the native driver function-pointer globals",
+)
+@pytest.mark.parametrize("payload", ("utf8", "empty", "null", "invalid", "warning"))
+def test_narrow_getdata_decoding_in_subprocess(conn_str, payload):
+ script = textwrap.dedent("""
+ import ctypes as c
+ import os
+ import sys
+ import mssql_python
+ from mssql_python import ddbc_bindings as ddbc
+
+ mode = sys.argv[1]
+ data = b"\\xff" if mode == "invalid" else (
+ b"" if mode == "empty" else "\\ufeffA\\x00\\U0001f600".encode("utf-8")
+ )
+ pointer, short, length = c.c_void_p, c.c_short, c.c_ssize_t
+ get_type = c.CFUNCTYPE(short, pointer, c.c_ushort, short, pointer, length, pointer)
+ diag_type = c.CFUNCTYPE(
+ short, short, pointer, short, pointer, pointer, pointer, short, pointer
+ )
+ library = c.CDLL(ddbc.module.__file__)
+ slot = pointer.in_dll(library, "SQLGetData_ptr")
+ diag_slot = pointer.in_dll(library, "SQLGetDiagRec_ptr")
+ calls, errors = [], []
+
+ @diag_type
+ def diagnostic(handle_type, handle, record, state, native, message, capacity, size):
+ try:
+ if record > 1:
+ return 100
+ text = "getdata warning".encode("utf-16le")
+ assert capacity > len(text) // 2
+ c.memmove(state, "01000\\0".encode("utf-16le"), 12)
+ c.memmove(message, text + b"\\0\\0", len(text) + 2)
+ c.cast(native, c.POINTER(c.c_int))[0] = 0
+ c.cast(size, c.POINTER(short))[0] = len(text) // 2
+ return 0
+ except BaseException as error:
+ errors.append(type(error).__name__)
+ return -1
+
+ @get_type
+ def getdata(handle, column, ctype, buffer, capacity, indicator):
+ try:
+ assert ctype == 1 # SQL_C_CHAR
+ calls.append(column)
+ result = original(handle, column, ctype, buffer, capacity, indicator)
+ assert result == 0 and capacity > len(data)
+ c.memmove(buffer, data + b"\\0", len(data) + 1)
+ c.cast(indicator, c.POINTER(length))[0] = -1 if mode == "null" else len(data)
+ return 1 if mode == "warning" else result
+ except BaseException as error:
+ errors.append(type(error).__name__)
+ return -1
+
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as connection:
+ with connection.cursor() as cursor:
+ cursor.execute("SELECT CAST('text' AS VARCHAR(30))")
+ saved, saved_diag = slot.value, diag_slot.value
+ assert saved and saved_diag
+ original = get_type(saved)
+ assert ddbc.DDBCSQLFetch(cursor.hstmt) == 0
+ marker = object()
+ row = [marker]
+ slot.value = c.cast(getdata, pointer).value
+ if mode == "warning":
+ diag_slot.value = c.cast(diagnostic, pointer).value
+ try:
+ ret = ddbc.DDBCSQLGetData(
+ cursor.hstmt, 1, row, "utf-8", "utf-16le", 1, cursor.messages
+ )
+ assert ret == (1 if mode == "warning" else 0)
+ finally:
+ slot.value, diag_slot.value = saved, saved_diag
+ expected = None if mode == "null" else (
+ data if mode == "invalid" else data.decode("utf-8")
+ )
+ assert row == [marker, expected]
+ assert type(row[1]) is type(expected)
+ assert calls == [1] and not errors, (calls, errors)
+ assert cursor.messages == (
+ [("[01000] (0)", "getdata warning")] if mode == "warning" else []
+ )
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+ """)
+ _run_fetch_script(conn_str, script, payload)
+
+
+@pytest.mark.parametrize("method", ["fetchone", "fetchmany", "fetchall"])
+@pytest.mark.parametrize("size", [None, 0, 1, 8191, 8192, 8193, 16384, 262144])
+def test_lob_binary_read_boundaries(db_connection, method, size):
+ expected = None if size is None else (bytes(range(256)) * ((size + 255) // 256))[:size]
+ literal = "NULL" if expected is None else "0x" + expected.hex()
+ with db_connection.cursor() as cursor:
+ cursor.execute(f"SELECT CAST({literal} AS varbinary(max))")
+ rows = _fetch_rows(cursor, method)
+ assert rows == [(expected,)]
+ assert type(rows[0][0]) is type(expected)
+ assert cursor.messages == []
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+
+
+@pytest.mark.parametrize("method", ["fetchone", "fetchmany", "fetchall"])
+@pytest.mark.parametrize("kind", ["nvarchar", "varchar-wide", "varchar-narrow"])
+def test_lob_text_read_boundaries(db_connection, method, kind):
+ # Keep embedded NUL inside a chunk, outside the existing trailing-NUL trimming.
+ suffix = "\U0001f642c\0af\u00e9-tail" if kind == "nvarchar" else "caf\u00e9\0tail"
+ expected = [None, ""] + [
+ "x" * boundary + suffix for boundary in (4093, 4094, 4095, 8187, 8190, 8191, 8192, 262144)
+ ]
+ original = db_connection.getdecoding(SQL_CHAR)
+ try:
+ narrow = kind == "varchar-narrow"
+ db_connection.setdecoding(
+ SQL_CHAR,
+ encoding="latin-1" if narrow else "utf-16le",
+ ctype=SQL_CHAR if narrow else SQL_WCHAR,
+ )
+ expressions = []
+ for index, value in enumerate(expected):
+ literal = "NULL" if value is None else "0x" + value.encode("utf-16le").hex()
+ expression = f"CAST({literal} AS nvarchar(max))"
+ if kind != "nvarchar":
+ expression = f"CAST({expression} COLLATE Latin1_General_100_BIN2 AS varchar(max))"
+ expressions.append(f"({index}, {expression})")
+ with db_connection.cursor() as cursor:
+ cursor.execute(
+ "SELECT payload FROM (VALUES "
+ + ", ".join(expressions)
+ + ") AS v(n, payload) ORDER BY n"
+ )
+ rows = _fetch_rows(cursor, method)
+ assert rows == [(value,) for value in expected]
+ assert all(row[0] is None or type(row[0]) is str for row in rows)
+ assert cursor.messages == []
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+ finally:
+ db_connection.setdecoding(SQL_CHAR, encoding=original["encoding"], ctype=original["ctype"])
+
+
+@pytest.mark.skipif(
+ sys.platform == "win32",
+ reason="Windows does not export the native driver function-pointer globals",
+)
+@pytest.mark.parametrize("kind", ["binary", "narrow", "wide"])
+@pytest.mark.parametrize(
+ "mode", ["known", "unknown", "negative", "split", "zero", "null", "error", "no-data"]
+)
+def test_lob_getdata_storage_contract(conn_str, kind, mode):
+ script = textwrap.dedent("""
+ import ctypes as c
+ import os
+ import sys
+ import mssql_python
+ from mssql_python import ddbc_bindings as ddbc
+
+ kind, mode = sys.argv[1:]
+ binary, wide = kind == "binary", kind == "wide"
+ ctype = -2 if binary else (-8 if wide else 1)
+ # VARCHAR(MAX) enters the LOB helper directly, including SQL_C_WCHAR.
+ # NVARCHAR(MAX) can first issue a bounded probe when its reported size is zero.
+ sqltype = "varbinary(max)" if binary else "varchar(max)"
+ codec = "utf-16le" if wide else "utf-8"
+ unit = 2 if wide else 1
+ terminator = b"" if binary else bytes(unit)
+ head_text = "A\\0B" + "x" * ((8192 - len(terminator)) // unit - 3)
+ tail_text = "Z\\0Y\\0\\0"
+ head, tail = head_text.encode(codec), tail_text.encode(codec)
+ expected = head + tail if binary else head_text + tail_text.rstrip("\\0")
+ if mode == "split" and not binary:
+ text = head_text[:-1] + "\\U0001f642" + tail_text
+ encoded = text.encode(codec)
+ head, tail = encoded[:len(head)], encoded[len(head):]
+ expected = text.rstrip("\\0")
+ elif mode == "zero":
+ expected = head if binary else head_text
+ elif mode == "null":
+ expected = None
+
+ pointer, short, length = c.c_void_p, c.c_short, c.c_ssize_t
+ get_type = c.CFUNCTYPE(short, pointer, c.c_ushort, short, pointer, length, pointer)
+ library = c.CDLL(ddbc.module.__file__)
+ slot = pointer.in_dll(library, "SQLGetData_ptr")
+ calls, errors = [], []
+
+ @get_type
+ def getdata(handle, column, actual_ctype, buffer, capacity, indicator):
+ try:
+ assert column == 1 and actual_ctype == ctype
+ assert capacity == 8192
+ assert c.string_at(buffer, capacity) == bytes(capacity)
+ calls.append(capacity)
+ assert len(calls) <= 2
+ first = len(calls) == 1
+ data = head if first else tail
+ c.memmove(buffer, data + terminator, len(data) + len(terminator))
+ reported = len(head) + len(tail) if first else len(tail)
+ if first and mode in ("unknown", "negative"):
+ reported = -4 if mode == "unknown" else -9
+ if not first:
+ if mode == "zero":
+ reported = 0
+ elif mode in ("null", "error", "no-data"):
+ reported = -1
+ c.cast(indicator, c.POINTER(length))[0] = reported
+ if not first and mode in ("error", "no-data"):
+ return -1 if mode == "error" else 100
+ return 1 if first else 0
+ except BaseException as error:
+ errors.append(repr(error))
+ return -1
+
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as connection:
+ with connection.cursor() as cursor:
+ cursor.execute(f"SELECT CAST(NULL AS {sqltype})")
+ assert ddbc.DDBCSQLFetch(cursor.hstmt) == 0
+ saved = slot.value
+ assert saved
+ marker = object()
+ row = [marker]
+ slot.value = c.cast(getdata, pointer).value
+ try:
+ try:
+ result = ddbc.DDBCSQLGetData(
+ cursor.hstmt, 1, row, "utf-8", "utf-16le", ctype, cursor.messages
+ )
+ except RuntimeError as error:
+ assert mode in ("error", "no-data"), str(error)
+ assert "Error fetching LOB" in str(error)
+ assert row == [marker]
+ else:
+ assert mode not in ("error", "no-data")
+ assert result == 0 and row == [marker, expected]
+ assert type(row[1]) is type(expected)
+ finally:
+ slot.value = saved
+ assert calls == [8192, 8192] and not errors, (calls, errors)
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+ """)
+ _run_fetch_script(conn_str, script, kind, mode)
diff --git a/tests/test_036_profiler_ci.py b/tests/test_036_profiler_ci.py
index 8f5e24eed..e71795159 100644
--- a/tests/test_036_profiler_ci.py
+++ b/tests/test_036_profiler_ci.py
@@ -1081,7 +1081,77 @@ def test_unix_profiler_step_does_not_put_database_password_on_command_line():
assert "Pwd=$DB_PASSWORD" in benchmark
-def test_build_check_rejects_foreign_provider_and_enabled_recording(tmp_path, monkeypatch):
+@pytest.fixture
+def release_cache(tmp_path):
+ cache = tmp_path / "mssql_python/pybind/build/CMakeCache.txt"
+ cache.parent.mkdir(parents=True)
+ cache.write_text(
+ "CMAKE_GENERATOR:INTERNAL=Unix Makefiles\n"
+ "CMAKE_CXX_FLAGS:STRING=\n"
+ "CMAKE_CXX_FLAGS_RELEASE:STRING=-O3 -DNDEBUG\n"
+ "CMAKE_BUILD_TYPE:STRING=Release\n",
+ encoding="utf-8",
+ )
+ return cache
+
+
+@pytest.mark.parametrize("profiling", [False, True])
+def test_benchmark_build_requests_and_records_release(
+ tmp_path, monkeypatch, release_cache, profiling
+):
+ def run(command, log, timeout, **options):
+ assert options["env"]["CMAKE_BUILD_TYPE"] == "Release"
+ assert options["env"]["ENABLE_PROFILING"] == str(int(profiling))
+ assert options["cwd"] == tmp_path / "mssql_python/pybind"
+ assert timeout == 37
+ log.write_text("build output\n", encoding="utf-8")
+
+ monkeypatch.setenv("CMAKE_BUILD_TYPE", "Debug")
+ monkeypatch.setattr(controller, "run_process", run)
+ log = tmp_path / "build.log"
+ controller.build(tmp_path, log, 37, profiling=profiling)
+ assert 'CMAKE_CXX_FLAGS_RELEASE": "-O3 -DNDEBUG"' in log.read_text(encoding="utf-8")
+
+
+@pytest.mark.parametrize("configuration", ["", "Debug", "Release", "Debug;Release"])
+def test_release_build_configuration(tmp_path, release_cache, configuration):
+ key = "CMAKE_CONFIGURATION_TYPES" if ";" in configuration else "CMAKE_BUILD_TYPE"
+ content = release_cache.read_text(encoding="utf-8")
+ release_cache.write_text(
+ content.replace("CMAKE_BUILD_TYPE:STRING=Release", f"{key}:STRING={configuration}"),
+ encoding="utf-8",
+ )
+ if ";" in configuration:
+ tag = f"py{sys.version_info.major}{sys.version_info.minor}"
+ windows_cache = release_cache.parent / "x64" / tag / release_cache.name
+ windows_cache.parent.mkdir(parents=True)
+ release_cache.replace(windows_cache)
+ release_cache = windows_cache
+ if "Release" not in configuration.split(";"):
+ with pytest.raises(ValueError, match="Release native build"):
+ controller.release_build_configuration(tmp_path)
+ else:
+ assert controller.release_build_configuration(tmp_path)[key] == configuration
+ release_cache.unlink()
+ with pytest.raises(FileNotFoundError):
+ controller.release_build_configuration(tmp_path)
+
+
+def test_release_cache_rejects_ambiguous_windows_builds(tmp_path, release_cache):
+ content = release_cache.read_text(encoding="utf-8")
+ tag = f"py{sys.version_info.major}{sys.version_info.minor}"
+ for arch in ("x64", "arm64"):
+ cache = release_cache.parent / arch / tag / release_cache.name
+ cache.parent.mkdir(parents=True)
+ cache.write_text(content, encoding="utf-8")
+ release_cache.unlink()
+ with pytest.raises(ValueError, match="Ambiguous"):
+ controller.release_build_configuration(tmp_path)
+
+
+def test_build_check_rejects_foreign_provider_and_enabled_recording(
+ tmp_path, monkeypatch, release_cache
+):
native = SimpleNamespace(
__file__=str(tmp_path / "binding.so"), profiling=SimpleNamespace(is_enabled=lambda: False)
)
@@ -1103,6 +1173,153 @@ def test_build_check_rejects_foreign_provider_and_enabled_recording(tmp_path, mo
controller.check_build(tmp_path, True)
+@pytest.mark.parametrize(
+ "failure", [None, "timeout", "identity", "deadline", "process-cleanup", "directory-cleanup"]
+)
+def test_ci_bundle_modes_identity_deadline_cleanup_and_headline(
+ report, tmp_path, monkeypatch, failure
+):
+ clock, measured, built = [0], [], []
+ directory = tmp_path / "build-roots"
+ directory.mkdir()
+ monkeypatch.setattr(controller.time, "monotonic", lambda: clock[0])
+ monkeypatch.setattr(controller, "resolve_revisions", lambda *a, **k: ("a" * 40, "b" * 40))
+ monkeypatch.setattr(controller, "git", lambda *a, **k: "b" * 40)
+ monkeypatch.setattr(controller.tempfile, "mkdtemp", lambda **k: str(directory))
+ monkeypatch.setattr(controller, "run_process", lambda *a, **k: None)
+ monkeypatch.setattr(
+ controller, "build", lambda path, *a, **k: built.append((path.name, k["profiling"]))
+ )
+ monkeypatch.setenv("BUILD_BUILDID", "42")
+ monkeypatch.setenv("SYSTEM_PULLREQUEST_SOURCECOMMITID", "c" * 40)
+ route = dict(version=1, methods=dict.fromkeys(("fetchone", "fetchmany", "fetchval"), False))
+
+ def source(path, revision):
+ return dict(source_commit=revision, python_cursor_sha256="d" * 64, row_route=route)
+
+ monkeypatch.setattr(controller, "python_source_identity", source)
+
+ def measure(path, output, scenarios, timeout, *, mode, revision, isolated):
+ measured.append(output.name)
+ assert isolated and scenarios is None
+ assert timeout == (controller.WORKER_TIMEOUT if mode == "diagnostic" else 30)
+ if output.name == "latency-base-1.json":
+ if failure == "timeout":
+ raise subprocess.TimeoutExpired("worker", timeout)
+ if failure == "process-cleanup":
+ raise controller.ProcessCleanupError("injected process cleanup failure")
+ if failure == "deadline":
+ clock[0] = controller.BENCHMARK_TIMEOUT + 1
+ side = "base" if revision == "a" * 40 else "candidate"
+ sample = copy.deepcopy(report["pairs"][0][side])
+ sample.update(
+ status="complete",
+ mode=mode,
+ provenance=dict(
+ **source(path, revision),
+ native_file=str(path / "native.so"),
+ native_sha256=("1" if mode == "latency" else "2") * 64,
+ native_profiling=mode != "latency",
+ guarded_row=True,
+ ),
+ )
+ if failure == "identity" and output.name == "latency-base-1.json":
+ sample["provenance"]["native_sha256"] = "3" * 64
+ if mode != "diagnostic":
+ sample["scenarios"] = {}
+ for name in reporting.FETCH_CASES:
+ shape, method = name.split("_")
+ timer = "ddbc::FetchMany_wrap" if method == "fetchmany" else "ddbc::FetchOne_wrap"
+ sample["scenarios"][name] = dict(
+ wall_ms=10,
+ work=f"Rows: 1000; shape: {shape}; API: {method}; EOF: 1",
+ cpp=(
+ {timer: dict(calls=1001, total_us=1001, min_us=1, max_us=1)}
+ if mode == "route"
+ else {}
+ ),
+ py={},
+ )
+ return sample
+
+ monkeypatch.setattr(controller, "measure", measure)
+ cleanup = MagicMock(wraps=controller.shutil.rmtree)
+ if failure == "directory-cleanup":
+ cleanup.side_effect = OSError("injected directory cleanup failure")
+ monkeypatch.setattr(controller.shutil, "rmtree", cleanup)
+ args = SimpleNamespace(
+ base=None,
+ candidate="HEAD",
+ output=tmp_path,
+ leg=report["leg"],
+ samples=5,
+ warmups=1,
+ reuse_candidate=True,
+ scenarios=None,
+ mode="diagnostic",
+ )
+ if failure:
+ with pytest.raises(RuntimeError):
+ controller.run_ci_report(args)
+ else:
+ controller.run_ci_report(args)
+ bundle = json.loads((tmp_path / "report.json").read_text(encoding="utf-8"))
+ modes, errors = reporting.ci_mode_reports(bundle)
+ if failure in ("process-cleanup", "directory-cleanup"):
+ assert bundle["cleanup_required"] == str(directory) and directory.exists()
+ assert cleanup.call_count == (failure == "directory-cleanup")
+ else:
+ cleanup.assert_called_once_with(str(directory))
+ assert not directory.exists()
+ if failure in ("process-cleanup", "deadline"):
+ assert measured == [
+ "latency-base-0.json",
+ "latency-candidate-0.json",
+ "latency-candidate-1.json",
+ "latency-base-1.json",
+ ]
+ assert all(mode["status"] == "incomplete" for mode in modes.values())
+ elif failure in ("timeout", "identity"):
+ assert modes["latency"]["status"] == "incomplete" and "latency" in errors
+ assert modes["route"]["status"] == modes["diagnostic"]["status"] == "complete"
+ elif failure is None:
+ assert not errors and all(len(mode["pairs"]) == 5 for mode in modes.values())
+ assert built == [("base-off", False), ("candidate-off", False), ("base-on", True)]
+ assert measured == [
+ f"{mode}-{side}-{index}.json"
+ for mode in ("latency", "route", "diagnostic")
+ for index in range(6)
+ for side in (("base", "candidate") if index % 2 == 0 else ("candidate", "base"))
+ ]
+ headline = reporting.render_ci_reports([bundle], "c" * 40, 42)
+ del bundle["fetch_measurements"]["latency"]
+ assert "latency" in reporting.ci_mode_reports(bundle)[1]
+ assert reporting.render_ci_reports([bundle], "c" * 40, 42) == headline
+ for field, value in (("source_commit", "e" * 40), ("mode", "latency")):
+ invalid = copy.deepcopy(bundle)
+ invalid["fetch_measurements"]["route"][field] = value
+ assert "route" in reporting.ci_mode_reports(invalid)[1]
+ assert reporting.render_ci_reports([invalid], "c" * 40, 42) == headline
+ for pair in bundle["fetch_measurements"]["route"]["pairs"]:
+ pair["candidate"]["provenance"]["native_sha256"] = "e" * 64
+ valid, errors = reporting.ci_mode_reports(bundle)
+ assert not valid and set(errors) == {"latency", "route", "diagnostic"}
+ assert "Shared ON" in errors["diagnostic"]
+
+
+@pytest.mark.parametrize("active,guarded", [(False, True), (True, False), (False, 1)])
+def test_default_row_route_distinguishes_capability_from_dispatch(active, guarded):
+ route = dict(version=1, methods=dict.fromkeys(("fetchone", "fetchmany", "fetchval"), active))
+ identity = dict(row_route=route, guarded_row=guarded)
+ if active or type(guarded) is not bool:
+ with pytest.raises(ValueError, match="binding contradicts"):
+ reporting.expected_constructors(identity, "fetchmany")
+ else:
+ assert all(
+ reporting.expected_constructors(identity, method) == 0 for method in route["methods"]
+ )
+
+
def test_head_moving_while_listing_comments_prevents_publish(monkeypatch):
calls = []
reads = 0
diff --git a/tests/test_fetch_settings_cache.py b/tests/test_fetch_settings_cache.py
index 78ad5229b..19d8c1a01 100644
--- a/tests/test_fetch_settings_cache.py
+++ b/tests/test_fetch_settings_cache.py
@@ -3,20 +3,26 @@
Licensed under the MIT license.
Regression and operation-count tests for fetch settings and diagnostic preservation.
-All integration queries are read-only and each test owns its connection.
+Read-only fetch regressions; native fault injection runs in isolated child processes.
"""
import datetime
+import decimal
+import gc
+import os
from pathlib import Path
import subprocess
import sys
+import textwrap
import uuid
import weakref
from unittest.mock import Mock, patch
import pytest
import mssql_python
+from mssql_python import ddbc_bindings as ddbc
from mssql_python.constants import ConstantsDDBC
+from mssql_python.cursor import Cursor
from mssql_python.row import Row
FETCH_METHODS = ("fetchone", "fetchmany", "fetchall")
@@ -188,18 +194,33 @@ def test_wchar_decoding_forwarded_to_live_fetch_bridge(connection, method, bridg
def test_decoding_cache_reuse_and_multiple_cursors(connection):
- with patch.object(connection, "getdecoding", wraps=connection.getdecoding) as reads:
+ with (
+ patch.object(connection, "getdecoding", wraps=connection.getdecoding) as reads,
+ patch.object(ddbc, "_FetchOptions", wraps=ddbc._FetchOptions) as options,
+ patch.object(ddbc, "_fetchone_with_options", wraps=ddbc._fetchone_with_options) as one,
+ patch.object(ddbc, "_fetchmany_with_options", wraps=ddbc._fetchmany_with_options) as many,
+ ):
with connection.cursor() as first, connection.cursor() as second:
assert reads.call_count == 4
+ assert options.call_count == 0
for cursor in (first, second):
+ cursor.execute("SELECT 42")
+ assert cursor.fetchall()[0][0] == 42
+ assert cursor._cached_fetch_options is None
cursor.execute("SELECT n FROM (VALUES (1), (2), (3)) AS v(n) ORDER BY n")
assert cursor.fetchone()[0] == 1
+ snapshot = cursor._cached_fetch_options
assert cursor.fetchmany(1)[0][0] == 2
assert cursor.fetchall()[0][0] == 3
+ assert one.call_args.args[-2] is many.call_args.args[-2] is snapshot
+ assert one.call_args.args[-1] is many.call_args.args[-1] is cursor.messages
+ assert cursor._cached_fetch_options is snapshot
assert reads.call_count == 4
+ assert options.call_count == one.call_count == many.call_count == 2
connection.setdecoding(mssql_python.SQL_CHAR, encoding="latin-1")
for index, cursor in enumerate((first, second), 1):
+ snapshot = cursor._cached_fetch_options
cursor.execute(
"SELECT CONVERT(VARCHAR(1), 0xE9) AS txt FROM (VALUES (1), (2), (3)) AS v(n)"
)
@@ -207,6 +228,35 @@ def test_decoding_cache_reuse_and_multiple_cursors(connection):
assert cursor.fetchmany(1)[0].txt == "\u00e9"
assert cursor.fetchall()[0].txt == "\u00e9"
assert reads.call_count == 4 + 2 * index
+ assert options.call_count == 2 + index
+ assert cursor._cached_fetch_options is not snapshot
+
+
+@pytest.mark.parametrize("method", ("fetchone", "fetchmany", "fetchval"))
+def test_native_fetch_options_refresh_failure_is_retried(connection, method):
+ with connection.cursor() as cursor:
+ cursor.execute("SELECT 42")
+ assert cursor.fetchone()[0] == 42
+ assert cursor._cached_fetch_options is not None
+ cursor.execute("SELECT CONVERT(VARCHAR(1), 0xE9) AS txt")
+ connection.setdecoding(mssql_python.SQL_CHAR, encoding="latin-1")
+
+ def fetch():
+ return cursor.fetchmany(1) if method == "fetchmany" else getattr(cursor, method)()
+
+ with patch.object(ddbc, "_FetchOptions", side_effect=MemoryError("options allocation")):
+ with pytest.raises(MemoryError, match="options allocation"):
+ fetch()
+ assert cursor._cached_decoding_generation == connection._decoding_generation
+ assert cursor._cached_fetch_options is None
+ assert cursor._next_row_index == 0
+ value = fetch()
+ actual = (
+ value[0][0] if method == "fetchmany" else value if method == "fetchval" else value[0]
+ )
+ assert actual == "\u00e9"
+ assert cursor._cached_fetch_options is not None
+ assert cursor.rowcount == 1
@pytest.mark.parametrize("method", FETCH_METHODS)
@@ -512,6 +562,273 @@ def test_direct_row_without_converters_is_zero_copy():
assert row._values is values
+@pytest.mark.parametrize("container", (list, tuple))
+def test_converter_leaf_inputs_and_copy(container):
+ text = "A\0\u00e9\U0001f600"
+ values = container(
+ [
+ text,
+ b"\0\xff",
+ 42,
+ decimal.Decimal("1.25"),
+ datetime.date(2026, 1, 2),
+ uuid.UUID(UUID_TEXT),
+ None,
+ ]
+ )
+ seen = []
+
+ def convert(value):
+ seen.append(value)
+ return value
+
+ with patch.object(
+ ddbc, "_apply_output_converters", wraps=ddbc._apply_output_converters
+ ) as leaf:
+ row = Row(values, {}, converter_map=[convert] * len(values), uuid_str_indices=(5,))
+ expected = [b"A\0\0\0\xe9\0\x3d\xd8\0\xde", *values[1:-1]]
+ assert seen == expected
+ assert [type(value) for value in seen] == [type(value) for value in expected]
+ assert row._values is not values
+ assert list(row) == [*expected[:5], UUID_TEXT, None]
+ assert isinstance(values[5], uuid.UUID)
+ assert leaf.call_count == int(container is list)
+ if container is list:
+ values[0] = "changed"
+ assert row[0] == expected[0]
+
+
+@pytest.mark.parametrize("container", (list, tuple))
+@pytest.mark.parametrize("map_container", (list, tuple))
+@pytest.mark.parametrize("map_length", (1, 4))
+def test_converter_leaf_map_length(container, map_container, map_length):
+ calls = []
+ converters = map_container([lambda value: calls.append(value)] * map_length)
+ row = Row(container([1, 2]), {}, converter_map=converters)
+ assert calls == ([1] if map_length == 1 else [1, 2])
+ assert list(row) == ([None, 2] if map_length == 1 else [None, None])
+
+
+@pytest.mark.parametrize("container", (list, tuple))
+def test_converter_leaf_dynamic_encode(container):
+ events = []
+ token = object()
+
+ class Text(str):
+ def encode(self, encoding):
+ events.append((str(self), encoding))
+ Text.encode = lambda self, encoding: token
+ return b"first"
+
+ class StringProxy:
+ __class__ = property(lambda self: str)
+
+ def encode(self, encoding):
+ events.append(("proxy", encoding))
+ return token
+
+ values = container([Text("one"), Text("two"), StringProxy()])
+ row = Row(values, {}, converter_map=[lambda value: value] * 3)
+ assert events == [("one", "utf-16-le"), ("proxy", "utf-16-le")]
+ assert row[0] == b"first"
+ assert row[1] is row[2] is token
+
+
+def test_converter_leaf_releases_owned_references():
+ references = []
+
+ class Value:
+ pass
+
+ class Text(str):
+ def encode(self, encoding):
+ value = Value()
+ references.append(weakref.ref(value))
+ return value
+
+ class Converter:
+ def __call__(self, value):
+ assert value is references[-1]()
+ result = Value()
+ references.append(weakref.ref(result))
+ return result
+
+ values, converters = [Text("text")], [Converter()]
+ references.extend([weakref.ref(values[0]), weakref.ref(converters[0])])
+ row = Row(values, {}, converter_map=converters)
+ assert references[2]() is None
+ assert references[3]() is row[0]
+ del values, converters, row
+ gc.collect()
+ assert all(reference() is None for reference in references)
+
+
+@pytest.mark.parametrize("container", (list, tuple))
+@pytest.mark.parametrize("stage", ("bool", "encode", "call"))
+@pytest.mark.parametrize("error_type", (ValueError, KeyboardInterrupt))
+def test_converter_leaf_exception_boundaries(container, stage, error_type):
+ class Text(str):
+ def encode(self, encoding):
+ if stage == "encode":
+ raise error_type("converter failure")
+ return super().encode(encoding)
+
+ class Converter:
+ def __bool__(self):
+ if stage == "bool":
+ raise error_type("converter failure")
+ return True
+
+ def __call__(self, value):
+ raise error_type("converter failure")
+
+ values = container([Text("text")])
+ if stage == "bool" or error_type is KeyboardInterrupt:
+ with pytest.raises(error_type, match="converter failure"):
+ Row(values, {}, converter_map=[Converter()])
+ else:
+ assert Row(values, {}, converter_map=[Converter()])[0] is values[0]
+
+
+@pytest.mark.parametrize("native", (False, True))
+@pytest.mark.parametrize("mutation", ("replace", "grow", "shrink"))
+def test_converter_leaf_live_lists(native, mutation):
+ events = []
+ values = [1, 2, None]
+
+ class Values(list):
+ pass
+
+ if not native:
+ values = Values(values)
+
+ class Converter:
+ def __bool__(self):
+ events.append("bool")
+ return True
+
+ def __call__(self, value):
+ events.append(value)
+ if value == 1:
+ if mutation == "replace":
+ values[1] = 20
+ converters[1] = lambda value: value + 100
+ elif mutation == "grow":
+ values.append(4)
+ converters.append(self)
+ else:
+ values.clear()
+ converters.clear()
+ return value + 10
+
+ converters = [Converter()] * 3
+ row = Row(values, {}, converter_map=converters)
+ assert (
+ list(row)
+ == {
+ "replace": [11, 120, None],
+ "grow": [11, 12, None],
+ "shrink": [11, 2, None],
+ }[mutation]
+ )
+ assert (
+ events
+ == {
+ "replace": ["bool", 1, "bool"],
+ "grow": ["bool", 1, "bool", 2, "bool", "bool", 4],
+ "shrink": ["bool", 1],
+ }[mutation]
+ )
+
+
+def test_converter_leaf_self_replacement_finalizer_order():
+ def run(container):
+ events = []
+
+ class Converter:
+ def __call__(self, value):
+ converters[0] = None
+ events.append(value)
+ return value
+
+ def __del__(self):
+ events.append("released")
+
+ converters = [Converter(), events.append, events.append]
+ row = Row(container([1, 2, 3]), {}, converter_map=converters)
+ return events, list(row)
+
+ assert run(list) == run(tuple)
+
+
+def test_converter_leaf_codec_lookup_in_subprocess():
+ script = textwrap.dedent("""
+ import codecs
+ import encodings
+ from mssql_python.row import Row
+
+ events = []
+ def encode(value, errors="strict"):
+ events.append((value, errors))
+ return b"custom", len(value)
+ def search(name):
+ if name == "utf_16_le":
+ return codecs.CodecInfo(name=name, encode=encode, decode=None)
+ codecs.unregister(encodings.search_function)
+ codecs.register(search)
+ codecs.register(encodings.search_function)
+ values = ["A\\0\\u00e9", "\\U0001f600"]
+ expected = [value.encode("utf-16-le") for value in values]
+ expected_events = events[:]
+ events.clear()
+ row = Row(values, {}, converter_map=[lambda value: value] * 2)
+ assert list(row) == expected
+ assert events == expected_events
+ """)
+ result = subprocess.run(
+ [sys.executable, "-c", script], capture_output=True, text=True, timeout=30
+ )
+ assert result.returncode == 0, result.stdout + result.stderr
+
+
+@pytest.mark.parametrize("method", FETCH_METHODS)
+def test_converter_leaf_registry_changes_during_rows(connection, method):
+ calls = []
+
+ def replacement(value):
+ return value + 100
+
+ def original(value):
+ calls.append(value)
+ connection.add_output_converter(mssql_python.SQL_INTEGER, replacement)
+ return value + 10
+
+ connection.add_output_converter(mssql_python.SQL_INTEGER, original)
+ with connection.cursor() as cursor:
+ cursor.execute("SELECT n AS a, n + 10 AS b FROM (VALUES (1), (2)) AS v(n) ORDER BY n")
+ rows = fetch_rows(cursor, method)
+ count = 1 if method == "fetchone" else 2
+ assert calls == [value for n in range(1, count + 1) for value in (n, n + 10)]
+ assert [list(row) for row in rows] == [[n + 10, n + 20] for n in range(1, count + 1)]
+ if method == "fetchone":
+ assert list(cursor.fetchone()) == [102, 112]
+ cursor.execute("SELECT 3 AS renamed")
+ assert cursor.fetchval() == 103
+ assert list(rows[0]) == [11, 21]
+ assert dict(rows[0]._mapping) == {"a": 11, "b": 21}
+
+
+@pytest.mark.parametrize("method", ("fetchone", "fetchval"))
+def test_converter_leaf_native_error_precedes_callbacks(connection, method):
+ calls = []
+ connection.add_output_converter(mssql_python.SQL_INTEGER, lambda value: calls.append(value))
+ with connection.cursor() as cursor:
+ cursor.execute("SELECT 1 AS a, CAST(0x00D8 AS NVARCHAR(10)) AS invalid_utf16")
+ with pytest.raises(UnicodeDecodeError):
+ getattr(cursor, method)()
+ assert calls == []
+
+
def test_decoding_cache_refresh_failure_is_retried(connection):
with connection.cursor() as cursor:
cursor.execute("SELECT CONVERT(VARCHAR(1), 0xE9) AS txt")
@@ -617,7 +934,7 @@ def assert_row_instance_capabilities(row):
return reference
-@pytest.mark.parametrize("construction", ("direct", "python_fast", "native"))
+@pytest.mark.parametrize("construction", ("direct", "python_fast", "native", "native_single"))
def test_constructed_row_preserves_instance_capabilities(construction):
values = [42]
column_map = {"number": 0}
@@ -625,6 +942,8 @@ def test_constructed_row_preserves_instance_capabilities(construction):
row = Row(values, column_map)
elif construction == "python_fast":
row = Row._fast_create(values, column_map, None)
+ elif construction == "native_single":
+ row = mssql_python.ddbc_bindings.construct_row(values, Row, column_map, None)
else:
row = mssql_python.ddbc_bindings.construct_rows([values], Row, column_map, None)[0]
assert row._values is values
@@ -643,8 +962,12 @@ def test_fetched_row_preserves_instance_capabilities(connection, method):
assert reference() is None
-@pytest.mark.parametrize("size", (0, 1, 3))
-def test_construct_rows_repeated_calls_release_references(size):
+@pytest.mark.parametrize(
+ ("size", "single"),
+ ((0, False), (1, False), (3, False), (1, True)),
+ ids=("0", "1", "3", "native-single"),
+)
+def test_construct_rows_repeated_calls_release_references(size, single):
values = [[index] for index in range(size)]
column_map = {"Number": 0}
column_map_lower = {"number": 0}
@@ -653,9 +976,16 @@ def test_construct_rows_repeated_calls_release_references(size):
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
- )
+ if single:
+ rows = [
+ mssql_python.ddbc_bindings.construct_row(
+ values[0], Row, column_map, cursor, column_map_lower, column_names
+ )
+ ]
+ else:
+ 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)
@@ -667,6 +997,14 @@ def test_construct_rows_repeated_calls_release_references(size):
def test_construct_rows_releases_partial_batch_on_attribute_error():
+ _assert_construct_row_attribute_failure(single=False)
+
+
+def test_construct_single_row_releases_references_on_attribute_error():
+ _assert_construct_row_attribute_failure(single=True)
+
+
+def _assert_construct_row_attribute_failure(single):
class FailingRow(Row):
__slots__ = ()
@@ -687,9 +1025,14 @@ def _column_names(self, names):
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
- )
+ if single:
+ mssql_python.ddbc_bindings.construct_row(
+ values[-1], FailingRow, column_map, cursor, None, column_names
+ )
+ else:
+ mssql_python.ddbc_bindings.construct_rows(
+ values, FailingRow, column_map, cursor, None, column_names
+ )
assert [sys.getrefcount(value) for value in tracked] == references
@@ -1111,3 +1454,1031 @@ 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
+
+
+def _run_fetch_script(conn_str, script, *arguments, timeout=45):
+ result = subprocess.run(
+ [sys.executable, "-c", script, *arguments],
+ env={**os.environ, "DB_CONNECTION_STRING": conn_str},
+ capture_output=True,
+ text=True,
+ timeout=timeout,
+ check=False,
+ )
+ assert result.returncode == 0, result.stdout + result.stderr
+
+
+@pytest.mark.skipif(
+ sys.platform == "win32",
+ reason="Windows does not export the native driver function-pointer globals",
+)
+@pytest.mark.parametrize("method", ["fetchone", "fetchmany", "fetchval"])
+def test_first_numeric_getdata_error_stops_before_second_column(conn_str, method):
+ script = textwrap.dedent("""
+ import ctypes as c
+ import os
+ import sys
+ from types import SimpleNamespace
+ from unittest.mock import patch
+ import mssql_python
+ from mssql_python import ddbc_bindings as ddbc
+
+ pointer, short, length = c.c_void_p, c.c_short, c.c_ssize_t
+ get_type = c.CFUNCTYPE(short, pointer, c.c_ushort, short, pointer, length, pointer)
+ library = c.CDLL(ddbc.module.__file__)
+ slot = pointer.in_dll(library, "SQLGetData_ptr")
+ calls = []
+
+ @get_type
+ def getdata(handle, column, ctype, buffer, capacity, indicator):
+ calls.append(column)
+ if column == 1:
+ return -1
+ return original(handle, column, ctype, buffer, capacity, indicator)
+
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as connection:
+ with connection.cursor() as cursor:
+ cursor.execute("SELECT CAST(7 AS INT), CAST(8 AS INT)")
+ position = cursor._next_row_index
+ saved = slot.value
+ assert saved
+ original = get_type(saved)
+ diagnostic = SimpleNamespace(sqlState="HY000", ddbcErrorMsg="first column failed")
+ with patch.object(ddbc, "DDBCSQLCheckError", return_value=diagnostic):
+ slot.value = c.cast(getdata, pointer).value
+ try:
+ method = sys.argv[1]
+ try:
+ getattr(cursor, method)(*([1] if method == "fetchmany" else []))
+ except mssql_python.DatabaseError as error:
+ assert "first column failed" in str(error)
+ else:
+ raise AssertionError("first-column SQL_ERROR was masked")
+ finally:
+ slot.value = saved
+ assert calls == [1], calls
+ assert cursor._next_row_index == position
+ cursor.execute("SELECT 42, 43")
+ assert tuple(cursor.fetchone()) == (42, 43)
+ """)
+ _run_fetch_script(conn_str, script, method)
+
+
+@pytest.mark.parametrize(
+ ("sql_type", "literal", "expected"),
+ (
+ ("INT", "-2147483648", -2147483648),
+ ("INT", "2147483647", 2147483647),
+ ("SMALLINT", "-32768", -32768),
+ ("SMALLINT", "32767", 32767),
+ ("BIGINT", "-9223372036854775808", -9223372036854775808),
+ ("BIGINT", "9223372036854775807", 9223372036854775807),
+ ("TINYINT", "255", 255),
+ ("BIT", "0", False),
+ ("REAL", "-1.25", -1.25),
+ ("FLOAT", "1.7976931348623157E308", 1.7976931348623157e308),
+ ("FLOAT", "-2.2250738585072014E-308", -2.2250738585072014e-308),
+ ("FLOAT", "-0.0", 0.0),
+ ),
+ ids=(
+ "int_min",
+ "int_max",
+ "smallint_min",
+ "smallint_max",
+ "bigint_min",
+ "bigint_max",
+ "tinyint_max",
+ "bit_zero",
+ "real_negative",
+ "float_max",
+ "float_tiny",
+ "float_zero",
+ ),
+)
+def test_single_numeric_row_bound_path_parity(cursor, sql_type, literal, expected):
+ query = f"SELECT CAST({literal} AS {sql_type}) AS a, CAST(NULL AS {sql_type}) AS b"
+ cursor.execute(query)
+ bound = cursor.fetchmany(2)[0]
+ cursor.execute(query)
+ single = cursor.fetchmany(1)[0]
+ assert tuple(single) == tuple(bound) == (expected, None)
+ assert type(single[0]) is type(bound[0]) is type(expected)
+
+
+def test_single_row_converters_run_in_column_order(cursor):
+ events = []
+
+ def convert(value):
+ events.append(value)
+ if value == 2:
+ raise ValueError("keep original second column")
+ return value + 10
+
+ cursor.connection.add_output_converter(mssql_python.SQL_INTEGER, convert)
+ try:
+ cursor.execute("SELECT 1 AS a, 2 AS b, CAST(NULL AS INT) AS c")
+ assert cursor.fetchval() == 11
+ assert events == [1, 2]
+ events.clear()
+ cursor.execute("SELECT 1 AS a, 2 AS b, CAST(NULL AS INT) AS c")
+ assert tuple(cursor.fetchmany(1)[0]) == (11, 2, None)
+ assert events == [1, 2]
+ finally:
+ cursor.connection.remove_output_converter(mssql_python.SQL_INTEGER)
+
+
+def test_single_row_wrapper_does_not_enter_batch_factory(cursor):
+ from mssql_python import ddbc_bindings
+
+ cursor.execute(
+ "SELECT n AS a, CAST(N'text' AS NVARCHAR(10)) AS b "
+ "FROM (VALUES (1),(2),(3),(4)) AS v(n) ORDER BY n"
+ )
+ with (
+ patch.object(ddbc_bindings, "construct_rows", wraps=ddbc_bindings.construct_rows) as batch,
+ patch.object(
+ ddbc_bindings, "DDBCSQLFetchRow", wraps=ddbc_bindings.DDBCSQLFetchRow
+ ) as fused,
+ ):
+ retained = cursor.fetchmany(1)[0]
+ assert tuple(retained) == (1, "text")
+ batch.assert_not_called()
+ assert [row[0] for row in cursor.fetchmany(2)] == [2, 3]
+ batch.assert_called_once()
+ assert cursor.fetchmany(2)[0][0] == 4
+ assert batch.call_count == 2
+ assert tuple(retained) == (1, "text")
+ fused.assert_not_called()
+
+
+@pytest.mark.parametrize("override_fast_create", (False, True))
+def test_fetchmany_preserves_substituted_row_class(cursor, override_fast_create):
+ import importlib
+ from mssql_python.row import Row
+
+ class DerivedRow(Row):
+ pass
+
+ def forbidden(*args):
+ raise AssertionError("batch wrapping must not call a substituted Row's factory")
+
+ if override_fast_create:
+ DerivedRow._fast_create = staticmethod(forbidden)
+ cursor_module = importlib.import_module("mssql_python.cursor")
+ cursor.execute("SELECT 1 AS a")
+ with patch.object(cursor_module, "Row", DerivedRow):
+ row = cursor.fetchmany(1)[0]
+ assert type(row) is DerivedRow
+ assert row.a == 1
+
+
+def test_fetchmany_preserves_replaced_fast_factory(cursor):
+ from mssql_python.row import Row
+
+ cursor.execute("SELECT 1 AS a")
+ with patch.object(Row, "_fast_create", side_effect=AssertionError("must use batch factory")):
+ assert cursor.fetchmany(1)[0].a == 1
+
+
+@pytest.mark.parametrize("via_arraysize", (False, True))
+@pytest.mark.parametrize("raises", (False, True))
+def test_fetchmany_size_subclass_equality_is_not_called(cursor, via_arraysize, raises):
+ from mssql_python import ddbc_bindings
+
+ calls = []
+
+ class Size(int):
+ def __eq__(self, other):
+ calls.append(other)
+ if raises:
+ raise AssertionError("size equality must not run after native fetch")
+ return super().__eq__(other)
+
+ cursor.execute("SELECT 1 AS a UNION ALL SELECT 2")
+ with patch.object(ddbc_bindings, "construct_rows", wraps=ddbc_bindings.construct_rows) as batch:
+ if via_arraysize:
+ cursor.arraysize = Size(1)
+ result = cursor.fetchmany()
+ else:
+ result = cursor.fetchmany(Size(1))
+ assert result[0][0] == 1
+ batch.assert_called_once()
+ assert calls == []
+ assert cursor.fetchone()[0] == 2
+
+
+def test_real_subclass_and_instance_fetchone_overrides(cursor):
+ calls = []
+
+ class DerivedCursor(Cursor):
+ def fetchone(self):
+ calls.append("derived")
+ return (71, 72)
+
+ with DerivedCursor(cursor.connection) as derived:
+ derived.execute("SELECT 1 AS a UNION ALL SELECT 2")
+ assert derived.fetchval() == 71
+ assert next(derived) == (71, 72)
+ assert calls == ["derived", "derived"]
+ # fetchmany must not acquire fetchone's Python override semantics.
+ assert derived.fetchmany(1)[0][0] == 1
+ assert calls == ["derived", "derived"]
+ with patch.object(derived, "fetchone", return_value=(81, 82)) as override:
+ assert derived.fetchval() == 81
+ assert next(derived) == (81, 82)
+ assert derived.fetchmany(1)[0][0] == 2
+ assert override.call_count == 2
+
+
+@pytest.mark.skipif(
+ sys.platform == "win32",
+ reason="Windows does not export the native driver function-pointer globals",
+)
+@pytest.mark.parametrize("first_fetch", ("fetchone", "fetchmany"))
+def test_count_generation_change_with_unbound_marker(conn_str, first_fetch):
+ """A count obtained across invalidation cannot authorize either cache."""
+ script = textwrap.dedent("""
+ import ctypes as c
+ import os
+ import sys
+ import mssql_python
+ from mssql_python import ddbc_bindings as ddbc
+
+ library = c.CDLL(ddbc.module.__file__)
+ pointer = c.c_void_p
+ count_type = c.CFUNCTYPE(c.c_short, pointer, c.POINTER(c.c_short))
+ unbind_type = c.CFUNCTYPE(c.c_short, pointer, c.c_ushort)
+ count_slot = pointer.in_dll(library, "SQLNumResultCols_ptr")
+ unbind_slot = pointer.in_dll(library, "SQLFreeStmt_ptr")
+ counts, unbinds, invalidations, callback_errors = [], [], [], []
+
+ @count_type
+ def counted(handle, value):
+ try:
+ counts.append(handle)
+ ret = original_count(handle, value)
+ if not invalidations:
+ invalidations.append(True)
+ # Change only the cache generation, not the result shape.
+ assert ddbc.DDBCSQLSetStmtAttr(cursor.hstmt, 0, 0) == 0
+ return ret
+ except BaseException as error:
+ callback_errors.append(type(error).__name__)
+ return -1
+
+ @unbind_type
+ def unbound(handle, option):
+ try:
+ if option == 2:
+ unbinds.append(handle)
+ return original_unbind(handle, option)
+ except BaseException as error:
+ callback_errors.append(type(error).__name__)
+ return -1
+
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as connection:
+ with connection.cursor() as cursor:
+ cursor.execute(
+ "SELECT n AS a, n + 10 AS b FROM (VALUES (1),(2),(3),(4)) v(n) ORDER BY n"
+ )
+ saved_count, saved_unbind = count_slot.value, unbind_slot.value
+ assert saved_count and saved_unbind
+ original_count = count_type(saved_count)
+ original_unbind = unbind_type(saved_unbind)
+ count_slot.value = c.cast(counted, pointer).value
+ unbind_slot.value = c.cast(unbound, pointer).value
+ try:
+ first = (
+ cursor.fetchmany(1)[0] if sys.argv[1] == "fetchmany"
+ else cursor.fetchone()
+ )
+ assert tuple(first) == (1, 11)
+ assert invalidations == [True]
+ before_count, before_unbind = len(counts), len(unbinds)
+ assert tuple(cursor.fetchone()) == (2, 12)
+ assert len(counts) == before_count + 1
+ assert len(unbinds) == before_unbind + 1
+ before_count, before_unbind = len(counts), len(unbinds)
+ assert cursor.fetchval() == 3
+ assert tuple(next(cursor)) == (4, 14)
+ assert cursor.fetchone() is None
+ assert len(counts) == before_count
+ assert len(unbinds) == before_unbind
+ # Direct count calls remain uncached even after a warm fetch.
+ assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2
+ assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2
+ assert len(counts) == before_count + 2
+ assert tuple(first) == (1, 11)
+ assert not callback_errors, callback_errors
+ assert not cursor.messages
+ finally:
+ count_slot.value, unbind_slot.value = saved_count, saved_unbind
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+ """)
+ _run_fetch_script(conn_str, script, first_fetch)
+
+
+@pytest.mark.parametrize("mutation", ("new", "setattr", "code", "abstract", "descriptor"))
+def test_fast_row_inplace_customization(monkeypatch, mutation):
+ import weakref
+ from mssql_python import ddbc_bindings
+ from mssql_python.row import Row
+
+ factory = Row._fast_create
+ events = []
+
+ def custom_new(cls):
+ events.append("new")
+ return object.__new__(cls)
+
+ def custom_setattr(self, name, value):
+ events.append(name)
+ object.__setattr__(self, name, value)
+
+ def custom_factory(values, column_map, cursor, column_map_lower=None, column_names=None):
+ raise RuntimeError("customized factory code")
+
+ def fail_descriptor(self, value):
+ events.append(weakref.ref(self))
+ raise RuntimeError("customized descriptor")
+
+ if mutation == "new":
+ monkeypatch.setattr(Row, "__new__", staticmethod(custom_new))
+ elif mutation == "setattr":
+ monkeypatch.setattr(Row, "__setattr__", custom_setattr)
+ elif mutation == "code":
+ monkeypatch.setattr(factory, "__code__", custom_factory.__code__)
+ elif mutation == "abstract":
+ monkeypatch.setattr(Row, "__abstractmethods__", frozenset({"required"}), raising=False)
+ else:
+ monkeypatch.setattr(Row, "_column_names", property(fset=fail_descriptor))
+
+ assert Row._fast_create is factory
+ with patch.object(ddbc_bindings, "construct_row", wraps=ddbc_bindings.construct_row) as native:
+ if mutation in ("code", "descriptor"):
+ with pytest.raises(RuntimeError, match="customized"):
+ factory([42], {"number": 0}, None)
+ elif mutation == "abstract":
+ with pytest.raises(TypeError, match="abstract"):
+ factory([42], {"number": 0}, None)
+ else:
+ assert factory([42], {"number": 0}, None).number == 42
+ native.assert_not_called()
+ if mutation == "new":
+ assert events == ["new"]
+ elif mutation == "setattr":
+ assert events == ["_values", "_column_map", "_cursor", "_column_map_lower", "_column_names"]
+ elif mutation == "descriptor":
+ assert len(events) == 1 and events[0]() is None
+
+
+@pytest.mark.parametrize("method", ("fetchone", "fetchmany", "fetchval"))
+@pytest.mark.parametrize("failure", ("maps", "factory"))
+def test_single_row_construction_failure_keeps_fetch_position(cursor, method, failure):
+ from mssql_python import ddbc_bindings
+
+ cursor.execute("SELECT n AS number FROM (VALUES (1), (2)) AS v(n) ORDER BY n")
+ from mssql_python.row import Row
+
+ bridge_name = "DDBCSQLFetchMany" if method == "fetchmany" else "DDBCSQLFetchOne"
+ bridge = getattr(ddbc_bindings, bridge_name)
+
+ def fetch():
+ value = cursor.fetchmany(1) if method == "fetchmany" else getattr(cursor, method)()
+ return value[0][0] if method == "fetchmany" else value if method == "fetchval" else value[0]
+
+ def fail(*args):
+ assert cursor.rowcount == 1
+ assert cursor.rownumber == 0
+ assert cursor._next_row_index == 1
+ raise RuntimeError("injected post-fetch failure")
+
+ failed_stage = Mock(side_effect=fail)
+ failure_patch = (
+ patch.object(cursor, "_get_column_and_converter_maps", failed_stage)
+ if failure == "maps"
+ else patch.object(Row, "_column_names", property(fset=failed_stage))
+ )
+ with patch.object(ddbc_bindings, bridge_name, wraps=bridge) as native_fetch, failure_patch:
+ with pytest.raises(RuntimeError, match="injected post-fetch failure"):
+ fetch()
+ native_fetch.assert_called_once()
+ failed_stage.assert_called_once()
+ assert fetch() == 2
+ assert cursor.rowcount == 2
+ assert cursor.rownumber == 1
+
+
+@pytest.mark.parametrize("target", ("__new__", "__setattr__"))
+def test_native_row_guard_does_not_invoke_descriptors(monkeypatch, target):
+ from mssql_python.cursor import _native_row_eligible
+ from mssql_python.row import Row
+
+ events = []
+
+ class Descriptor:
+ def __get__(self, instance, owner):
+ events.append("lookup")
+ if target == "__new__":
+ if len(events) > 1:
+ raise RuntimeError("duplicate allocator lookup")
+
+ def allocate(cls):
+ events.append("allocate")
+ return object.__new__(cls)
+
+ return allocate
+ return lambda name, value: object.__setattr__(instance, name, value)
+
+ monkeypatch.setattr(Row, target, Descriptor())
+ assert not _native_row_eligible(Row)
+ assert events == []
+ row = Row._fast_create([42], {"number": 0}, None)
+ assert row.number == 42
+ assert events == (["lookup", "allocate"] if target == "__new__" else ["lookup"] * 5)
+
+
+def test_native_row_guard_does_not_invoke_metaclass_hooks():
+ from mssql_python.cursor import _native_row_eligible
+ from mssql_python.row import Row
+
+ class Meta(type):
+ def __getattribute__(cls, name):
+ raise AssertionError("guard must not inspect substituted class through metaclass")
+
+ class CustomRow(Row, metaclass=Meta):
+ pass
+
+ assert not _native_row_eligible(CustomRow)
+
+
+@pytest.mark.parametrize("method", ("fetchone", "fetchmany", "fetchval"))
+@pytest.mark.parametrize("when", ("maps", "final_argument"))
+def test_single_row_fusion_handles_post_fetch_factory_change(cursor, method, when):
+ from mssql_python import ddbc_bindings
+ from mssql_python.row import Row
+
+ original_maps = cursor._get_column_and_converter_maps
+ factory = Row._fast_create
+
+ def replacement(values, column_map, cursor, column_map_lower=None, column_names=None):
+ raise RuntimeError("factory changed after native advancement")
+
+ def maps():
+ factory.__code__ = replacement.__code__
+ return original_maps()
+
+ def names(self):
+ factory.__code__ = replacement.__code__
+ return self.__dict__["_cached_result_columns"]
+
+ cursor.execute("SELECT n AS number FROM (VALUES (1), (2)) AS v(n) ORDER BY n")
+ cursor._get_column_and_converter_maps()
+ change = (
+ patch.object(cursor, "_get_column_and_converter_maps", side_effect=maps)
+ if when == "maps"
+ else patch.object(type(cursor), "_cached_result_columns", property(names), create=True)
+ )
+ old_code = factory.__code__
+ try:
+ with (
+ change,
+ patch.object(
+ ddbc_bindings, "DDBCSQLFetchRow", wraps=ddbc_bindings.DDBCSQLFetchRow
+ ) as fused,
+ ):
+ with pytest.raises(RuntimeError, match="factory changed after native advancement"):
+ cursor.fetchmany(1) if method == "fetchmany" else getattr(cursor, method)()
+ fused.assert_not_called()
+ assert Row._fast_create is factory
+ assert cursor.rowcount == 1 and cursor.rownumber == 0
+ finally:
+ factory.__code__ = old_code
+ assert cursor.fetchone().number == 2
+
+
+@pytest.mark.parametrize("method", ("fetchone", "fetchmany", "fetchval"))
+def test_single_row_late_allocator_preserves_factory_global_lookup(cursor, method):
+ from mssql_python import ddbc_bindings
+ from mssql_python.row import Row
+
+ original_row = Row
+
+ class ChangedRow(original_row):
+ pass
+
+ events = []
+ factory = original_row._fast_create
+ assert "__new__" not in vars(original_row)
+
+ class Allocator:
+ def __get__(self, instance, owner):
+ events.append("new_lookup")
+ factory.__globals__["Row"] = ChangedRow
+
+ def allocate(cls):
+ events.append("allocate:" + cls.__name__)
+ return object.__new__(cls)
+
+ return allocate
+
+ def names(self):
+ events.append("names_lookup")
+ original_row.__new__ = Allocator()
+ return self.__dict__["_cached_result_columns"]
+
+ cursor.execute("SELECT n AS number FROM (VALUES (1), (2)) AS v(n) ORDER BY n")
+ cursor._get_column_and_converter_maps()
+ try:
+ with (
+ patch.object(type(cursor), "_cached_result_columns", property(names), create=True),
+ patch.object(
+ ddbc_bindings, "DDBCSQLFetchRow", wraps=ddbc_bindings.DDBCSQLFetchRow
+ ) as fused,
+ ):
+ result = cursor.fetchmany(1) if method == "fetchmany" else getattr(cursor, method)()
+ fused.assert_not_called()
+ if method == "fetchval":
+ assert result == 1
+ else:
+ row = result[0] if method == "fetchmany" else result
+ assert type(row) is ChangedRow
+ assert row.number == 1
+ assert events == ["names_lookup", "new_lookup", "allocate:ChangedRow"]
+ assert cursor.rowcount == 1 and cursor.rownumber == 0
+ finally:
+ factory.__globals__["Row"] = original_row
+ if "__new__" in vars(original_row):
+ delattr(original_row, "__new__")
+ assert cursor.fetchone().number == 2
+
+
+def test_getdata_appends_to_existing_list_without_python_append(cursor):
+ class Destination(list):
+ def append(self, value):
+ raise AssertionError("native append must not dispatch to list overrides")
+
+ class Marker:
+ pass
+
+ marker = Marker()
+ reference = weakref.ref(marker)
+ destination = Destination([marker])
+ alias = destination
+ cursor.execute(
+ "SELECT CAST(-2147483648 AS INT), CAST(-32768 AS SMALLINT), "
+ "CAST(-9223372036854775808 AS BIGINT), CAST(255 AS TINYINT), "
+ "CAST(1 AS BIT), CAST(-1.25 AS REAL), CAST(1.5 AS FLOAT), "
+ "CAST(NULL AS INT), CAST(NULL AS SMALLINT), CAST(NULL AS BIGINT), "
+ "CAST(NULL AS TINYINT), CAST(NULL AS BIT), CAST(NULL AS REAL), "
+ "CAST(NULL AS FLOAT), CAST(12.50 AS DECIMAL(5,2)), "
+ "CAST('2024-02-29' AS DATE), CAST(0x0001FF AS VARBINARY(3))"
+ )
+ assert ddbc.DDBCSQLFetch(cursor.hstmt) == 0
+ assert ddbc.DDBCSQLGetData(cursor.hstmt, 17, destination, "utf-16le", "utf-16le", -8) == 0
+ expected = [
+ -2147483648,
+ -32768,
+ -9223372036854775808,
+ 255,
+ True,
+ -1.25,
+ 1.5,
+ *([None] * 7),
+ decimal.Decimal("12.50"),
+ datetime.date(2024, 2, 29),
+ b"\x00\x01\xff",
+ ]
+ assert destination is alias and destination[0] is marker
+ assert destination[1:] == expected
+ assert list(map(type, destination[1:])) == list(map(type, expected))
+ del marker, destination, alias
+ gc.collect()
+ assert reference() is None
+ cursor.execute("SELECT 42")
+ assert cursor.fetchval() == 42
+
+
+@pytest.mark.skipif(
+ sys.platform == "win32",
+ reason="Windows does not export the native driver function-pointer globals",
+)
+@pytest.mark.parametrize(
+ "mode",
+ (
+ "numeric",
+ "mixed",
+ "lob",
+ "prefix",
+ "generation",
+ "count_error",
+ "nextset",
+ "warning",
+ "decode_error",
+ ),
+)
+def test_fetchmany_one_full_count_calls_in_subprocess(conn_str, mode):
+ script = textwrap.dedent("""
+ import ctypes as c
+ import os
+ import sys
+ import mssql_python
+ from mssql_python import ddbc_bindings as ddbc
+
+ mode = sys.argv[1]
+ pointer, short, ushort = c.c_void_p, c.c_short, c.c_ushort
+ count_type = c.CFUNCTYPE(short, pointer, c.POINTER(short))
+ diag_type = c.CFUNCTYPE(
+ short, short, pointer, short, pointer, pointer, pointer, short, pointer
+ )
+ library = c.CDLL(ddbc.module.__file__)
+ count_slot = pointer.in_dll(library, "SQLNumResultCols_ptr")
+ diag_slot = pointer.in_dll(library, "SQLGetDiagRec_ptr")
+ calls, errors, injected = [], [], []
+ warning_pending = False
+
+ @count_type
+ def counted(handle, count):
+ global warning_pending
+ try:
+ calls.append(handle)
+ if mode == "count_error" and not injected:
+ injected.append(True)
+ return -1
+ result = original_count(handle, count)
+ if mode == "generation" and not injected:
+ injected.append(True)
+ assert ddbc.DDBCSQLSetStmtAttr(cursor.hstmt, 0, 0) == 0
+ if mode == "warning" and not injected:
+ injected.append(True)
+ warning_pending = True
+ return 1
+ return result
+ except BaseException as error:
+ errors.append(type(error).__name__)
+ return -1
+
+ @diag_type
+ def diagnostic(handle_type, handle, record, state, native, message, capacity, size):
+ global warning_pending
+ try:
+ if not warning_pending:
+ return original_diag(
+ handle_type, handle, record, state, native, message, capacity, size
+ )
+ if record > 1:
+ warning_pending = False
+ return 100
+ text = "count warning".encode("utf-16le")
+ assert capacity > len(text) // 2
+ c.memmove(state, "01000\\0".encode("utf-16le"), 12)
+ c.memmove(message, text + b"\\0\\0", len(text) + 2)
+ c.cast(native, c.POINTER(c.c_int))[0] = 0
+ c.cast(size, c.POINTER(short))[0] = len(text) // 2
+ return 0
+ except BaseException as error:
+ errors.append(type(error).__name__)
+ return -1
+
+ with mssql_python.connect(os.environ["DB_CONNECTION_STRING"]) as connection:
+ with connection.cursor() as cursor:
+ value_sql = (
+ "CASE WHEN n = 2 THEN CAST(0x00D8 AS NVARCHAR(10)) "
+ "ELSE CAST(N'text' AS NVARCHAR(10)) END" if mode == "decode_error" else
+ "CAST(N'text' AS NVARCHAR(MAX))" if mode == "lob" else
+ "CAST(N'text' AS NVARCHAR(10))" if mode in ("mixed", "prefix") else
+ "n + 10"
+ )
+ query = (
+ f"SELECT n AS a, {value_sql} AS b FROM "
+ "(VALUES (1),(2),(3),(4),(5),(6)) v(n) ORDER BY n"
+ )
+ cursor.execute(query)
+ saved_count, saved_diag = count_slot.value, diag_slot.value
+ assert saved_count and saved_diag
+ original_count, original_diag = count_type(saved_count), diag_type(saved_diag)
+ count_slot.value = c.cast(counted, pointer).value
+ diag_slot.value = c.cast(diagnostic, pointer).value
+ try:
+ expected = 1
+ if mode == "prefix":
+ assert ddbc.DDBCSQLFetch(cursor.hstmt) == 0
+ prefix = []
+ assert ddbc.DDBCSQLGetData(
+ cursor.hstmt, 1, prefix, "utf-16le", "utf-16le", -8
+ ) == 0
+ assert prefix == [1] and calls == []
+ expected = 2
+ if mode == "count_error":
+ try:
+ cursor.fetchmany(1)
+ except mssql_python.DatabaseError:
+ pass
+ else:
+ raise AssertionError("count failure was not propagated")
+ assert len(calls) == 1
+ calls.clear()
+ first = cursor.fetchmany(1)[0]
+ assert first[0] == expected and len(first) == 2
+ assert first[1] == (
+ "text" if mode in ("mixed", "prefix", "lob", "decode_error")
+ else expected + 10
+ )
+ # Cold eager count and DescribeColumns' independent count.
+ assert len(calls) == 2, (mode, calls)
+ if mode == "decode_error":
+ try:
+ cursor.fetchone()
+ except UnicodeDecodeError:
+ pass
+ else:
+ raise AssertionError("invalid UTF-16 must fail GetData decoding")
+ assert len(calls) == 2
+ assert tuple(cursor.fetchmany(1)[0]) == (3, "text")
+ assert len(calls) == 4 # failure invalidated count and metadata
+ expected = 3
+ cold = len(calls)
+ assert cursor.fetchmany(1)[0][0] == expected + 1
+ assert len(calls) == cold + (1 if mode == "generation" else 0)
+ warm = len(calls)
+ assert cursor.fetchval() == expected + 2
+ assert cursor.fetchmany(1)[0][0] == expected + 3
+ assert len(calls) == warm
+ assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2
+ assert ddbc.DDBCSQLNumResultCols(cursor.hstmt) == 2
+ assert len(calls) == warm + 2 # public API is still uncached
+ assert not errors, errors
+ if mode == "warning":
+ assert cursor.messages == [("[01000] (0)", "count warning")]
+ else:
+ assert not cursor.messages
+ while cursor.fetchmany(1):
+ pass
+ assert len(calls) == warm + 2 # EOF does not reacquire the count
+ if mode == "nextset":
+ cursor.execute("SELECT 7 AS a; SELECT 8 AS a, 9 AS b")
+ calls.clear()
+ assert tuple(cursor.fetchmany(1)[0]) == (7,)
+ assert len(calls) == 2
+ assert cursor.nextset()
+ calls.clear()
+ assert tuple(cursor.fetchmany(1)[0]) == (8, 9)
+ assert len(calls) == 2
+ assert cursor.fetchmany(1) == []
+ assert len(calls) == 2
+ finally:
+ count_slot.value, diag_slot.value = saved_count, saved_diag
+ assert not errors, errors
+ """)
+ _run_fetch_script(conn_str, script, mode, timeout=60)