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)