diff --git a/DEVELOPING.md b/DEVELOPING.md index c9a8a5e..6795d8e 100644 --- a/DEVELOPING.md +++ b/DEVELOPING.md @@ -39,27 +39,16 @@ 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 Interactive Extras + +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. 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 c111792..b635a39 100644 --- a/README.md +++ b/README.md @@ -121,53 +121,72 @@ 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) +### Interactive Extras -The package supports the [sparksql-magic](https://github.com/cryeo/sparksql-magic) library for executing Spark SQL queries directly in Jupyter notebooks. +For notebooks and other interactive use, install the `interactive` extra: -**Installation**: To use magic commands, install the required dependencies manually: -```bash -pip install google-cloud-spark-connect -pip install IPython sparksql-magic +```sh +pip install 'google-cloud-spark-connect[interactive]' ``` -1. Load the magic extension: - ```python - %load_ext sparksql_magic - ``` +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 + 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 +``` + +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 + +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`: -2. Configure default settings (optional): +```python +c.ManagedSparkConnect.enable_extras = False +``` + +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 2356758..c1d9a33 100644 --- a/google/cloud/managed_spark_connect/__init__.py +++ b/google/cloud/managed_spark_connect/__init__.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. import importlib.metadata +import importlib.util import warnings from .session import ManagedSparkSession @@ -28,3 +29,14 @@ ) except Exception: pass + +# 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 new file mode 100644 index 0000000..55f45d0 --- /dev/null +++ b/google/cloud/managed_spark_connect/_ipython.py @@ -0,0 +1,424 @@ +# 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 importlib.util +import logging +import os +import sys +from typing import Any, Dict, List, 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")) + +# 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(): + # 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 the installed interactive 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 interactive 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 _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 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 + ): + 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 + loaded.append(f"{target_name}()") + 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 + 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 in state.loaded_extensions or ( + _is_extension_loaded(ip, ext_name, magic_kind, magic_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", + ext_name, + exc_info=True, + ) + failed = True + + if loaded: + _print_extras_loaded(loaded) + 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 interactive 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 79632f5..eb66725 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 54363ae..38ca3c7 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 5cf7026..0980639 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 89637d5..c8c4619 100644 --- a/setup.py +++ b/setup.py @@ -36,4 +36,17 @@ "tqdm>=4.67", "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 794f6fa..81bf363 100644 --- a/tests/unit/test_init.py +++ b/tests/unit/test_init.py @@ -135,5 +135,868 @@ 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("Loaded interactive extras", output) + self.assertIn("explore_dataframe()", 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): + 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, and since we didn't load + # it, it isn't listed either. + mock_ip.push.assert_not_called() + 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") + 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 interactive extras", output) + # Three separate failures (the injection and both extensions) must not + # produce three separate messages. + self.assertEqual(output.count("Failed to load interactive extras"), 1) + + @mock.patch("IPython.get_ipython") + def test_failure_message_when_colabsqlviz_fails_to_import( + 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() + + # colabsqlviz is installed but broken, which is worth reporting. + self.assertNotIn("explore_dataframe", mock_ip.user_ns) + self.assertIn( + "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( + [ + mock.call("google.cloud.managed_spark_magics"), + mock.call("sparksql_magic"), + ] + ) + + @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 + + 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"]) + + 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( + "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", + ) + + +_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()