From 38305da5d587ae124e54cceec6b6b7f14d803109 Mon Sep 17 00:00:00 2001 From: Zach Maddox Date: Wed, 30 Sep 2026 14:02:20 -0400 Subject: [PATCH] defer pyspark loading - regression fix --- src/datacustomcode/spark/column_hints.py | 61 +++++++ tests/spark/test_column_hints_deferred.py | 200 ++++++++++++++++++++++ 2 files changed, 261 insertions(+) create mode 100644 tests/spark/test_column_hints_deferred.py diff --git a/src/datacustomcode/spark/column_hints.py b/src/datacustomcode/spark/column_hints.py index 9ea4edb..1f49d3e 100644 --- a/src/datacustomcode/spark/column_hints.py +++ b/src/datacustomcode/spark/column_hints.py @@ -28,7 +28,14 @@ """ from __future__ import annotations +import importlib.abc import logging +import sys +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from importlib.machinery import ModuleSpec + from types import ModuleType logger = logging.getLogger(__name__) @@ -125,9 +132,63 @@ def __getattr__with_hint(self, name): # type: ignore[no-untyped-def] def install_column_casing_hints() -> None: + """Arrange for the column-casing hint hooks to be installed. + + Never imports pyspark directly. If pyspark is already loaded, the hooks + are installed immediately; otherwise, installation is deferred until + ``pyspark.sql`` is actually imported by someone. Idempotent; never raises. + """ + if "pyspark.sql" in sys.modules: + _install_hooks_now() + return + if any(isinstance(finder, _PySparkImportHook) for finder in sys.meta_path): + return + try: + sys.meta_path.insert(0, _PySparkImportHook()) + except Exception as exc: # pragma: no cover - defensive + logger.debug(f"Could not defer column-casing hint installation: {exc}") + + +def _install_hooks_now() -> None: """Install the column-casing hint hooks. Idempotent; never raises.""" try: _install_analysis_exception_hook() _install_getattr_hook() except Exception as exc: # pragma: no cover - defensive logger.debug(f"Could not install column-casing hint hooks: {exc}") + + +class _PySparkImportHook(importlib.abc.MetaPathFinder): + """Installs the hint hooks right after ``pyspark.sql`` finishes loading. + + ``install_column_casing_hints`` must not import pyspark itself: this + package is also used in the pyspark-free Function code path that must + not pull pyspark in as a side effect of importing ``datacustomcode``. + Wrapping the real loader lets us defer the pyspark import until whoever + actually needs it (the SDK's own lazy accessors, or user code) imports + ``pyspark.sql`` on their own. + """ + + def find_spec( + self, fullname: str, path: object, target: ModuleType | None = None + ) -> ModuleSpec | None: + if fullname != "pyspark.sql": + return None + + for finder in sys.meta_path: + if finder is self: + continue + find_spec = getattr(finder, "find_spec", None) + if find_spec is None: + continue + spec = find_spec(fullname, path, target) + if spec is not None and spec.loader is not None: + original_exec_module = spec.loader.exec_module + + def exec_module(module: ModuleType, _orig=original_exec_module) -> None: + _orig(module) + _install_hooks_now() + + spec.loader.exec_module = exec_module # type: ignore[method-assign] + return spec # type: ignore[no-any-return] + return None diff --git a/tests/spark/test_column_hints_deferred.py b/tests/spark/test_column_hints_deferred.py new file mode 100644 index 0000000..83542f8 --- /dev/null +++ b/tests/spark/test_column_hints_deferred.py @@ -0,0 +1,200 @@ +# Copyright (c) 2025, Salesforce, Inc. +# SPDX-License-Identifier: Apache-2 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tests for deferred (pyspark-free) installation of the column-casing hooks. + +These exercise ``install_column_casing_hints`` before ``pyspark.sql`` has +been imported. ``test_column_hints.py`` always imports ``pyspark.sql`` at +module scope, so it can never observe the deferred branch; the subprocess +tests here start from a clean interpreter where pyspark genuinely hasn't +been imported yet. +""" +from __future__ import annotations + +import subprocess +import sys +import textwrap + +import pytest + +from datacustomcode.spark import column_hints +from datacustomcode.spark.column_hints import _PySparkImportHook + + +def _run(script: str) -> None: + """Run *script* in a fresh interpreter; fail loudly on error.""" + result = subprocess.run( + [sys.executable, "-c", textwrap.dedent(script)], + capture_output=True, + text=True, + timeout=60, + ) + if result.returncode != 0: + pytest.fail( + f"subprocess failed (rc={result.returncode}):\n" + f"stdout:\n{result.stdout}\nstderr:\n{result.stderr}" + ) + + +def test_install_before_pyspark_import_does_not_import_pyspark() -> None: + _run( + """ + import sys + from datacustomcode.spark.column_hints import install_column_casing_hints + + assert "pyspark" not in sys.modules + install_column_casing_hints() + assert "pyspark" not in sys.modules + assert "pyspark.sql" not in sys.modules + """ + ) + + +def test_install_before_pyspark_import_registers_finder() -> None: + _run( + """ + import sys + from datacustomcode.spark.column_hints import ( + _PySparkImportHook, + install_column_casing_hints, + ) + + install_column_casing_hints() + hooks = [f for f in sys.meta_path if isinstance(f, _PySparkImportHook)] + assert len(hooks) == 1 + """ + ) + + +def test_repeated_deferred_install_does_not_stack_finders() -> None: + _run( + """ + import sys + from datacustomcode.spark.column_hints import ( + _PySparkImportHook, + install_column_casing_hints, + ) + + install_column_casing_hints() + install_column_casing_hints() + install_column_casing_hints() + hooks = [f for f in sys.meta_path if isinstance(f, _PySparkImportHook)] + assert len(hooks) == 1 + """ + ) + + +@pytest.mark.spark +def test_hooks_install_when_pyspark_sql_is_later_imported() -> None: + _run( + """ + from datacustomcode.spark.column_hints import ( + _WRAPPED_MARKER, + install_column_casing_hints, + ) + + install_column_casing_hints() + + import pyspark.sql # noqa: F401 - triggers the deferred install + from pyspark.errors.exceptions import captured + from pyspark.sql import DataFrame + + assert getattr(captured.convert_exception, _WRAPPED_MARKER, False) + assert getattr(DataFrame.__getattr__, _WRAPPED_MARKER, False) + """ + ) + + +@pytest.mark.spark +def test_install_installs_immediately_when_pyspark_already_loaded() -> None: + _run( + """ + import sys + + import pyspark.sql # noqa: F401 - pyspark is loaded before install + from datacustomcode.spark.column_hints import ( + _PySparkImportHook, + _WRAPPED_MARKER, + install_column_casing_hints, + ) + + install_column_casing_hints() + + from pyspark.errors.exceptions import captured + from pyspark.sql import DataFrame + + assert getattr(captured.convert_exception, _WRAPPED_MARKER, False) + assert getattr(DataFrame.__getattr__, _WRAPPED_MARKER, False) + assert not any(isinstance(f, _PySparkImportHook) for f in sys.meta_path) + """ + ) + + +class _FakeLoader: + def __init__(self) -> None: + self.exec_calls: list[object] = [] + + def exec_module(self, module: object) -> None: + self.exec_calls.append(module) + + +class _FakeSpec: + def __init__(self, loader: _FakeLoader) -> None: + self.loader = loader + + +class _FakeFinder: + """Stands in for the real pyspark finder on ``sys.meta_path``.""" + + def __init__(self, spec: object | None) -> None: + self._spec = spec + + def find_spec(self, fullname, path, target=None): # type: ignore[no-untyped-def] + return self._spec + + +def test_find_spec_ignores_other_module_names() -> None: + hook = _PySparkImportHook() + assert hook.find_spec("pyspark", None) is None + assert hook.find_spec("some.other.module", None) is None + + +def test_find_spec_returns_none_when_no_other_finder_can_resolve(monkeypatch) -> None: + hook = _PySparkImportHook() + monkeypatch.setattr(sys, "meta_path", [hook, _FakeFinder(None)]) + assert hook.find_spec("pyspark.sql", None) is None + + +def test_find_spec_wraps_loader_exec_module_and_installs_hooks_after( + monkeypatch, +) -> None: + loader = _FakeLoader() + spec = _FakeSpec(loader) + hook = _PySparkImportHook() + monkeypatch.setattr(sys, "meta_path", [hook, _FakeFinder(spec)]) + + calls: list[str] = [] + monkeypatch.setattr( + column_hints, "_install_hooks_now", lambda: calls.append("installed") + ) + + found = hook.find_spec("pyspark.sql", None) + assert found is spec + # The loader's exec_module has been swapped for a wrapper. + assert loader.exec_module is not _FakeLoader.exec_module.__get__(loader) + + fake_module = object() + spec.loader.exec_module(fake_module) + assert loader.exec_calls == [fake_module] # original behavior preserved + assert calls == ["installed"] # hint hooks installed right after