Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 61 additions & 0 deletions src/datacustomcode/spark/column_hints.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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
200 changes: 200 additions & 0 deletions tests/spark/test_column_hints_deferred.py
Original file line number Diff line number Diff line change
@@ -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
Loading