From 62a606f3c6c11a6cd1a93a9b7011d32266c9d4ea Mon Sep 17 00:00:00 2001 From: Dave Borowitz Date: Fri, 18 Sep 2026 14:29:51 -0700 Subject: [PATCH 1/2] feat: Automatically load notebook extras on import with opt-out controls Automatically initialize interactive notebook extras when importing google.cloud.managed_spark_connect inside an IPython kernel: - Inject colabsqlviz's explore_dataframe() into IPython user_ns - Load the %dpip line magic extension (google.cloud.managed_spark_magics) - Load the %%sparksql cell magic extension (sparksql_magic) - Add google-colabsqlviz>=0.3.0 and sparksql-magic>=0.0.3 to dependencies Add opt-out and runtime configuration controls: - Environment variable: MANAGED_SPARK_CONNECT_ENABLE_EXTRAS=false - IPython traitlet: ManagedSparkConnect.enable_extras = False (supports both persistent file-based config and runtime toggling via %config, tracking and only undoing changes that managed_spark_connect itself performed). --- DEVELOPING.md | 29 +- README.md | 71 +- .../cloud/managed_spark_connect/__init__.py | 6 + .../cloud/managed_spark_connect/_ipython.py | 368 ++++++++++ google/cloud/managed_spark_magics/__init__.py | 9 + google/cloud/managed_spark_magics/magics.py | 6 +- requirements-dev.txt | 2 + setup.py | 7 + tests/unit/test_init.py | 669 ++++++++++++++++++ 9 files changed, 1114 insertions(+), 53 deletions(-) create mode 100644 google/cloud/managed_spark_connect/_ipython.py diff --git a/DEVELOPING.md b/DEVELOPING.md index c9a8a5ea..36ca1fa3 100644 --- a/DEVELOPING.md +++ b/DEVELOPING.md @@ -39,27 +39,14 @@ env \ pytest --tb=auto -v ``` -## Testing with Magic Support - -To run tests with magic functionality, install the required dependencies manually: - -```sh -pip install . -pip install IPython sparksql-magic -``` - -Then run tests as normal. Any magic-related tests will automatically detect and use the available dependencies. - -## Testing without Magic Support - -To run tests without the magic dependencies, simply install the base package: - -```sh -pip install . -pytest -``` - -Tests that require magic functionality will be automatically skipped if the dependencies are not available. +## Testing the Notebook Extras + +The notebook extras (`explore_dataframe`, `%dpip`, `%%sparksql`) are regular +`install_requires` dependencies, so `pip install .` is enough and there is no +"without magic support" configuration to test. The tests in +`tests/unit/test_init.py` assume `google-colabsqlviz`, `sparksql-magic`, +`ipython` and `traitlets` are importable and will fail, not skip, if they are +not. The integration tests in particular can take a while to run. To speed up the testing cycle, you can run them in parallel. You can do so using the `xdist` diff --git a/README.md b/README.md index c1117926..f94de49c 100644 --- a/README.md +++ b/README.md @@ -121,53 +121,62 @@ To create or connect to a named session: 5. A session with a given ID that is in a TERMINATED state cannot be reused. It must be deleted before a new session with the same ID can be created. -### Using Spark SQL Magic Commands (Jupyter Notebooks) +### Jupyter Notebook Extras -The package supports the [sparksql-magic](https://github.com/cryeo/sparksql-magic) library for executing Spark SQL queries directly in Jupyter notebooks. +When you import the package inside an IPython kernel, it automatically sets up +a few interactive conveniences--no separate install or setup required: -**Installation**: To use magic commands, install the required dependencies manually: -```bash -pip install google-cloud-spark-connect -pip install IPython sparksql-magic +- `explore_dataframe()` from + [google-colabsqlviz](https://pypi.org/project/google-colabsqlviz/) is injected + into your notebook globals. +- The `%dpip` line magic is loaded. +- The `%%sparksql` cell magic from + [sparksql-magic](https://github.com/cryeo/sparksql-magic) is loaded. + +```python +import google.cloud.managed_spark_connect # extras load here ``` -1. Load the magic extension: - ```python - %load_ext sparksql_magic - ``` +The extras won't override anything you've already set up, for example if `explore_dataframe` is already present, or `%%sparksql` magic is loaded from somewhere else, these are left alone. + +#### Opting out + +Set the environment variable before starting the kernel: + +```sh +export MANAGED_SPARK_CONNECT_ENABLE_EXTRAS=false +``` + +Alternatively, configure it through IPython, either persistently in +`~/.ipython/profile_default/ipython_config.py`: + +```python +c.ManagedSparkConnect.enable_extras = False +``` -2. Configure default settings (optional): +Or at runtime: + +```python +%config ManagedSparkConnect.enable_extras = False +``` + +An explicit IPython setting takes precedence over the environment variable. + +#### Using `%%sparksql` + +1. Configure default settings (optional): ```python %config SparkSql.limit=20 ``` -3. Execute SQL queries: +2. Execute SQL queries: ```python %%sparksql SELECT * FROM your_table ``` -4. Advanced usage with options: - ```python - # Cache results and create a view - %%sparksql --cache --view result_view df - SELECT * FROM your_table WHERE condition = true - ``` - -Available options: -- `--cache` / `-c`: Cache the DataFrame -- `--eager` / `-e`: Cache with eager loading -- `--view VIEW` / `-v VIEW`: Create a temporary view -- `--limit N` / `-l N`: Override default row display limit -- `variable_name`: Store result in a variable - See [sparksql-magic](https://github.com/cryeo/sparksql-magic) for more examples. -**Note**: Magic commands are optional. If you only need basic ManagedSparkSession functionality without Jupyter magic support, install only the base package: -```bash -pip install google-cloud-spark-connect -``` - ## Migrating from dataproc-spark-connect The `dataproc-spark-connect` package has been renamed to `google-cloud-spark-connect`. This is a breaking change with no compatibility shims — you need to update your code in the following places when you switch to the new package. diff --git a/google/cloud/managed_spark_connect/__init__.py b/google/cloud/managed_spark_connect/__init__.py index 23567583..68366d35 100644 --- a/google/cloud/managed_spark_connect/__init__.py +++ b/google/cloud/managed_spark_connect/__init__.py @@ -14,6 +14,10 @@ import importlib.metadata import warnings +from ._ipython import ( + ManagedSparkConnect, + _init_extras, +) from .session import ManagedSparkSession old_package_names = ["google-spark-connect", "dataproc-spark-connect"] @@ -28,3 +32,5 @@ ) except Exception: pass + +_init_extras() diff --git a/google/cloud/managed_spark_connect/_ipython.py b/google/cloud/managed_spark_connect/_ipython.py new file mode 100644 index 00000000..fb9ce181 --- /dev/null +++ b/google/cloud/managed_spark_connect/_ipython.py @@ -0,0 +1,368 @@ +# Copyright 2026 Google LLC +# +# 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. + +from dataclasses import dataclass, field +import logging +import os +import sys +from typing import Any, Dict, Optional, Set + +from traitlets import Bool, default, observe +from traitlets.config import Config, Configurable + +logger = logging.getLogger(__name__) + +ENV_ENABLE_EXTRAS = "MANAGED_SPARK_CONNECT_ENABLE_EXTRAS" + +_ENV_FALSE_VALUES = frozenset(("0", "false", "no", "off")) + +_EXTRAS_EXTENSIONS = ( + ("google.cloud.managed_spark_magics", "line", "dpip"), + ("sparksql_magic", "cell", "sparksql"), +) + + +def _is_env_extras_enabled() -> bool: + val = os.getenv(ENV_ENABLE_EXTRAS) + if val is None or not val.strip(): + # Treat an empty/whitespace-only value the same as an unset variable, + # so that `MANAGED_SPARK_CONNECT_ENABLE_EXTRAS=` is not silently + # interpreted as "enabled". + return True + return val.strip().lower() not in _ENV_FALSE_VALUES + + +class ManagedSparkConnect(Configurable): + """IPython configuration for google.cloud.managed_spark_connect extras.""" + + # No static default_value: the dynamic `_default_enable_extras` below + # always wins. + enable_extras = Bool( + help=( + "Whether to automatically load notebook extras " + "(explore_dataframe, %dpip, %%sparksql)." + ), + ).tag(config=True) + + def __init__(self, shell: Any = None, **kwargs: Any) -> None: + self._initialized = False + self._shell = shell + super().__init__(**kwargs) + self._initialized = True + + @default("enable_extras") + def _default_enable_extras(self) -> bool: + return _is_env_extras_enabled() + + @observe("enable_extras") + def _observe_enable_extras(self, change: Dict[str, Any]) -> None: + # Traitlets applies file-based config inside `super().__init__()`, and + # the resulting notification fires before the shell has been wired up. + # Skip it: `_init_extras` loads the extras once construction is done. + if not getattr(self, "_initialized", False): + return + shell = self._shell or self.parent + if shell is None: + return + if change["new"]: + _load_extras_for_shell(shell) + else: + _unload_extras_for_shell(shell) + + +@dataclass +class _ShellExtrasState: + shell: Any + injected_explore_dataframe: Any = None + loaded_extensions: Set[str] = field(default_factory=set) + config_instance: Optional[ManagedSparkConnect] = None + + +# Keyed by id() rather than being a WeakKeyDictionary. A weak key would not +# release anything here: the value reaches the key via +# `config_instance._shell` and traitlets' own `Configurable.parent`, and +# `InteractiveShell.__init__` hands a bound method to `atexit`, so a shell is +# retained for the life of the process regardless. `state.shell` is what makes +# id() keying safe, since CPython reuses addresses; see `_get_shell_state`. +_SHELL_STATES: Dict[int, _ShellExtrasState] = {} + + +def _get_shell_state(ip: Any) -> _ShellExtrasState: + shell_id = id(ip) + state = _SHELL_STATES.get(shell_id) + if state is None or state.shell is not ip: + state = _ShellExtrasState(shell=ip) + _SHELL_STATES[shell_id] = state + return state + + +def _is_same_class(obj: Any, cls: type) -> bool: + """Compares by module+qualname so that reloaded classes still match.""" + obj_cls = type(obj) + return ( + obj_cls.__module__ == cls.__module__ + and obj_cls.__qualname__ == cls.__qualname__ + ) + + +def _get_or_create_config(ip: Any) -> ManagedSparkConnect: + state = _get_shell_state(ip) + ip_config = getattr(ip, "config", None) + + if state.config_instance is None: + parent = ip if isinstance(ip, Configurable) else None + config = ( + ip_config + if (parent is None and isinstance(ip_config, Config)) + else None + ) + # Precedence is left entirely to traitlets: an explicit + # `c.ManagedSparkConnect.enable_extras` from the shell's config wins, + # otherwise `_default_enable_extras` consults the environment. + cfg = ManagedSparkConnect(shell=ip, parent=parent, config=config) + state.config_instance = cfg + else: + cfg = state.config_instance + cfg._shell = ip + + configurables = getattr(ip, "configurables", None) + if isinstance(configurables, list): + # Remove any stale ManagedSparkConnect instances (e.g. across module + # reloads) + configurables[:] = [ + c + for c in configurables + if c is cfg or not _is_same_class(c, ManagedSparkConnect) + ] + if cfg not in configurables: + configurables.append(cfg) + + return cfg + + +def _import_explore_dataframe() -> Any: + """Returns colabsqlviz's explore_dataframe(), or None if unavailable.""" + try: + from google.colabsqlviz.explore_dataframe import explore_dataframe + + return explore_dataframe + except Exception: + logger.debug("Failed to import explore_dataframe", exc_info=True) + return None + + +def _safe_repr(value: Any, max_len: int = 60) -> str: + """repr() that tolerates objects whose __repr__ raises.""" + try: + val_repr = repr(value) + except Exception: + return f"<{type(value).__name__} with a failing __repr__>" + if len(val_repr) > max_len: + return f"{val_repr[:max_len]}..." + return val_repr + + +def _print_extras_failure() -> None: + """Tells the user something went wrong, without diagnosing what. + + Every individual failure below is logged at debug level and otherwise + swallowed, which leaves a broken install looking exactly like extras being + switched off. One line, the same regardless of cause, is enough to point + somebody at the logs. + """ + print( + "\033[94m⚠️ [google.cloud.managed_spark_connect]\033[0m" + " Failed to load notebook extras." + " For details, enable debug logging and restart the kernel." + ) + + +def _is_extension_loaded( + ip: Any, ext_name: str, magic_kind: str, magic_name: str +) -> bool: + mod = sys.modules.get(ext_name) + if mod is not None and getattr( + getattr(mod, "__spec__", None), "_initializing", False + ): + return True + + ext_mgr = getattr(ip, "extension_manager", None) + loaded = getattr(ext_mgr, "loaded", None) + if isinstance(loaded, (set, dict, list, tuple)) and ext_name in loaded: + return True + + magics_mgr = getattr(ip, "magics_manager", None) + magics = getattr(magics_mgr, "magics", None) + if isinstance(magics, dict): + kind_dict = magics.get(magic_kind) + if isinstance(kind_dict, dict) and magic_name in kind_dict: + return True + + return False + + +def _load_extras_for_shell(ip: Any) -> None: + state = _get_shell_state(ip) + target_name = "explore_dataframe" + failed = False + + # 1. Inject explore_dataframe from google-colabsqlviz + try: + user_ns = getattr(ip, "user_ns", None) + if isinstance(user_ns, dict): + if ( + state.injected_explore_dataframe is not None + and user_ns.get(target_name) is state.injected_explore_dataframe + ): + pass + elif target_name not in user_ns: + from google.colabsqlviz.explore_dataframe import ( + explore_dataframe, + ) + + ip.push({target_name: explore_dataframe}) + state.injected_explore_dataframe = explore_dataframe + + msg = ( + "\033[94m👉 [google.cloud.managed_spark_connect]\033[0m" + f" Injected \033[1m{target_name}()\033[0m into globals." + f" Use \033[1m{target_name}(df)\033[0m to interactively explore your data." + ) + print(msg) + else: + current_val = user_ns[target_name] + ours = _import_explore_dataframe() + if ours is not None and current_val is ours: + # Already bound to the exact function we would have + # injected, e.g. because the user imported it themselves. + # Warning about that would be pure noise. + pass + else: + val_repr = _safe_repr(current_val) + print( + "\033[94m⚠️ [google.cloud.managed_spark_connect]\033[0m" + f" Did not inject \033[1m{target_name}()\033[0m" + f" because it is already defined as: {val_repr}" + ) + except Exception: + logger.debug( + "Failed to inject %s into IPython user_ns", + target_name, + exc_info=True, + ) + failed = True + + # 2. Load magic extensions silently + ext_mgr = getattr(ip, "extension_manager", None) + if ext_mgr is not None and hasattr(ext_mgr, "load_extension"): + for ext_name, magic_kind, magic_name in _EXTRAS_EXTENSIONS: + try: + if ext_name not in state.loaded_extensions and not ( + _is_extension_loaded(ip, ext_name, magic_kind, magic_name) + ): + res = ext_mgr.load_extension(ext_name) + if res is None: + state.loaded_extensions.add(ext_name) + except Exception: + logger.debug( + "Failed to load IPython extension %s", + ext_name, + exc_info=True, + ) + failed = True + + if failed: + _print_extras_failure() + + +def _unload_extras_for_shell(ip: Any) -> None: + state = _get_shell_state(ip) + target_name = "explore_dataframe" + + # 1. Remove explore_dataframe only if we injected it and it wasn't overwritten + try: + if state.injected_explore_dataframe is not None: + user_ns = getattr(ip, "user_ns", None) + if ( + isinstance(user_ns, dict) + and user_ns.get(target_name) is state.injected_explore_dataframe + ): + del user_ns[target_name] + state.injected_explore_dataframe = None + except Exception: + logger.debug( + "Failed to remove %s from IPython user_ns", + target_name, + exc_info=True, + ) + + # 2. Unload only the extensions that we loaded + ext_mgr = getattr(ip, "extension_manager", None) + magics_mgr = getattr(ip, "magics_manager", None) + magics = getattr(magics_mgr, "magics", None) + loaded = getattr(ext_mgr, "loaded", None) + + for ext_name, magic_kind, magic_name in _EXTRAS_EXTENSIONS: + if ext_name in state.loaded_extensions: + try: + if ext_mgr is not None and hasattr(ext_mgr, "unload_extension"): + ext_mgr.unload_extension(ext_name) + except Exception: + logger.debug( + "Failed to unload IPython extension %s", + ext_name, + exc_info=True, + ) + try: + if isinstance(magics, dict): + kind_dict = magics.get(magic_kind) + if isinstance(kind_dict, dict): + kind_dict.pop(magic_name, None) + if isinstance(loaded, set): + loaded.discard(ext_name) + except Exception: + logger.debug( + "Failed to clean up magic %s for extension %s", + magic_name, + ext_name, + exc_info=True, + ) + state.loaded_extensions.discard(ext_name) + + +def _init_extras() -> None: + """Initialize notebook extras on package import unless opted out.""" + ip = None + try: + from IPython import get_ipython + + ip = get_ipython() + if ip is None: + return + + # `cfg.enable_extras` already folds in both the environment variable + # (via the trait's dynamic default) and any explicit configuration, and + # it is the same value `%config` toggles later, so it is the only thing + # that should be consulted here. + cfg = _get_or_create_config(ip) + if cfg.enable_extras: + _load_extras_for_shell(ip) + except Exception: + logger.debug( + "Failed to initialize IPython extras on import", exc_info=True + ) + if ip is not None: + # Only worth saying inside a shell: outside one there are no extras + # to miss, and printing would be noise in ordinary scripts. + _print_extras_failure() diff --git a/google/cloud/managed_spark_magics/__init__.py b/google/cloud/managed_spark_magics/__init__.py index 79632f57..eb667258 100644 --- a/google/cloud/managed_spark_magics/__init__.py +++ b/google/cloud/managed_spark_magics/__init__.py @@ -17,3 +17,12 @@ def load_ipython_extension(ipython): ipython.register_magics(ManagedSparkMagics) + + +def unload_ipython_extension(ipython): + magics_mgr = getattr(ipython, "magics_manager", None) + magics = getattr(magics_mgr, "magics", None) + if isinstance(magics, dict): + line_magics = magics.get("line") + if isinstance(line_magics, dict): + line_magics.pop("dpip", None) diff --git a/google/cloud/managed_spark_magics/magics.py b/google/cloud/managed_spark_magics/magics.py index 54363ae1..38ca3c7c 100644 --- a/google/cloud/managed_spark_magics/magics.py +++ b/google/cloud/managed_spark_magics/magics.py @@ -16,7 +16,6 @@ import shlex from IPython.core.magic import (Magics, magics_class, line_magic) -from google.cloud.managed_spark_connect import ManagedSparkSession @magics_class @@ -35,6 +34,11 @@ def dpip(self, line): Custom magic to install pip packages as Spark Connect artifacts. Usage: %dpip install pandas numpy """ + # Imported here rather than at module scope: google.cloud + # .managed_spark_connect loads this extension from its own __init__, + # so a module-level import back into it would be circular. + from google.cloud.managed_spark_connect import ManagedSparkSession + try: args = shlex.split(line) diff --git a/requirements-dev.txt b/requirements-dev.txt index 5cf7026e..09806396 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -1,5 +1,6 @@ google-api-core>=2.19 google-cloud-dataproc>=5.18 +google-colabsqlviz>=0.3.0 ipython~=9.1 ipywidgets>=8.0.0 packaging>=20.0 @@ -8,4 +9,5 @@ pyspark[connect]~=4.0.0 setuptools>=72.0 sparksql-magic>=0.0.3 tqdm>=4.67 +traitlets>=5.1 websockets>=14.0 diff --git a/setup.py b/setup.py index 89637d54..5b45fdcc 100644 --- a/setup.py +++ b/setup.py @@ -31,9 +31,16 @@ install_requires=[ "google-api-core>=2.19", "google-cloud-dataproc>=5.18", + "google-colabsqlviz>=0.3.0", + # Imported directly by managed_spark_connect._ipython and + # managed_spark_magics; previously these only arrived transitively via + # google-colabsqlviz and sparksql-magic. + "ipython>=8.0", "packaging>=20.0", "pyspark[connect]~=4.0.0", + "sparksql-magic>=0.0.3", "tqdm>=4.67", + "traitlets>=5.1", "websockets>=14.0", ], ) diff --git a/tests/unit/test_init.py b/tests/unit/test_init.py index 794f6fa8..c3eb2a92 100644 --- a/tests/unit/test_init.py +++ b/tests/unit/test_init.py @@ -135,5 +135,674 @@ def test_invalid_runtime_version_logs_warning(self, mock_logger): ) +class TestDependencies(unittest.TestCase): + + def test_colabsqlviz_available(self): + import importlib.metadata + from packaging.version import Version + import google.colabsqlviz + from google.colabsqlviz.explore_dataframe import ( + display, + explore_dataframe, + ) + + self.assertIsNotNone(google.colabsqlviz) + self.assertTrue(callable(display)) + self.assertTrue(callable(explore_dataframe)) + version = importlib.metadata.version("google-colabsqlviz") + self.assertGreaterEqual(Version(version), Version("0.3.0")) + + def test_sparksql_magic_available(self): + import importlib.metadata + from packaging.version import Version + import sparksql_magic + + self.assertIsNotNone(sparksql_magic) + self.assertTrue(callable(sparksql_magic.load_ipython_extension)) + version = importlib.metadata.version("sparksql-magic") + self.assertGreaterEqual(Version(version), Version("0.0.3")) + + def test_managed_spark_magics_available(self): + import google.cloud.managed_spark_magics as magics_ext + + self.assertTrue(callable(magics_ext.load_ipython_extension)) + self.assertTrue(callable(magics_ext.unload_ipython_extension)) + + +def _make_mock_shell(user_ns=None, config=None): + """Builds a mock shell whose push() mutates user_ns, like the real one. + + Without the side effect, mock shells silently swallow ip.push() and only + pass if the code under test also writes to user_ns directly. + """ + mock_ip = mock.Mock() + mock_ip.user_ns = {} if user_ns is None else user_ns + mock_ip.config = config + mock_ip.configurables = [] + mock_ip.push.side_effect = mock_ip.user_ns.update + mock_ip.extension_manager.load_extension.return_value = None + return mock_ip + + +class TestExtrasUnit(unittest.TestCase): + + def setUp(self): + import os + from google.cloud.managed_spark_connect import _ipython + + env_patcher = mock.patch.dict(os.environ) + env_patcher.start() + self.addCleanup(env_patcher.stop) + os.environ.pop(_ipython.ENV_ENABLE_EXTRAS, None) + _ipython._SHELL_STATES.clear() + self.addCleanup(_ipython._SHELL_STATES.clear) + + @mock.patch("sys.stdout", new_callable=unittest.mock.MagicMock) + @mock.patch("IPython.get_ipython", return_value=None) + def test_no_extras_when_not_in_ipython(self, mock_get_ipython, mock_stdout): + from google.cloud.managed_spark_connect._ipython import _init_extras + + _init_extras() + mock_get_ipython.assert_called_once() + mock_stdout.write.assert_not_called() + + @mock.patch("IPython.get_ipython") + def test_successful_extras_loading(self, mock_get_ipython): + import io + from google.colabsqlviz.explore_dataframe import explore_dataframe + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + mock_ip.push.assert_called_once_with( + {"explore_dataframe": explore_dataframe} + ) + self.assertIs(mock_ip.user_ns["explore_dataframe"], explore_dataframe) + mock_ip.extension_manager.load_extension.assert_has_calls( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + output = stdout_capture.getvalue() + self.assertIn("[google.cloud.managed_spark_connect]", output) + self.assertIn("Injected", output) + self.assertIn("explore_dataframe()", output) + self.assertNotIn("dpip", output) + self.assertNotIn("sparksql", output) + + @mock.patch("IPython.get_ipython") + def test_fallback_when_already_defined_short_repr(self, mock_get_ipython): + import io + from google.cloud.managed_spark_connect._ipython import ( + _get_or_create_config, + _init_extras, + ) + + mock_ip = _make_mock_shell(user_ns={"explore_dataframe": 42}) + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + mock_ip.push.assert_not_called() + output = stdout_capture.getvalue() + self.assertIn("[google.cloud.managed_spark_connect]", output) + self.assertIn("Did not inject", output) + self.assertIn("explore_dataframe()", output) + self.assertIn("because it is already defined as: 42", output) + self.assertNotIn("42...", output) + + # Disabling extras via traitlet must not remove user-defined explore_dataframe + _get_or_create_config(mock_ip).enable_extras = False + self.assertEqual(mock_ip.user_ns.get("explore_dataframe"), 42) + + @mock.patch("IPython.get_ipython") + def test_fallback_when_already_defined_long_repr(self, mock_get_ipython): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + long_val = "x" * 100 + mock_ip = _make_mock_shell(user_ns={"explore_dataframe": long_val}) + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + mock_ip.push.assert_not_called() + output = stdout_capture.getvalue() + expected_repr = f"{repr(long_val)[:60]}..." + self.assertIn("[google.cloud.managed_spark_connect]", output) + self.assertIn("Did not inject", output) + self.assertIn("explore_dataframe()", output) + self.assertIn( + f"because it is already defined as: {expected_repr}", + output, + ) + + @mock.patch("IPython.get_ipython") + def test_fallback_when_already_defined_unreprable(self, mock_get_ipython): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + class Unreprable: + + def __repr__(self): + raise ValueError("no repr for you") + + mock_ip = _make_mock_shell(user_ns={"explore_dataframe": Unreprable()}) + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + mock_ip.push.assert_not_called() + output = stdout_capture.getvalue() + self.assertIn("Did not inject", output) + self.assertIn("Unreprable", output) + # A broken __repr__ must not abort the rest of the initialization. + mock_ip.extension_manager.load_extension.assert_has_calls( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + + @mock.patch("IPython.get_ipython") + def test_no_warning_when_already_defined_same_function( + self, mock_get_ipython + ): + import io + from google.colabsqlviz.explore_dataframe import explore_dataframe + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell( + user_ns={"explore_dataframe": explore_dataframe} + ) + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + # The name already refers to the exact function we would have injected, + # so there is nothing to warn the user about. + mock_ip.push.assert_not_called() + self.assertEqual(stdout_capture.getvalue(), "") + self.assertIs(mock_ip.user_ns["explore_dataframe"], explore_dataframe) + + @mock.patch("IPython.get_ipython") + def test_opt_out_via_env_var(self, mock_get_ipython): + import os + from google.cloud.managed_spark_connect._ipython import ( + ENV_ENABLE_EXTRAS, + _init_extras, + ) + + for falsy_val in ("false", "FALSE", "0", "no", "off", " false "): + with self.subTest(val=falsy_val): + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + with mock.patch.dict( + os.environ, {ENV_ENABLE_EXTRAS: falsy_val} + ): + _init_extras() + + mock_ip.push.assert_not_called() + mock_ip.extension_manager.load_extension.assert_not_called() + self.assertNotIn("explore_dataframe", mock_ip.user_ns) + + @mock.patch("IPython.get_ipython") + def test_opt_out_via_traitlet_config(self, mock_get_ipython): + from traitlets.config import Config + from google.cloud.managed_spark_connect._ipython import _init_extras + + cfg = Config() + cfg.ManagedSparkConnect.enable_extras = False + mock_ip = _make_mock_shell(config=cfg) + mock_get_ipython.return_value = mock_ip + + _init_extras() + + mock_ip.push.assert_not_called() + mock_ip.extension_manager.load_extension.assert_not_called() + self.assertNotIn("explore_dataframe", mock_ip.user_ns) + + @mock.patch("IPython.get_ipython") + def test_explicit_config_beats_env_var(self, mock_get_ipython): + import io + import os + from traitlets.config import Config + from google.cloud.managed_spark_connect._ipython import ( + ENV_ENABLE_EXTRAS, + _init_extras, + ) + + # An explicit `c.ManagedSparkConnect.enable_extras = True` is a more + # specific signal than the environment variable, so it must win. + cfg = Config() + cfg.ManagedSparkConnect.enable_extras = True + mock_ip = _make_mock_shell(config=cfg) + mock_get_ipython.return_value = mock_ip + + with mock.patch.dict(os.environ, {ENV_ENABLE_EXTRAS: "false"}): + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + self.assertIn("explore_dataframe", mock_ip.user_ns) + mock_ip.extension_manager.load_extension.assert_has_calls( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + + @mock.patch("IPython.get_ipython") + def test_blank_env_var_is_treated_as_unset(self, mock_get_ipython): + import io + import os + from google.cloud.managed_spark_connect._ipython import ( + ENV_ENABLE_EXTRAS, + _init_extras, + ) + + for blank_val in ("", " "): + with self.subTest(val=blank_val): + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + with mock.patch.dict( + os.environ, {ENV_ENABLE_EXTRAS: blank_val} + ): + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + self.assertIn("explore_dataframe", mock_ip.user_ns) + + @mock.patch("IPython.get_ipython") + def test_traitlet_disable_preserves_overwritten_explore_dataframe( + self, mock_get_ipython + ): + import io + from google.cloud.managed_spark_connect._ipython import ( + _get_or_create_config, + _init_extras, + ) + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + custom_fn = lambda df: "custom" + mock_ip.user_ns["explore_dataframe"] = custom_fn + + _get_or_create_config(mock_ip).enable_extras = False + self.assertIs(mock_ip.user_ns.get("explore_dataframe"), custom_fn) + mock_ip.extension_manager.unload_extension.assert_has_calls( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + + @mock.patch("IPython.get_ipython", side_effect=RuntimeError("Kernel error")) + def test_exception_silently_caught(self, mock_get_ipython): + from google.cloud.managed_spark_connect._ipython import _init_extras + + try: + _init_extras() + except Exception as e: + self.fail(f"_init_extras raised an exception: {e}") + + @mock.patch("IPython.get_ipython", side_effect=RuntimeError("Kernel error")) + def test_no_failure_message_outside_a_shell(self, mock_get_ipython): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + # There is no shell, so there are no extras to miss. + self.assertEqual(stdout_capture.getvalue(), "") + + @mock.patch("IPython.get_ipython") + def test_failure_message_printed_once(self, mock_get_ipython): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell() + mock_ip.extension_manager.load_extension.side_effect = RuntimeError( + "boom" + ) + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch( + "google.cloud.managed_spark_connect._ipython" + "._import_explore_dataframe", + side_effect=RuntimeError("boom"), + ): + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + output = stdout_capture.getvalue() + self.assertIn("Failed to load notebook extras", output) + # Three separate failures (the injection and both extensions) must not + # produce three separate messages. + self.assertEqual(output.count("Failed to load notebook extras"), 1) + + @mock.patch("IPython.get_ipython") + def test_failure_message_when_colabsqlviz_is_missing( + self, mock_get_ipython + ): + import builtins + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + real_import = builtins.__import__ + + def fail_colabsqlviz(name, *args, **kwargs): + if name.startswith("google.colabsqlviz"): + raise ImportError("no colabsqlviz here") + return real_import(name, *args, **kwargs) + + stdout_capture = io.StringIO() + with mock.patch("builtins.__import__", side_effect=fail_colabsqlviz): + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + self.assertNotIn("explore_dataframe", mock_ip.user_ns) + self.assertIn( + "Failed to load notebook extras", stdout_capture.getvalue() + ) + # The magics are independent of colabsqlviz and should still load. + mock_ip.extension_manager.load_extension.assert_has_calls( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + + def test_module_import_calls_init_extras(self): + import importlib + import google.cloud.managed_spark_connect as msc + + self.assertIsNotNone(msc.ManagedSparkConnect) + self.assertFalse(hasattr(msc, "enable_extras")) + self.assertFalse(hasattr(msc, "disable_extras")) + + # reload() re-executes `from ._ipython import _init_extras`, which + # copies whatever is patched at that moment into the package namespace. + # Undoing the patch does not undo that copy, so reload again outside it + # or every later test in this process sees the mock. + self.addCleanup(importlib.reload, msc) + with mock.patch( + "google.cloud.managed_spark_connect._ipython._init_extras" + ) as mock_init: + importlib.reload(msc) + mock_init.assert_called_once() + + +class TestExtrasInteractiveShell(unittest.TestCase): + + def setUp(self): + import os + from IPython.core.interactiveshell import InteractiveShell + from google.cloud.managed_spark_connect import _ipython + + env_patcher = mock.patch.dict(os.environ) + env_patcher.start() + self.addCleanup(env_patcher.stop) + os.environ.pop(_ipython.ENV_ENABLE_EXTRAS, None) + _ipython._SHELL_STATES.clear() + + self.ip = InteractiveShell.instance() + self.ip.user_ns.pop("explore_dataframe", None) + self.ip.magics_manager.magics.get("line", {}).pop("dpip", None) + self.ip.magics_manager.magics.get("cell", {}).pop("sparksql", None) + self.ip.extension_manager.loaded.discard( + "google.cloud.managed_spark_magics" + ) + self.ip.extension_manager.loaded.discard("sparksql_magic") + if "ManagedSparkConnect" in self.ip.config: + del self.ip.config["ManagedSparkConnect"] + self.ip.configurables[:] = [ + c + for c in self.ip.configurables + if c.__class__.__name__ != "ManagedSparkConnect" + ] + + self._get_ipython_patcher = mock.patch( + "IPython.get_ipython", return_value=self.ip + ) + self._get_ipython_patcher.start() + + def tearDown(self): + from IPython.core.interactiveshell import InteractiveShell + from google.cloud.managed_spark_connect import _ipython + + self._get_ipython_patcher.stop() + self.ip.user_ns.pop("explore_dataframe", None) + self.ip.magics_manager.magics.get("line", {}).pop("dpip", None) + self.ip.magics_manager.magics.get("cell", {}).pop("sparksql", None) + self.ip.extension_manager.loaded.discard( + "google.cloud.managed_spark_magics" + ) + self.ip.extension_manager.loaded.discard("sparksql_magic") + if "ManagedSparkConnect" in self.ip.config: + del self.ip.config["ManagedSparkConnect"] + self.ip.configurables[:] = [ + c + for c in self.ip.configurables + if c.__class__.__name__ != "ManagedSparkConnect" + ] + InteractiveShell.clear_instance() + _ipython._SHELL_STATES.clear() + + def test_real_shell_loads_all_extras_and_toggles_via_config(self): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + self.assertIn("explore_dataframe", self.ip.user_ns) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = False" + ) + + self.assertNotIn("explore_dataframe", self.ip.user_ns) + self.assertNotIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertNotIn("sparksql", self.ip.magics_manager.magics["cell"]) + + with mock.patch("sys.stdout", io.StringIO()): + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = True" + ) + + self.assertIn("explore_dataframe", self.ip.user_ns) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + + def test_real_shell_config_magic_pre_and_post_import(self): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + # 1. Opt out via %config BEFORE _init_extras() + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = False" + ) + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + self.assertNotIn("explore_dataframe", self.ip.user_ns) + self.assertNotIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertNotIn("sparksql", self.ip.magics_manager.magics["cell"]) + + # 2. Enable via %config AFTER _init_extras() + with mock.patch("sys.stdout", io.StringIO()): + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = True" + ) + + self.assertIn("explore_dataframe", self.ip.user_ns) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + + # 3. Disable via %config AFTER _init_extras() + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = False" + ) + + self.assertNotIn("explore_dataframe", self.ip.user_ns) + self.assertNotIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertNotIn("sparksql", self.ip.magics_manager.magics["cell"]) + + def test_real_shell_partial_undo_preserves_preloaded_extension(self): + import io + from google.cloud.managed_spark_connect._ipython import _init_extras + + # User manually loads %dpip before managed_spark_connect initializes extras + self.ip.extension_manager.load_extension( + "google.cloud.managed_spark_magics" + ) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertNotIn("sparksql", self.ip.magics_manager.magics["cell"]) + + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + self.assertIn("explore_dataframe", self.ip.user_ns) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + + # Disabling extras via %config must keep %dpip (pre-loaded by user) while removing explore_dataframe and %%sparksql + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = False" + ) + + self.assertNotIn("explore_dataframe", self.ip.user_ns) + self.assertNotIn("sparksql", self.ip.magics_manager.magics["cell"]) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + + def test_real_shell_config_magic_overrides_env_opt_out(self): + import io + import os + from google.cloud.managed_spark_connect._ipython import ( + ENV_ENABLE_EXTRAS, + _init_extras, + ) + + # %config is evaluated before the import, so the merged shell config + # must beat the environment variable rather than being vetoed by it. + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = True" + ) + with mock.patch.dict(os.environ, {ENV_ENABLE_EXTRAS: "false"}): + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + self.assertIn("explore_dataframe", self.ip.user_ns) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + + def test_real_shell_env_opt_out_then_enable_via_config_magic(self): + import io + import os + from google.cloud.managed_spark_connect._ipython import ( + ENV_ENABLE_EXTRAS, + _init_extras, + ) + + with mock.patch.dict(os.environ, {ENV_ENABLE_EXTRAS: "false"}): + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + self.assertNotIn("explore_dataframe", self.ip.user_ns) + + # Turning the trait on must actually load, and the trait must + # report the same thing the shell is really in. + with mock.patch("sys.stdout", io.StringIO()): + self.ip.run_line_magic( + "config", "ManagedSparkConnect.enable_extras = True" + ) + + self.assertIn("explore_dataframe", self.ip.user_ns) + self.assertIn("dpip", self.ip.magics_manager.magics["line"]) + self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + + +_IMPORT_ORDER_PROBE = """ +import io, sys +from contextlib import redirect_stdout +from IPython.core.interactiveshell import InteractiveShell + +ip = InteractiveShell.instance() +with redirect_stdout(io.StringIO()): + if sys.argv[1] == "magics_first": + import google.cloud.managed_spark_magics + elif sys.argv[1] == "load_ext_first": + ip.run_line_magic("load_ext", "google.cloud.managed_spark_magics") + import google.cloud.managed_spark_connect + +print("explore_dataframe" in ip.user_ns) +print("dpip" in ip.magics_manager.magics["line"]) +print("sparksql" in ip.magics_manager.magics["cell"]) +""" + + +class TestExtrasImportOrder(unittest.TestCase): + """Each import order must end up with the same set of extras loaded. + + Has to run out of process: once a module is in sys.modules, the import + order that produced it cannot be replayed. + """ + + def _run(self, order): + import os + import subprocess + import sys + + env = dict(os.environ) + env.pop("MANAGED_SPARK_CONNECT_ENABLE_EXTRAS", None) + result = subprocess.run( + [sys.executable, "-c", _IMPORT_ORDER_PROBE, order], + capture_output=True, + text=True, + env=env, + ) + self.assertEqual(result.returncode, 0, result.stderr) + return result.stdout.split() + + def test_all_import_orders_load_all_extras(self): + for order in ("connect_first", "magics_first", "load_ext_first"): + with self.subTest(order=order): + # Importing managed_spark_magics first used to re-enter + # _init_extras() while that module was still initializing, + # which silently dropped %dpip. + self.assertEqual( + self._run(order), + ["True", "True", "True"], + "expected explore_dataframe, %dpip and %%sparksql", + ) + + if __name__ == "__main__": unittest.main() From 096ffec2b96bbf9d341549bb2a6310b0aa97d28b Mon Sep 17 00:00:00 2001 From: Dave Borowitz Date: Tue, 29 Sep 2026 14:47:24 -0700 Subject: [PATCH 2/2] Move interactive extras dependencies behind [interactive] extra google-colabsqlviz, ipython, sparksql-magic and traitlets are now only installed with google-cloud-spark-connect[interactive], so that users who don't use notebooks don't pay for them. For consistency, the feature is now called "interactive extras" rather than "notebook extras" throughout. At runtime, each extra is initialized only if its dependency is installed; one that isn't installed is skipped silently, and only installed-but-broken ones produce the "Failed to load interactive extras" message. The package itself only imports _ipython when traitlets is available. Since the set of extras now depends on what happens to be installed, replace the explore_dataframe()-specific message with a single line listing every extra that was loaded, e.g.: Loaded interactive extras: explore_dataframe(), %dpip, %%sparksql. To disable: %config ManagedSparkConnect.enable_extras = False Extras that were already present (e.g. loaded by the user) are not listed. --- DEVELOPING.md | 12 +- README.md | 16 +- .../cloud/managed_spark_connect/__init__.py | 16 +- .../cloud/managed_spark_connect/_ipython.py | 90 ++++++-- setup.py | 20 +- tests/unit/test_init.py | 214 +++++++++++++++++- 6 files changed, 321 insertions(+), 47 deletions(-) diff --git a/DEVELOPING.md b/DEVELOPING.md index 36ca1fa3..6795d8ed 100644 --- a/DEVELOPING.md +++ b/DEVELOPING.md @@ -39,14 +39,16 @@ env \ pytest --tb=auto -v ``` -## Testing the Notebook Extras +## Testing the Interactive Extras -The notebook extras (`explore_dataframe`, `%dpip`, `%%sparksql`) are regular -`install_requires` dependencies, so `pip install .` is enough and there is no -"without magic support" configuration to test. The tests in +The dependencies of the interactive extras (`explore_dataframe`, `%dpip`, +`%%sparksql`) live in the `interactive` extra +(`pip install '.[interactive]'`), and `requirements-dev.txt` already includes +them. The tests in `tests/unit/test_init.py` assume `google-colabsqlviz`, `sparksql-magic`, `ipython` and `traitlets` are importable and will fail, not skip, if they are -not. +not. Missing dependencies are simulated within the tests themselves, by +setting `sys.modules[name] = None`. The integration tests in particular can take a while to run. To speed up the testing cycle, you can run them in parallel. You can do so using the `xdist` diff --git a/README.md b/README.md index f94de49c..b635a39f 100644 --- a/README.md +++ b/README.md @@ -121,10 +121,16 @@ To create or connect to a named session: 5. A session with a given ID that is in a TERMINATED state cannot be reused. It must be deleted before a new session with the same ID can be created. -### Jupyter Notebook Extras +### Interactive Extras -When you import the package inside an IPython kernel, it automatically sets up -a few interactive conveniences--no separate install or setup required: +For notebooks and other interactive use, install the `interactive` extra: + +```sh +pip install 'google-cloud-spark-connect[interactive]' +``` + +When you import the package inside an IPython kernel, it then automatically sets +up a few interactive conveniences: - `explore_dataframe()` from [google-colabsqlviz](https://pypi.org/project/google-colabsqlviz/) is injected @@ -137,6 +143,10 @@ a few interactive conveniences--no separate install or setup required: import google.cloud.managed_spark_connect # extras load here ``` +Each extra is only set up if its dependency is importable, so without the +`interactive` extra you get whichever of them your environment already happens +to provide. On import, a single line lists the extras that were loaded. + The extras won't override anything you've already set up, for example if `explore_dataframe` is already present, or `%%sparksql` magic is loaded from somewhere else, these are left alone. #### Opting out diff --git a/google/cloud/managed_spark_connect/__init__.py b/google/cloud/managed_spark_connect/__init__.py index 68366d35..c1d9a338 100644 --- a/google/cloud/managed_spark_connect/__init__.py +++ b/google/cloud/managed_spark_connect/__init__.py @@ -12,12 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. import importlib.metadata +import importlib.util import warnings -from ._ipython import ( - ManagedSparkConnect, - _init_extras, -) from .session import ManagedSparkSession old_package_names = ["google-spark-connect", "dataproc-spark-connect"] @@ -33,4 +30,13 @@ except Exception: pass -_init_extras() +# traitlets, like every other dependency of the interactive extras, only comes +# with the [interactive] extra. IPython depends on it, so without it there can +# be no shell and nothing to initialize. +if importlib.util.find_spec("traitlets") is not None: + from ._ipython import ( + ManagedSparkConnect, + _init_extras, + ) + + _init_extras() diff --git a/google/cloud/managed_spark_connect/_ipython.py b/google/cloud/managed_spark_connect/_ipython.py index fb9ce181..55f45d04 100644 --- a/google/cloud/managed_spark_connect/_ipython.py +++ b/google/cloud/managed_spark_connect/_ipython.py @@ -13,10 +13,11 @@ # limitations under the License. from dataclasses import dataclass, field +import importlib.util import logging import os import sys -from typing import Any, Dict, Optional, Set +from typing import Any, Dict, List, Optional, Set from traitlets import Bool, default, observe from traitlets.config import Config, Configurable @@ -27,12 +28,34 @@ _ENV_FALSE_VALUES = frozenset(("0", "false", "no", "off")) +# Module that must be installed for explore_dataframe() to be injected. +_COLABSQLVIZ_MODULE = "google.colabsqlviz" + +# (extension module, magic kind, magic name). Each extension is only loaded if +# its module is installed. _EXTRAS_EXTENSIONS = ( ("google.cloud.managed_spark_magics", "line", "dpip"), ("sparksql_magic", "cell", "sparksql"), ) +def _is_module_available(name: str) -> bool: + """Whether `name` is installed, without importing it. + + All of the extras' dependencies are optional (they come from the + [interactive] extra), so a module that isn't installed means that extra is + deliberately absent, not broken. Anything other than a clean "not found" is + reported as available so that the subsequent import surfaces the real error. + """ + try: + return importlib.util.find_spec(name) is not None + except ModuleNotFoundError: + # A parent package is missing. + return False + except Exception: + return True + + def _is_env_extras_enabled() -> bool: val = os.getenv(ENV_ENABLE_EXTRAS) if val is None or not val.strip(): @@ -50,7 +73,7 @@ class ManagedSparkConnect(Configurable): # always wins. enable_extras = Bool( help=( - "Whether to automatically load notebook extras " + "Whether to automatically load the installed interactive extras " "(explore_dataframe, %dpip, %%sparksql)." ), ).tag(config=True) @@ -183,7 +206,7 @@ def _print_extras_failure() -> None: """ print( "\033[94m⚠️ [google.cloud.managed_spark_connect]\033[0m" - " Failed to load notebook extras." + " Failed to load interactive extras." " For details, enable debug logging and restart the kernel." ) @@ -212,15 +235,44 @@ def _is_extension_loaded( return False +def _print_extras_loaded(loaded: List[str]) -> None: + """Lists the extras that were just loaded, in a single line. + + Which extras get loaded depends on what happens to be installed, which the + user may not know (or may have forgotten), so name them explicitly rather + than leaving them to be discovered. + """ + names = ", ".join(f"\033[1m{name}\033[0m" for name in loaded) + print( + "\033[94m👉 [google.cloud.managed_spark_connect]\033[0m" + f" Loaded interactive extras: {names}." + " To disable: %config ManagedSparkConnect.enable_extras = False" + ) + + +def _magic_display_name(magic_kind: str, magic_name: str) -> str: + prefix = "%%" if magic_kind == "cell" else "%" + return f"{prefix}{magic_name}" + + def _load_extras_for_shell(ip: Any) -> None: state = _get_shell_state(ip) target_name = "explore_dataframe" failed = False + # Only what this call loaded; anything that was already present, whether + # from us or the user, is not news. + loaded: List[str] = [] # 1. Inject explore_dataframe from google-colabsqlviz try: user_ns = getattr(ip, "user_ns", None) - if isinstance(user_ns, dict): + if not _is_module_available(_COLABSQLVIZ_MODULE): + logger.debug( + "Not injecting %s: %s is not installed", + target_name, + _COLABSQLVIZ_MODULE, + ) + elif isinstance(user_ns, dict): if ( state.injected_explore_dataframe is not None and user_ns.get(target_name) is state.injected_explore_dataframe @@ -233,13 +285,7 @@ def _load_extras_for_shell(ip: Any) -> None: ip.push({target_name: explore_dataframe}) state.injected_explore_dataframe = explore_dataframe - - msg = ( - "\033[94m👉 [google.cloud.managed_spark_connect]\033[0m" - f" Injected \033[1m{target_name}()\033[0m into globals." - f" Use \033[1m{target_name}(df)\033[0m to interactively explore your data." - ) - print(msg) + loaded.append(f"{target_name}()") else: current_val = user_ns[target_name] ours = _import_explore_dataframe() @@ -263,17 +309,25 @@ def _load_extras_for_shell(ip: Any) -> None: ) failed = True - # 2. Load magic extensions silently + # 2. Load magic extensions ext_mgr = getattr(ip, "extension_manager", None) if ext_mgr is not None and hasattr(ext_mgr, "load_extension"): for ext_name, magic_kind, magic_name in _EXTRAS_EXTENSIONS: try: - if ext_name not in state.loaded_extensions and not ( + if ext_name in state.loaded_extensions or ( _is_extension_loaded(ip, ext_name, magic_kind, magic_name) ): - res = ext_mgr.load_extension(ext_name) - if res is None: - state.loaded_extensions.add(ext_name) + continue + if not _is_module_available(ext_name): + logger.debug( + "Not loading IPython extension %s: not installed", + ext_name, + ) + continue + res = ext_mgr.load_extension(ext_name) + if res is None: + state.loaded_extensions.add(ext_name) + loaded.append(_magic_display_name(magic_kind, magic_name)) except Exception: logger.debug( "Failed to load IPython extension %s", @@ -282,6 +336,8 @@ def _load_extras_for_shell(ip: Any) -> None: ) failed = True + if loaded: + _print_extras_loaded(loaded) if failed: _print_extras_failure() @@ -342,7 +398,7 @@ def _unload_extras_for_shell(ip: Any) -> None: def _init_extras() -> None: - """Initialize notebook extras on package import unless opted out.""" + """Initialize interactive extras on package import unless opted out.""" ip = None try: from IPython import get_ipython diff --git a/setup.py b/setup.py index 5b45fdcc..c8c46190 100644 --- a/setup.py +++ b/setup.py @@ -31,16 +31,22 @@ install_requires=[ "google-api-core>=2.19", "google-cloud-dataproc>=5.18", - "google-colabsqlviz>=0.3.0", - # Imported directly by managed_spark_connect._ipython and - # managed_spark_magics; previously these only arrived transitively via - # google-colabsqlviz and sparksql-magic. - "ipython>=8.0", "packaging>=20.0", "pyspark[connect]~=4.0.0", - "sparksql-magic>=0.0.3", "tqdm>=4.67", - "traitlets>=5.1", "websockets>=14.0", ], + extras_require={ + # Everything here is optional at runtime: managed_spark_connect._ipython + # only initializes the extras whose dependencies are importable. + "interactive": [ + "google-colabsqlviz>=0.3.0", + # Imported directly by managed_spark_connect._ipython and + # managed_spark_magics; previously these only arrived transitively + # via google-colabsqlviz and sparksql-magic. + "ipython>=8.0", + "sparksql-magic>=0.0.3", + "traitlets>=5.1", + ], + }, ) diff --git a/tests/unit/test_init.py b/tests/unit/test_init.py index c3eb2a92..81bf363f 100644 --- a/tests/unit/test_init.py +++ b/tests/unit/test_init.py @@ -231,10 +231,13 @@ def test_successful_extras_loading(self, mock_get_ipython): ) output = stdout_capture.getvalue() self.assertIn("[google.cloud.managed_spark_connect]", output) - self.assertIn("Injected", output) + self.assertIn("Loaded interactive extras", output) self.assertIn("explore_dataframe()", output) - self.assertNotIn("dpip", output) - self.assertNotIn("sparksql", output) + self.assertIn("%dpip", output) + self.assertIn("%%sparksql", output) + self.assertIn("ManagedSparkConnect.enable_extras = False", output) + # One summary line, not one per extra. + self.assertEqual(len(output.splitlines()), 1) @mock.patch("IPython.get_ipython") def test_fallback_when_already_defined_short_repr(self, mock_get_ipython): @@ -334,9 +337,13 @@ def test_no_warning_when_already_defined_same_function( _init_extras() # The name already refers to the exact function we would have injected, - # so there is nothing to warn the user about. + # so there is nothing to warn the user about, and since we didn't load + # it, it isn't listed either. mock_ip.push.assert_not_called() - self.assertEqual(stdout_capture.getvalue(), "") + output = stdout_capture.getvalue() + self.assertNotIn("Did not inject", output) + self.assertNotIn("explore_dataframe", output) + self.assertIn("Loaded interactive extras", output) self.assertIs(mock_ip.user_ns["explore_dataframe"], explore_dataframe) @mock.patch("IPython.get_ipython") @@ -498,13 +505,13 @@ def test_failure_message_printed_once(self, mock_get_ipython): _init_extras() output = stdout_capture.getvalue() - self.assertIn("Failed to load notebook extras", output) + self.assertIn("Failed to load interactive extras", output) # Three separate failures (the injection and both extensions) must not # produce three separate messages. - self.assertEqual(output.count("Failed to load notebook extras"), 1) + self.assertEqual(output.count("Failed to load interactive extras"), 1) @mock.patch("IPython.get_ipython") - def test_failure_message_when_colabsqlviz_is_missing( + def test_failure_message_when_colabsqlviz_fails_to_import( self, mock_get_ipython ): import builtins @@ -526,9 +533,10 @@ def fail_colabsqlviz(name, *args, **kwargs): with mock.patch("sys.stdout", stdout_capture): _init_extras() + # colabsqlviz is installed but broken, which is worth reporting. self.assertNotIn("explore_dataframe", mock_ip.user_ns) self.assertIn( - "Failed to load notebook extras", stdout_capture.getvalue() + "Failed to load interactive extras", stdout_capture.getvalue() ) # The magics are independent of colabsqlviz and should still load. mock_ip.extension_manager.load_extension.assert_has_calls( @@ -538,6 +546,146 @@ def fail_colabsqlviz(name, *args, **kwargs): ] ) + @mock.patch("IPython.get_ipython") + def test_colabsqlviz_not_installed_is_skipped(self, mock_get_ipython): + import io + import sys + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + # A None entry in sys.modules makes both find_spec() and import behave + # as if the module were not installed. + with mock.patch.dict(sys.modules, {"google.colabsqlviz": None}): + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + mock_ip.push.assert_not_called() + self.assertNotIn("explore_dataframe", mock_ip.user_ns) + mock_ip.extension_manager.load_extension.assert_has_calls( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + # Only what was actually loaded is listed, and a missing optional + # dependency is not a failure. + output = stdout_capture.getvalue() + self.assertEqual(len(output.splitlines()), 1) + self.assertIn("Loaded interactive extras", output) + self.assertIn("%dpip", output) + self.assertIn("%%sparksql", output) + self.assertNotIn("explore_dataframe", output) + self.assertNotIn("Failed to load interactive extras", output) + + @mock.patch("IPython.get_ipython") + def test_colabsqlviz_not_installed_does_not_warn_about_existing_name( + self, mock_get_ipython + ): + import io + import sys + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell(user_ns={"explore_dataframe": 42}) + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch.dict(sys.modules, {"google.colabsqlviz": None}): + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + # We would not have injected anything, so there is no conflict to + # warn about. + self.assertNotIn("Did not inject", stdout_capture.getvalue()) + self.assertEqual(mock_ip.user_ns["explore_dataframe"], 42) + + @mock.patch("IPython.get_ipython") + def test_sparksql_magic_not_installed_is_skipped(self, mock_get_ipython): + import io + import sys + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch.dict(sys.modules, {"sparksql_magic": None}): + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + self.assertIn("explore_dataframe", mock_ip.user_ns) + mock_ip.extension_manager.load_extension.assert_called_once_with( + "google.cloud.managed_spark_magics" + ) + output = stdout_capture.getvalue() + self.assertEqual(len(output.splitlines()), 1) + self.assertIn("explore_dataframe()", output) + self.assertIn("%dpip", output) + self.assertNotIn("sparksql", output) + self.assertNotIn("Failed to load interactive extras", output) + + @mock.patch("IPython.get_ipython") + def test_nothing_printed_when_nothing_is_loaded(self, mock_get_ipython): + import io + import sys + from google.cloud.managed_spark_connect._ipython import _init_extras + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + stdout_capture = io.StringIO() + with mock.patch.dict( + sys.modules, + { + "google.colabsqlviz": None, + "google.cloud.managed_spark_magics": None, + "sparksql_magic": None, + }, + ): + with mock.patch("sys.stdout", stdout_capture): + _init_extras() + + mock_ip.extension_manager.load_extension.assert_not_called() + self.assertEqual(stdout_capture.getvalue(), "") + + @mock.patch("IPython.get_ipython") + def test_summary_not_repeated_on_second_load(self, mock_get_ipython): + import io + from google.cloud.managed_spark_connect._ipython import ( + _init_extras, + _load_extras_for_shell, + ) + + mock_ip = _make_mock_shell() + mock_get_ipython.return_value = mock_ip + + with mock.patch("sys.stdout", io.StringIO()): + _init_extras() + + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): + _load_extras_for_shell(mock_ip) + + # Everything is already loaded, so there is nothing new to report. + self.assertEqual(stdout_capture.getvalue(), "") + + def test_is_module_available(self): + from google.cloud.managed_spark_connect._ipython import ( + _is_module_available, + ) + + self.assertTrue(_is_module_available("json")) + self.assertTrue(_is_module_available("google.colabsqlviz")) + self.assertFalse(_is_module_available("no_such_module_for_test")) + # Missing leaf in an existing (namespace) package. + self.assertFalse(_is_module_available("google.no_such_module_for_test")) + # Missing parent package. + self.assertFalse( + _is_module_available("no_such_module_for_test.submodule") + ) + def test_module_import_calls_init_extras(self): import importlib import google.cloud.managed_spark_connect as msc @@ -687,12 +835,18 @@ def test_real_shell_partial_undo_preserves_preloaded_extension(self): self.assertIn("dpip", self.ip.magics_manager.magics["line"]) self.assertNotIn("sparksql", self.ip.magics_manager.magics["cell"]) - with mock.patch("sys.stdout", io.StringIO()): + stdout_capture = io.StringIO() + with mock.patch("sys.stdout", stdout_capture): _init_extras() self.assertIn("explore_dataframe", self.ip.user_ns) self.assertIn("dpip", self.ip.magics_manager.magics["line"]) self.assertIn("sparksql", self.ip.magics_manager.magics["cell"]) + # %dpip was the user's doing, so the summary must not claim it. + output = stdout_capture.getvalue() + self.assertIn("explore_dataframe()", output) + self.assertIn("%%sparksql", output) + self.assertNotIn("%dpip", output) # Disabling extras via %config must keep %dpip (pre-loaded by user) while removing explore_dataframe and %%sparksql self.ip.run_line_magic( @@ -804,5 +958,45 @@ def test_all_import_orders_load_all_extras(self): ) +_NO_INTERACTIVE_DEPS_PROBE = """ +import sys + +# Simulate an install without the [interactive] extra. A None entry makes both +# importlib.util.find_spec() and import treat the module as absent. +for name in ("IPython", "traitlets", "google.colabsqlviz", "sparksql_magic"): + sys.modules[name] = None + +import google.cloud.managed_spark_connect as msc + +print(msc.ManagedSparkSession is not None) +print(hasattr(msc, "ManagedSparkConnect")) +print("google.cloud.managed_spark_connect._ipython" in sys.modules) +""" + + +class TestWithoutInteractiveExtra(unittest.TestCase): + """The base package must not need anything from the [interactive] extra. + + Runs out of process so that nothing already imported by the test runner + (IPython in particular) leaks in. + """ + + def test_import_without_interactive_deps(self): + import os + import subprocess + import sys + + env = dict(os.environ) + env.pop("MANAGED_SPARK_CONNECT_ENABLE_EXTRAS", None) + result = subprocess.run( + [sys.executable, "-c", _NO_INTERACTIVE_DEPS_PROBE], + capture_output=True, + text=True, + env=env, + ) + self.assertEqual(result.returncode, 0, result.stderr) + self.assertEqual(result.stdout.split(), ["True", "False", "False"]) + + if __name__ == "__main__": unittest.main()