diff --git a/google/cloud/managed_spark_connect/client/proxy.py b/google/cloud/managed_spark_connect/client/proxy.py index cf68043..0aae5fa 100755 --- a/google/cloud/managed_spark_connect/client/proxy.py +++ b/google/cloud/managed_spark_connect/client/proxy.py @@ -18,11 +18,13 @@ import logging import socket import threading +import time import websockets.sync.client as websocketclient from google import auth as googleauth from google.auth.transport import requests as googleauthrequests +from google.cloud.managed_spark_connect import execution_timer parser = argparse.ArgumentParser() parser.add_argument("port") @@ -48,12 +50,27 @@ def recv(self, buff_size): # # We set that timeout to 60 seconds to prevent any scenarios where we wind up stuck waiting for a message from a websocket connection # that never comes. - msg = self._conn.recv(timeout=60) - return bytes.fromhex(msg) + start = time.monotonic() + nbytes = 0 + try: + msg = self._conn.recv(timeout=60) + result = bytes.fromhex(msg) + nbytes = len(result) + return result + finally: + execution_timer.record_transport( + "down", nbytes, time.monotonic() - start + ) def send(self, msg_bytes): - msg = bytes.hex(msg_bytes) - self._conn.send(msg) + start = time.monotonic() + try: + msg = bytes.hex(msg_bytes) + self._conn.send(msg) + finally: + execution_timer.record_transport( + "up", len(msg_bytes), time.monotonic() - start + ) def close(self): return self._conn.close() diff --git a/google/cloud/managed_spark_connect/execution_timer.py b/google/cloud/managed_spark_connect/execution_timer.py new file mode 100644 index 0000000..61608f2 --- /dev/null +++ b/google/cloud/managed_spark_connect/execution_timer.py @@ -0,0 +1,269 @@ +# 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. + +import logging +import threading +import time + +logger = logging.getLogger(__name__) + +_MIB = 1024 * 1024 +_KIB = 1024 + +# How many of the slowest round trips to name individually in the log +# line before collapsing the rest into a "+N more" tail. +_MAX_ROUND_TRIPS_SHOWN = 3 + +_ZERO_TOTALS = { + "session_creation_seconds": 0.0, + "bytes_down": 0, + "bytes_up": 0, + "transport_blocked": 0.0, +} + +_lock = threading.Lock() +_cell_start = None +_registered = False +_totals = dict(_ZERO_TOTALS) +# (kind, seconds) for every round trip to the Spark Connect endpoint in +# the current cell, where kind is "execute", "fetch", or "analyze". +# Bounded by one cell's worth of activity: start_cell() clears it, and +# nothing is appended at all outside IPython. +_round_trips = [] + + +def start_cell(): + """Marks the start of a cell, discarding any prior measurements.""" + global _cell_start + with _lock: + _cell_start = time.monotonic() + _totals.update(_ZERO_TOTALS) + _round_trips.clear() + + +def record(kind, start, end): + """Adds one Managed Spark round trip to the current cell. + + ``kind`` is one of ``"execute"``, ``"fetch"``, or ``"analyze"`` and + is recorded verbatim alongside the duration; it is not validated, + since every caller is this module's own package. Dropped when no + cell is in progress, which keeps this inert and bounded outside + IPython, where nothing ever resets the totals. + """ + with _lock: + if _cell_start is None or end <= _cell_start: + return + _round_trips.append((kind, end - max(start, _cell_start))) + + +def record_session_creation(seconds): + """Adds cold session provisioning time to the current cell. + + Takes an already-computed duration rather than a start/end pair, + since the caller measures it around a long provisioning loop. + Dropped when no cell is in progress. + """ + with _lock: + if _cell_start is None: + return + _totals["session_creation_seconds"] += seconds + + +def record_transport(direction, nbytes, blocked): + """Adds one websocket frame's transport accounting to the cell. + + Called from the bridge's hot path once per frame, so this stays + cheap: one lock acquire and dict updates, nothing else. An + unknown ``direction`` is ignored rather than raising. Dropped + when no cell is in progress. + """ + with _lock: + if _cell_start is None: + return + if direction == "down": + _totals["bytes_down"] += nbytes + elif direction == "up": + _totals["bytes_up"] += nbytes + else: + return + _totals["transport_blocked"] += blocked + + +def summary(): + """Returns a dict of every accumulated total plus ``cell_seconds``. + + ``managed_spark_seconds`` is the sum of every round trip's + duration plus ``session_creation_seconds``: provisioning a + Managed Spark session is time spent in Managed Spark, same as + any round trip. Both ``managed_spark_seconds`` and + ``transport_blocked`` are clamped to at most ``cell_seconds``. + When no cell is in progress, every value is zero and + ``round_trips`` is empty. + """ + with _lock: + if _cell_start is None: + result = dict(_ZERO_TOTALS) + result["round_trips"] = [] + result["managed_spark_seconds"] = 0.0 + result["cell_seconds"] = 0.0 + return result + cell_seconds = time.monotonic() - _cell_start + result = dict(_totals) + result["round_trips"] = list(_round_trips) + managed_spark_seconds = ( + sum(seconds for _, seconds in result["round_trips"]) + + result["session_creation_seconds"] + ) + result["managed_spark_seconds"] = min(managed_spark_seconds, cell_seconds) + # The bridge forwards bytes on multiple daemon threads, and each + # one accumulates its blocked time into this same counter, so + # concurrent blocking double counts and can exceed the wall time + # of the cell that contains it. Clamp for the same reason as + # managed_spark_seconds above. + result["transport_blocked"] = min(result["transport_blocked"], cell_seconds) + result["cell_seconds"] = cell_seconds + return result + + +def register_cell_timing(): + """Registers IPython cell hooks that log per-cell Managed Spark timing. + + Returns False (and registers nothing) outside of IPython, or if + called more than once. Handlers never raise into the notebook. + """ + global _registered + + if _registered: + return False + + try: + from IPython import get_ipython + + shell = get_ipython() + except ImportError: + return False + + if shell is None: + return False + + shell.events.register("pre_run_cell", _on_pre_run_cell) + shell.events.register("post_run_cell", _on_post_run_cell) + _registered = True + + if _cell_start is None: + # Registration happens mid-cell on the very first cell of a + # notebook: getOrCreate() calls this before pre_run_cell has + # ever fired, since no handler existed yet to fire it into. + # Without starting the cell here, the session creation cost + # about to be recorded would find no cell in progress and be + # silently dropped. Known imprecision: cell_seconds for that + # first cell is measured from registration rather than the + # cell's true start, so it slightly under-reports the cell's + # own wall time -- an acceptable trade for capturing the + # session-creation cost, which is the dominant term. + start_cell() + + return True + + +def _on_pre_run_cell(info=None): + try: + start_cell() + except Exception: + pass + + +def _format_round_trips_clause(round_trips): + """Names up to the three slowest round trips, slowest first. + + Any remaining round trips are collapsed into a trailing "+N more" + rather than named individually, so a cell with many small fetches + doesn't turn the log line into a wall of text. + """ + slowest_first = sorted(round_trips, key=lambda rt: rt[1], reverse=True) + shown = slowest_first[:_MAX_ROUND_TRIPS_SHOWN] + text = ", ".join("%s %.2fs" % (kind, seconds) for kind, seconds in shown) + extra = len(slowest_first) - len(shown) + if extra > 0: + text += ", +%d more" % extra + return text + + +def _format_detail_clauses(data): + """Builds the parenthesised detail clauses for the cell log line. + + Each clause is omitted when its underlying value is empty/zero, + so a cell with no session creation, round trips, or transport + activity yields an empty list. + """ + clauses = [] + + if data["session_creation_seconds"]: + clauses.append( + "session creation %.2fs" % data["session_creation_seconds"] + ) + + if data["round_trips"]: + clauses.append(_format_round_trips_clause(data["round_trips"])) + + bytes_down = data["bytes_down"] + if bytes_down: + blocked = data["transport_blocked"] + if bytes_down >= _MIB: + # Big enough for the rate to mean something. + mib_down = bytes_down / _MIB + rate = mib_down / blocked if blocked else 0.0 + clauses.append( + "transport %.1f MiB down in %.2fs, %.1f MiB/s" + % (mib_down, blocked, rate) + ) + else: + # A few KiB of control traffic would round to "0.0 MiB" + # and "0.0 MiB/s", which reads as broken. Report the + # smaller unit instead, and drop the rate: it would be + # meaningless at this scale, while the blocked time + # itself is still the useful signal (e.g. a long wait on + # a cold cluster that moved almost no data). + kib_down = bytes_down / _KIB + clauses.append( + "transport %.1f KiB down in %.2fs" % (kib_down, blocked) + ) + + return clauses + + +def _on_post_run_cell(result=None): + try: + data = summary() + round_trips = data["round_trips"] + # Session creation deliberately does not require any round + # trips (it happens before the Spark session exists), so + # without this a cell that only calls getOrCreate() would + # still stay silent. + if not round_trips and data["session_creation_seconds"] == 0: + return + message = "Cell took %.2fs, %.2fs in Managed Spark" % ( + data["cell_seconds"], + data["managed_spark_seconds"], + ) + if round_trips: + count = len(round_trips) + unit = "round trip" if count == 1 else "round trips" + message += " across %d %s" % (count, unit) + clauses = _format_detail_clauses(data) + if clauses: + message = "%s (%s)" % (message, "; ".join(clauses)) + logger.info(message) + except Exception: + pass diff --git a/google/cloud/managed_spark_connect/session.py b/google/cloud/managed_spark_connect/session.py index 4f76a49..babc909 100644 --- a/google/cloud/managed_spark_connect/session.py +++ b/google/cloud/managed_spark_connect/session.py @@ -42,6 +42,7 @@ from google.auth.exceptions import DefaultCredentialsError from google.cloud.managed_spark_connect.client import ManagedSparkChannelBuilder from google.cloud.managed_spark_connect.exceptions import ManagedSparkConnectException +from google.cloud.managed_spark_connect import execution_timer from google.cloud.managed_spark_connect.pypi_artifacts import PyPiArtifacts from google.cloud.dataproc_v1 import ( AuthenticationConfig, @@ -333,6 +334,9 @@ def __create_spark_connect_session_from_s8s( # Register handler for Cell Execution Progress bar session._register_progress_execution_handler() + # Register handlers for per-cell Managed Spark timing + execution_timer.register_cell_timing() + ManagedSparkSession._set_default_and_active_session(session) return session @@ -377,7 +381,7 @@ def __create(self) -> "ManagedSparkSession": ManagedSparkSession._active_session_uses_custom_id = ( self._custom_session_id is not None ) - s8s_creation_start_time = time.time() + s8s_creation_start_time = time.monotonic() stop_create_session_pbar_event = threading.Event() @@ -492,8 +496,12 @@ def create_session_pbar(): finally: stop_create_session_pbar_event.set() - logger.debug( - f"Managed Spark Session created: {session_id} in {int(time.time() - s8s_creation_start_time)} seconds" + s8s_creation_duration = ( + time.monotonic() - s8s_creation_start_time + ) + execution_timer.record_session_creation(s8s_creation_duration) + logger.info( + f"Managed Spark Session created: {session_id} in {int(s8s_creation_duration)} seconds" ) return self.__create_spark_connect_session_from_s8s( session_response, session_config.name @@ -624,6 +632,16 @@ def _get_exiting_active_session( def getOrCreate(self) -> "ManagedSparkSession": with ManagedSparkSession._lock: + # Must precede session creation: on the first cell of a + # notebook, this is what starts the cell clock, since + # pre_run_cell fired before any handler was registered + # to catch it. Without registering (and starting the + # cell) here first, the provisioning cost paid below + # would have no cell to be attributed to and would be + # silently dropped. Idempotent, so this is a no-op on + # every call after the first. + execution_timer.register_cell_timing() + if environment.is_dataproc_batch(): # For Dataproc batch workloads, connect to the already initialized local SparkSession from pyspark.sql import SparkSession as PySparkSQLSession @@ -989,6 +1007,7 @@ def __init__( execute_and_fetch_as_iterator_base_method = ( self.client._execute_and_fetch_as_iterator ) + analyze_base_method = self.client._analyze def execute_plan_request_wrapped_method(*args, **kwargs): req = execute_plan_request_base_method(*args, **kwargs) @@ -1006,7 +1025,11 @@ def execute_plan_request_wrapped_method(*args, **kwargs): def execute_wrapped_method(client_self, req, *args, **kwargs): if not self._sql_lazy_transformation(req): self._display_operation_link(req.operation_id) - execute_base_method(req, *args, **kwargs) + start = time.monotonic() + try: + execute_base_method(req, *args, **kwargs) + finally: + execution_timer.record("execute", start, time.monotonic()) self.client._execute = MethodType(execute_wrapped_method, self.client) @@ -1015,14 +1038,24 @@ def execute_and_fetch_as_iterator_wrapped_method( ): if not self._sql_lazy_transformation(req): self._display_operation_link(req.operation_id) - return execute_and_fetch_as_iterator_base_method( + iterator = execute_and_fetch_as_iterator_base_method( req, *args, **kwargs ) + return self._timed_iterator(iterator) self.client._execute_and_fetch_as_iterator = MethodType( execute_and_fetch_as_iterator_wrapped_method, self.client ) + def analyze_wrapped_method(client_self, method, **kwargs): + start = time.monotonic() + try: + return analyze_base_method(method, **kwargs) + finally: + execution_timer.record("analyze", start, time.monotonic()) + + self.client._analyze = MethodType(analyze_wrapped_method, self.client) + # Patching clearProgressHandlers method to not remove Managed Spark Progress Handler clearProgressHandlers_base_method = self.clearProgressHandlers @@ -1120,6 +1153,18 @@ def handler( self.registerProgressHandler(handler) + @staticmethod + def _timed_iterator(iterator): + # This being a generator is load-bearing: `start` is not read until + # the first next(), so the clock runs from when the caller begins + # driving the RPC rather than from when the generator was built. + # Hoisting it out of the generator body would time every fetch as 0s. + start = time.monotonic() + try: + yield from iterator + finally: + execution_timer.record("fetch", start, time.monotonic()) + @staticmethod def _sql_lazy_transformation(req): # Select SQL command diff --git a/tests/unit/test_execution_timer.py b/tests/unit/test_execution_timer.py new file mode 100644 index 0000000..62669ab --- /dev/null +++ b/tests/unit/test_execution_timer.py @@ -0,0 +1,739 @@ +# 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. + +import threading +import unittest +from unittest import mock + +from google.cloud.managed_spark_connect import execution_timer +from google.cloud.managed_spark_connect.execution_timer import ( + record, + record_session_creation, + record_transport, + register_cell_timing, + start_cell, + summary, +) + +_MODULE = "google.cloud.managed_spark_connect.execution_timer" + +_ZERO_SUMMARY = dict(execution_timer._ZERO_TOTALS) +_ZERO_SUMMARY["round_trips"] = [] +_ZERO_SUMMARY["managed_spark_seconds"] = 0.0 +_ZERO_SUMMARY["cell_seconds"] = 0.0 + + +def _start_cell_at(t): + with mock.patch(f"{_MODULE}.time.monotonic", return_value=t): + start_cell() + + +def _summary_at(t): + with mock.patch(f"{_MODULE}.time.monotonic", return_value=t): + return summary() + + +def _make_summary(**overrides): + """Builds a full summary()-shaped dict, zero except for overrides.""" + data = dict(_ZERO_SUMMARY) + data.update(overrides) + return data + + +def _registered_handler(shell, name): + """Returns the handler that was registered for ``name`` on ``shell``.""" + for call in shell.events.register.call_args_list: + if call.args[0] == name: + return call.args[1] + raise AssertionError(f"no handler registered for {name!r}") + + +class ExecutionTimerTests(unittest.TestCase): + """Saves and restores the module-level timing state around each test.""" + + def setUp(self): + self._orig_cell_start = execution_timer._cell_start + self._orig_totals = dict(execution_timer._totals) + self._orig_round_trips = list(execution_timer._round_trips) + self._orig_registered = execution_timer._registered + execution_timer._cell_start = None + execution_timer._totals.clear() + execution_timer._totals.update(execution_timer._ZERO_TOTALS) + execution_timer._round_trips.clear() + execution_timer._registered = False + + def tearDown(self): + execution_timer._cell_start = self._orig_cell_start + execution_timer._totals.clear() + execution_timer._totals.update(self._orig_totals) + execution_timer._round_trips[:] = self._orig_round_trips + execution_timer._registered = self._orig_registered + + def _assert_all_zero(self, data): + self.assertEqual(data, _ZERO_SUMMARY) + + # -- record() ---------------------------------------------------- + + def test_record_stores_kind_and_duration(self): + _start_cell_at(0.0) + record("execute", 1.0, 2.0) + record("fetch", 3.0, 4.5) + data = _summary_at(10.0) + self.assertEqual( + data["round_trips"], [("execute", 1.0), ("fetch", 1.5)] + ) + + def test_round_trips_len_is_the_count(self): + _start_cell_at(0.0) + record("execute", 1.0, 2.0) + record("fetch", 3.0, 4.0) + record("analyze", 5.0, 6.0) + self.assertEqual(len(_summary_at(10.0)["round_trips"]), 3) + + def test_disjoint_intervals_sum_into_managed_spark_seconds(self): + _start_cell_at(0.0) + record("execute", 1.0, 2.0) + record("fetch", 3.0, 4.0) + data = _summary_at(10.0) + self.assertEqual(data["cell_seconds"], 10.0) + self.assertEqual(data["managed_spark_seconds"], 2.0) + self.assertEqual(len(data["round_trips"]), 2) + + def test_summary_with_no_cell_started(self): + self._assert_all_zero(summary()) + + def test_record_before_start_cell_is_dropped(self): + record("execute", 1.0, 2.0) + _start_cell_at(5.0) + self.assertEqual(_summary_at(6.0)["round_trips"], []) + + def test_record_before_start_cell_stays_dropped_after_later_start(self): + record("execute", 1.0, 2.0) + _start_cell_at(5.0) + _summary_at(6.0) + _start_cell_at(10.0) + self.assertEqual(_summary_at(11.0)["round_trips"], []) + + def test_start_cell_clears_previous_cells_round_trips(self): + _start_cell_at(0.0) + record("execute", 1.0, 2.0) + _start_cell_at(10.0) + data = _summary_at(11.0) + self.assertEqual(data["managed_spark_seconds"], 0.0) + self.assertEqual(data["round_trips"], []) + + def test_interval_starting_before_cell_start_is_clipped(self): + _start_cell_at(10.0) + record("execute", 5.0, 15.0) + data = _summary_at(20.0) + self.assertEqual(data["cell_seconds"], 10.0) + self.assertEqual(data["managed_spark_seconds"], 5.0) + self.assertEqual(data["round_trips"], [("execute", 5.0)]) + + def test_interval_ending_at_cell_start_is_dropped(self): + _start_cell_at(10.0) + record("execute", 3.0, 10.0) + data = _summary_at(20.0) + self.assertEqual(data["managed_spark_seconds"], 0.0) + self.assertEqual(data["round_trips"], []) + + def test_interval_ending_before_cell_start_is_dropped(self): + _start_cell_at(10.0) + record("execute", 3.0, 8.0) + data = _summary_at(20.0) + self.assertEqual(data["managed_spark_seconds"], 0.0) + self.assertEqual(data["round_trips"], []) + + def test_managed_spark_seconds_never_exceeds_cell_seconds(self): + _start_cell_at(0.0) + record("execute", 0.5, 3.0) + record("fetch", 1.0, 4.5) + data = _summary_at(5.0) + self.assertLessEqual( + data["managed_spark_seconds"], data["cell_seconds"] + ) + + def test_overlapping_intervals_summed_and_clamped_to_cell_duration(self): + _start_cell_at(0.0) + record("execute", 0.0, 3.0) + record("fetch", 0.0, 3.0) + data = _summary_at(4.0) + self.assertEqual(data["cell_seconds"], 4.0) + self.assertEqual(data["managed_spark_seconds"], 4.0) + self.assertEqual(len(data["round_trips"]), 2) + + def test_concurrent_record_calls_lose_no_updates(self): + start_cell() + + num_threads = 20 + records_per_thread = 50 + + def worker(): + for _ in range(records_per_thread): + start = execution_timer._cell_start + 0.001 + end = start + 0.001 + record("execute", start, end) + + threads = [threading.Thread(target=worker) for _ in range(num_threads)] + for t in threads: + t.start() + for t in threads: + t.join() + + self.assertEqual( + len(summary()["round_trips"]), num_threads * records_per_thread + ) + + # -- record_session_creation() ------------------------------------ + + def test_session_creation_counted_toward_managed_spark_seconds(self): + _start_cell_at(0.0) + record_session_creation(64.7) + record("execute", 1.0, 2.0) + data = _summary_at(70.0) + self.assertEqual(data["session_creation_seconds"], 64.7) + self.assertEqual(data["managed_spark_seconds"], 65.7) + + def test_record_session_creation_applies_clipping_like_record(self): + _start_cell_at(10.0) + record_session_creation(500.0) + data = _summary_at(20.0) + # session creation is unconditional but managed_spark_seconds + # is still clamped to the cell's own wall time. + self.assertEqual(data["managed_spark_seconds"], 10.0) + + def test_record_session_creation_dropped_when_no_cell_in_progress(self): + record_session_creation(64.7) + self._assert_all_zero(summary()) + + # -- record_transport() ------------------------------------------- + + def test_record_transport_accumulates_down_and_up_separately(self): + _start_cell_at(0.0) + record_transport("down", 100, 0.5) + record_transport("down", 200, 0.25) + record_transport("up", 50, 0.1) + data = _summary_at(10.0) + self.assertEqual(data["bytes_down"], 300) + self.assertEqual(data["bytes_up"], 50) + self.assertAlmostEqual(data["transport_blocked"], 0.85) + + def test_record_transport_unknown_direction_ignored(self): + _start_cell_at(0.0) + try: + record_transport("sideways", 100, 0.5) + except Exception as exc: # pragma: no cover - failure path + self.fail(f"record_transport raised: {exc}") + data = _summary_at(10.0) + self.assertEqual(data["bytes_down"], 0) + self.assertEqual(data["bytes_up"], 0) + self.assertEqual(data["transport_blocked"], 0.0) + + def test_record_transport_dropped_when_no_cell_in_progress(self): + record_transport("down", 100, 0.5) + self._assert_all_zero(summary()) + + def test_transport_blocked_clamped_to_cell_seconds_in_summary(self): + """Concurrent forwarding threads each accumulate blocked time + into the same counter, so it can exceed the cell's own wall + time; summary() must clamp it the same way it clamps + managed_spark_seconds. + """ + _start_cell_at(0.0) + record_transport("down", 100, 10.0) + data = _summary_at(1.53) + self.assertEqual(data["cell_seconds"], 1.53) + self.assertEqual(data["transport_blocked"], 1.53) + + # -- registration --------------------------------------------------- + + @mock.patch("IPython.get_ipython", return_value=None) + def test_returns_false_and_registers_nothing_without_ipython( + self, mock_get_ipython + ): + self.assertFalse(register_cell_timing()) + + @mock.patch("IPython.get_ipython", side_effect=ImportError) + def test_returns_false_when_ipython_not_installed(self, mock_get_ipython): + self.assertFalse(register_cell_timing()) + + def test_first_call_true_second_call_false_one_handler_each(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + self.assertTrue(register_cell_timing()) + self.assertFalse(register_cell_timing()) + + pre_run_calls = [ + call + for call in shell.events.register.call_args_list + if call.args[0] == "pre_run_cell" + ] + post_run_calls = [ + call + for call in shell.events.register.call_args_list + if call.args[0] == "post_run_cell" + ] + self.assertEqual(len(pre_run_calls), 1) + self.assertEqual(len(post_run_calls), 1) + + def test_register_cell_timing_starts_cell_when_none_in_progress(self): + """Regression test for the first-cell bug: registration happens + mid-cell (getOrCreate() calls it before pre_run_cell has ever + fired), so registering must itself start the cell clock or the + rest of that cell is unmeasurable. + """ + self.assertIsNone(execution_timer._cell_start) + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + self.assertIsNotNone(execution_timer._cell_start) + + def test_register_cell_timing_does_not_restart_live_cell(self): + """A cell already in progress (pre_run_cell already fired, e.g. + on the second+ cell of a notebook) must not have its totals + clobbered by registration. + """ + _start_cell_at(0.0) + record("execute", 1.0, 2.0) + + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + data = _summary_at(3.0) + self.assertEqual(len(data["round_trips"]), 1) + self.assertEqual(data["managed_spark_seconds"], 1.0) + + def test_pre_run_cell_handler_swallows_exceptions(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + pre_run = _registered_handler(shell, "pre_run_cell") + with mock.patch( + f"{_MODULE}.start_cell", side_effect=RuntimeError("boom") + ): + try: + pre_run() + except Exception as exc: # pragma: no cover - failure path + self.fail(f"pre_run_cell handler raised: {exc}") + + def test_post_run_cell_handler_swallows_exceptions(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + with mock.patch(f"{_MODULE}.summary", side_effect=RuntimeError("boom")): + try: + post_run() + except Exception as exc: # pragma: no cover - failure path + self.fail(f"post_run_cell handler raised: {exc}") + + def test_handlers_accept_optional_ipython_argument(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + pre_run = _registered_handler(shell, "pre_run_cell") + post_run = _registered_handler(shell, "post_run_cell") + + # Called with an argument, as IPython would. + pre_run(mock.Mock()) + post_run(mock.Mock()) + self.assertIsNotNone(execution_timer._cell_start) + + # -- end-to-end: real record() through the formatted log line ----- + + def test_all_three_kinds_reach_formatted_line(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + with mock.patch(f"{_MODULE}.time.monotonic", return_value=0.0): + start_cell() + record("execute", 0.0, 1.0) + record("fetch", 1.0, 2.0) + record("analyze", 2.0, 3.0) + + post_run = _registered_handler(shell, "post_run_cell") + with mock.patch(f"{_MODULE}.time.monotonic", return_value=3.0): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + + self.assertIn("execute", cm.output[0]) + self.assertIn("fetch", cm.output[0]) + self.assertIn("analyze", cm.output[0]) + + def test_transport_blocked_clamp_reflected_in_formatted_line(self): + """End-to-end: the clamp applied in summary() must show up in + the logged transport clause, so blocked time never appears to + exceed the cell's own reported duration. + """ + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + with mock.patch(f"{_MODULE}.time.monotonic", return_value=0.0): + start_cell() + record("execute", 0.1, 0.2) + record_transport("down", 2 * 1024 * 1024, 10.0) + + post_run = _registered_handler(shell, "post_run_cell") + with mock.patch(f"{_MODULE}.time.monotonic", return_value=1.53): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + + self.assertIn("transport 2.0 MiB down in 1.53s,", cm.output[0]) + + # -- formatting: base sentence --------------------------------------- + + def test_post_run_cell_logs_nothing_when_no_round_trips_and_no_session_creation( + self, + ): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary(cell_seconds=1.0) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertNoLogs(_MODULE, level="INFO"): + post_run() + + def test_base_sentence_omits_across_clause_with_zero_round_trips(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=61.79, + managed_spark_seconds=61.79, + session_creation_seconds=61.79, + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertNotIn("across", cm.output[0]) + + def test_post_run_cell_singular_round_trip_wording(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=1.0, + managed_spark_seconds=1.0, + round_trips=[("execute", 1.0)], + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertIn("1 round trip", cm.output[0]) + self.assertNotIn("1 round trips", cm.output[0]) + + def test_post_run_cell_plural_round_trips_wording(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=1.0, + managed_spark_seconds=1.0, + round_trips=[("execute", 0.5), ("fetch", 0.5)], + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertIn("2 round trips", cm.output[0]) + + # -- formatting: round trip list ---------------------------------- + + def test_round_trip_list_shows_three_slowest_slowest_first(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=10.0, + managed_spark_seconds=10.0, + round_trips=[ + ("execute", 0.10), + ("fetch", 5.00), + ("analyze", 2.00), + ], + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertIn("fetch 5.00s, analyze 2.00s, execute 0.10s", cm.output[0]) + + def test_round_trip_list_appends_plus_n_more_when_over_three(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=25.00, + managed_spark_seconds=22.70, + round_trips=[ + ("fetch", 20.10), + ("analyze", 1.20), + ("execute", 0.80), + ("fetch", 0.50), + ("execute", 0.10), + ], + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertIn( + "fetch 20.10s, analyze 1.20s, execute 0.80s, +2 more", + cm.output[0], + ) + + def test_exact_log_line_five_round_trips_plus_two_more(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=25.00, + managed_spark_seconds=22.70, + round_trips=[ + ("fetch", 20.10), + ("analyze", 1.20), + ("execute", 0.80), + ("fetch", 0.50), + ("execute", 0.10), + ], + ) + expected = ( + "Cell took 25.00s, 22.70s in Managed Spark across 5 round " + "trips (fetch 20.10s, analyze 1.20s, execute 0.80s, +2 more)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + def test_log_format_joins_multiple_clauses_with_semicolon(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=10.0, + managed_spark_seconds=8.0, + round_trips=[("execute", 1.0)], + session_creation_seconds=2.0, + ) + expected = ( + "Cell took 10.00s, 8.00s in Managed Spark across 1 round " + "trip (session creation 2.00s; execute 1.00s)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + # -- formatting: transport clause --------------------------------- + + def test_log_format_transport_rate_with_zero_blocked_does_not_raise(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=1.0, + managed_spark_seconds=1.0, + round_trips=[("execute", 1.0)], + bytes_down=1024 * 1024, + transport_blocked=0.0, + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + try: + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + except ZeroDivisionError as exc: # pragma: no cover + self.fail(f"post_run_cell raised: {exc}") + self.assertIn("0.0 MiB/s", cm.output[0]) + + def test_log_format_transport_kib_omits_rate(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=45.05, + managed_spark_seconds=45.05, + round_trips=[("execute", 1.0), ("fetch", 1.0)], + bytes_down=8.2 * 1024, + transport_blocked=43.24, + ) + expected = ( + "Cell took 45.05s, 45.05s in Managed Spark across 2 round " + "trips (execute 1.00s, fetch 1.00s; transport 8.2 KiB down " + "in 43.24s)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + def test_log_format_transport_exact_one_mib_uses_mib_form(self): + """Boundary: exactly 1 MiB (not merely close to it) must still + use the MiB form with a rate, not the KiB form. + """ + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=2.0, + managed_spark_seconds=2.0, + round_trips=[("execute", 2.0)], + bytes_down=1024 * 1024, + transport_blocked=2.0, + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertIn( + "transport 1.0 MiB down in 2.00s, 0.5 MiB/s", cm.output[0] + ) + + def test_log_format_transport_clause_omitted_when_bytes_down_zero(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=1.0, + managed_spark_seconds=1.0, + round_trips=[("execute", 1.0)], + bytes_down=0, + transport_blocked=5.0, + ) + expected = ( + "Cell took 1.00s, 1.00s in Managed Spark across 1 round trip " + "(execute 1.00s)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + self.assertNotIn("transport", cm.output[0]) + + # -- the four exact reviewer-facing strings, byte-for-byte -------- + + def test_exact_log_line_session_creation_only(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=61.79, + managed_spark_seconds=61.79, + session_creation_seconds=61.79, + ) + expected = "Cell took 61.79s, 61.79s in Managed Spark (session creation 61.79s)" + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + def test_exact_log_line_execute_and_fetch(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=47.44, + managed_spark_seconds=47.43, + round_trips=[("execute", 45.63), ("fetch", 1.80)], + bytes_down=10.4 * 1024, + transport_blocked=45.63, + ) + expected = ( + "Cell took 47.44s, 47.43s in Managed Spark across 2 round " + "trips (execute 45.63s, fetch 1.80s; transport 10.4 KiB " + "down in 45.63s)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + def test_exact_log_line_analyze_only(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=1.67, + managed_spark_seconds=1.67, + round_trips=[("analyze", 1.67)], + bytes_down=1.4 * 1024, + transport_blocked=1.67, + ) + expected = ( + "Cell took 1.67s, 1.67s in Managed Spark across 1 round " + "trip (analyze 1.67s; transport 1.4 KiB down in 1.67s)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + def test_exact_log_line_fetch_only_mib(self): + shell = mock.MagicMock() + with mock.patch("IPython.get_ipython", return_value=shell): + register_cell_timing() + + post_run = _registered_handler(shell, "post_run_cell") + data = _make_summary( + cell_seconds=23.46, + managed_spark_seconds=22.48, + round_trips=[("fetch", 22.48)], + bytes_down=62.1 * 1024 * 1024, + transport_blocked=23.32, + ) + expected = ( + "Cell took 23.46s, 22.48s in Managed Spark across 1 round " + "trip (fetch 22.48s; transport 62.1 MiB down in 23.32s, " + "2.7 MiB/s)" + ) + with mock.patch(f"{_MODULE}.summary", return_value=data): + with self.assertLogs(_MODULE, level="INFO") as cm: + post_run() + self.assertEqual(cm.output, [f"INFO:{_MODULE}:{expected}"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/test_proxy.py b/tests/unit/test_proxy.py index 2333939..fb9499a 100644 --- a/tests/unit/test_proxy.py +++ b/tests/unit/test_proxy.py @@ -14,10 +14,15 @@ import socket import threading import time +from unittest import mock import pytest -from google.cloud.managed_spark_connect.client.proxy import connect_sockets +from google.cloud.managed_spark_connect import execution_timer +from google.cloud.managed_spark_connect.client.proxy import ( + bridged_socket, + connect_sockets, +) @pytest.fixture @@ -149,3 +154,71 @@ def test_proxy_with_timeouts(client_wait, proxy_server_conn, test_message): retry_on_timeouts(proxy_server_conn.recv, 1024).decode() ) assert "\n".join(sent) == "\n".join(received) + + +@pytest.fixture +def reset_execution_timer(): + orig_cell_start = execution_timer._cell_start + orig_totals = dict(execution_timer._totals) + execution_timer._cell_start = None + execution_timer._totals = dict(execution_timer._ZERO_TOTALS) + yield + execution_timer._cell_start = orig_cell_start + execution_timer._totals = orig_totals + + +def test_recv_records_bytes_down_and_blocked_time(reset_execution_timer): + mock_conn = mock.MagicMock() + mock_conn.recv.return_value = "48656c6c6f" # "Hello" in hex + + execution_timer.start_cell() + result = bridged_socket(mock_conn).recv(1024) + summary = execution_timer.summary() + + assert result == b"Hello" + assert summary["bytes_down"] == 5 + assert summary["bytes_up"] == 0 + assert summary["transport_blocked"] >= 0.0 + + +def test_send_records_bytes_up_and_blocked_time(reset_execution_timer): + mock_conn = mock.MagicMock() + + execution_timer.start_cell() + bridged_socket(mock_conn).send(b"Hello") + summary = execution_timer.summary() + + mock_conn.send.assert_called_once_with("48656c6c6f") + assert summary["bytes_up"] == 5 + assert summary["bytes_down"] == 0 + assert summary["transport_blocked"] >= 0.0 + + +def test_recv_raising_still_records_blocked_time_and_reraises( + reset_execution_timer, +): + mock_conn = mock.MagicMock() + mock_conn.recv.side_effect = TimeoutError("boom") + + execution_timer.start_cell() + with pytest.raises(TimeoutError, match="boom"): + bridged_socket(mock_conn).recv(1024) + summary = execution_timer.summary() + + assert summary["bytes_down"] == 0 + assert summary["transport_blocked"] >= 0.0 + + +def test_send_raising_still_records_blocked_time_and_reraises( + reset_execution_timer, +): + mock_conn = mock.MagicMock() + mock_conn.send.side_effect = TimeoutError("boom") + + execution_timer.start_cell() + with pytest.raises(TimeoutError, match="boom"): + bridged_socket(mock_conn).send(b"Hello") + summary = execution_timer.summary() + + assert summary["bytes_up"] == 5 + assert summary["transport_blocked"] >= 0.0 diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index dbd4811..8a18062 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -23,6 +23,7 @@ NotFound, ) from google.cloud.managed_spark_connect import ManagedSparkSession +from google.cloud.managed_spark_connect import execution_timer from google.cloud.managed_spark_connect.exceptions import ManagedSparkConnectException from google.cloud.managed_spark_connect.session import ( _is_valid_label_value, @@ -2771,5 +2772,725 @@ def test_session_skip_terminated(self, mock_session_controller_client): mock_client.get_session.assert_called_once() +class ManagedSparkSessionExecuteTimingTests(unittest.TestCase): + """Tests that _execute / _execute_and_fetch_as_iterator feed the + per-cell execution_timer with real Managed Spark round-trip time. + """ + + def setUp(self): + self.original_environment = dict(os.environ) + os.environ.clear() + os.environ["GOOGLE_CLOUD_PROJECT"] = "test-project" + os.environ["GOOGLE_CLOUD_REGION"] = "test-region" + + self._orig_cell_start = execution_timer._cell_start + self._orig_totals = dict(execution_timer._totals) + self._orig_registered = execution_timer._registered + execution_timer._cell_start = None + execution_timer._totals = dict(execution_timer._ZERO_TOTALS) + execution_timer._registered = False + + def tearDown(self): + execution_timer._cell_start = self._orig_cell_start + execution_timer._totals = self._orig_totals + execution_timer._registered = self._orig_registered + os.environ.clear() + os.environ.update(self.original_environment) + + @staticmethod + def _sql_request(query, operation_id="op-1"): + return ExecutePlanRequest( + session_id="mock-session-id", + client_type="mock-client-type", + plan=Plan( + command=Command( + sql_command=SqlCommand(input=Relation(sql=SQL(query=query))) + ) + ), + tags=["mock-tag"], + user_context=UserContext(user_id="mock-user"), + operation_id=operation_id, + ) + + @staticmethod + def _plain_request(operation_id="op-1"): + return ExecutePlanRequest( + session_id="mock-session-id", + client_type="mock-client-type", + tags=["mock-tag"], + user_context=UserContext(user_id="mock-user"), + operation_id=operation_id, + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("pyspark.sql.connect.client.SparkConnectClient._execute") + def test_execute_records_one_interval_spanning_base_call( + self, + mock_base_execute, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_execute.return_value = None + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + with mock.patch( + "google.cloud.managed_spark_connect.session.time.monotonic", + side_effect=[0.0, 10.0, 12.5, 20.0], + ): + execution_timer.start_cell() + result = client._execute(self._plain_request()) + summary = execution_timer.summary() + + self.assertIsNone(result) + mock_base_execute.assert_called_once() + self.assertEqual(summary["round_trips"], [("execute", 2.5)]) + self.assertEqual(summary["managed_spark_seconds"], 2.5) + self.assertEqual(summary["cell_seconds"], 20.0) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("pyspark.sql.connect.client.SparkConnectClient._execute") + def test_execute_records_and_propagates_exception( + self, + mock_base_execute, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_execute.side_effect = RuntimeError("boom") + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + with mock.patch( + "google.cloud.managed_spark_connect.session.time.monotonic", + side_effect=[0.0, 5.0, 7.0, 9.0], + ): + execution_timer.start_cell() + with self.assertRaisesRegex(RuntimeError, "boom"): + client._execute(self._plain_request()) + summary = execution_timer.summary() + + self.assertEqual(summary["round_trips"], [("execute", 2.0)]) + self.assertEqual(summary["managed_spark_seconds"], 2.0) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch( + "pyspark.sql.connect.client.SparkConnectClient._execute_and_fetch_as_iterator" + ) + def test_iterator_records_nothing_until_consumed( + self, + mock_base_iterator, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_iterator.return_value = iter([1, 2, 3]) + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + execution_timer.start_cell() + gen = client._execute_and_fetch_as_iterator(self._plain_request()) + round_trips = execution_timer.summary()["round_trips"] + + self.assertEqual(round_trips, []) + gen.close() + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch( + "pyspark.sql.connect.client.SparkConnectClient._execute_and_fetch_as_iterator" + ) + def test_iterator_yields_all_values_unchanged( + self, + mock_base_iterator, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_iterator.return_value = iter([1, 2, 3]) + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + execution_timer.start_cell() + gen = client._execute_and_fetch_as_iterator(self._plain_request()) + values = list(gen) + + self.assertEqual(values, [1, 2, 3]) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch( + "pyspark.sql.connect.client.SparkConnectClient._execute_and_fetch_as_iterator" + ) + def test_iterator_records_one_interval_after_full_consumption( + self, + mock_base_iterator, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_iterator.return_value = iter([1, 2, 3]) + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + execution_timer.start_cell() + gen = client._execute_and_fetch_as_iterator(self._plain_request()) + list(gen) + round_trips = execution_timer.summary()["round_trips"] + + self.assertEqual(len(round_trips), 1) + self.assertEqual(round_trips[0][0], "fetch") + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch( + "pyspark.sql.connect.client.SparkConnectClient._execute_and_fetch_as_iterator" + ) + def test_iterator_abandoned_partway_still_records( + self, + mock_base_iterator, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_iterator.return_value = iter([1, 2, 3, 4, 5]) + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + execution_timer.start_cell() + gen = client._execute_and_fetch_as_iterator(self._plain_request()) + self.assertEqual(next(gen), 1) + gen.close() + round_trips = execution_timer.summary()["round_trips"] + + self.assertEqual(len(round_trips), 1) + self.assertEqual(round_trips[0][0], "fetch") + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("pyspark.sql.connect.client.SparkConnectClient._execute") + def test_lazy_select_is_timed_but_link_not_displayed( + self, + mock_base_execute, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_execute.return_value = None + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + session._display_operation_link = mock.Mock() + + execution_timer.start_cell() + client._execute(self._sql_request("SELECT 1")) + round_trips = execution_timer.summary()["round_trips"] + + self.assertEqual(len(round_trips), 1) + self.assertEqual(round_trips[0][0], "execute") + session._display_operation_link.assert_not_called() + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("pyspark.sql.connect.client.SparkConnectClient._analyze") + def test_analyze_records_one_analyze_round_trip( + self, + mock_base_analyze, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_analyze.return_value = "analyze-result" + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + with mock.patch( + "google.cloud.managed_spark_connect.session.time.monotonic", + side_effect=[0.0, 5.0, 8.0, 15.0], + ): + execution_timer.start_cell() + result = client._analyze("schema") + summary = execution_timer.summary() + + self.assertEqual(result, "analyze-result") + mock_base_analyze.assert_called_once_with("schema") + self.assertEqual(summary["round_trips"], [("analyze", 3.0)]) + self.assertEqual(summary["managed_spark_seconds"], 3.0) + self.assertEqual(summary["cell_seconds"], 15.0) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("pyspark.sql.connect.client.SparkConnectClient._analyze") + def test_analyze_records_and_propagates_exception( + self, + mock_base_analyze, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + mock_base_analyze.side_effect = RuntimeError("boom") + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + session = ManagedSparkSession.builder.getOrCreate() + client = session.client + + with mock.patch( + "google.cloud.managed_spark_connect.session.time.monotonic", + side_effect=[0.0, 4.0, 6.0, 9.0], + ): + execution_timer.start_cell() + with self.assertRaisesRegex(RuntimeError, "boom"): + client._analyze("schema") + summary = execution_timer.summary() + + self.assertEqual(summary["round_trips"], [("analyze", 2.0)]) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + def test_session_creation_is_recorded_and_counts_toward_managed_spark_seconds( + self, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + with mock.patch( + "google.cloud.managed_spark_connect.session.time.monotonic", + side_effect=[0.0, 5.0, 35.0], + ): + execution_timer.start_cell() + session = ManagedSparkSession.builder.getOrCreate() + + summary = execution_timer.summary() + + self.assertEqual(summary["session_creation_seconds"], 30.0) + self.assertEqual(summary["managed_spark_seconds"], 30.0) + self.assertEqual(summary["round_trips"], []) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("IPython.get_ipython") + def test_first_cell_session_creation_is_captured_not_dropped( + self, + mock_get_ipython, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ): + """Regression test for the first-cell bug: on the very first + cell of a notebook, IPython has not yet fired pre_run_cell by + the time getOrCreate() runs, because no handler existed yet to + catch it -- so nothing has called start_cell(). Simulate that + ordering by *not* calling execution_timer.start_cell() before + going through the builder path, and confirm the session + creation time is still captured rather than silently dropped. + """ + mock_get_ipython.return_value = mock.MagicMock() + + session = None + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + try: + self.assertIsNone(execution_timer._cell_start) + with mock.patch( + "google.cloud.managed_spark_connect.session.time.monotonic", + side_effect=[0.0, 5.0, 35.0], + ): + session = ManagedSparkSession.builder.getOrCreate() + + summary = execution_timer.summary() + + self.assertEqual(summary["session_creation_seconds"], 30.0) + self.assertEqual(summary["managed_spark_seconds"], 30.0) + self.assertEqual(summary["round_trips"], []) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session + ) + + +class ManagedSparkSessionCellTimingRegistrationTests(unittest.TestCase): + """Tests that session creation wires up the IPython cell timing hooks + exactly once, even across multiple session creations. + """ + + def setUp(self): + self.original_environment = dict(os.environ) + os.environ.clear() + os.environ["GOOGLE_CLOUD_PROJECT"] = "test-project" + os.environ["GOOGLE_CLOUD_REGION"] = "test-region" + + self._orig_registered = execution_timer._registered + execution_timer._registered = False + + def tearDown(self): + execution_timer._registered = self._orig_registered + os.environ.clear() + os.environ.update(self.original_environment) + + @mock.patch( + "google.cloud.managed_spark_connect.environment.is_interactive", + return_value=False, + ) + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") + @mock.patch( + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" + ) + @mock.patch( + "google.cloud.managed_spark_connect.session.is_s8s_session_active" + ) + @mock.patch("IPython.get_ipython") + def test_creating_two_sessions_registers_hooks_once( + self, + mock_get_ipython, + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + mock_is_interactive, + ): + shell = mock.MagicMock() + mock_get_ipython.return_value = shell + + mock_session_controller_client_instance = ( + ManagedSparkSessionBuilderTests._setup_session_creation_mocks( + mock_is_s8s_session_active, + mock_session_id, + mock_client_config, + mock_session_controller_client, + mock_credentials, + ) + ) + + session1 = ManagedSparkSession.builder.getOrCreate() + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session1 + ) + + # Reset the create_session operation so a second creation succeeds. + second_session_response = Session() + second_session_response.runtime_info.endpoints = { + "Spark Connect Server": "sc://spark-connect-server.example.com:443" + } + second_session_response.uuid = "c002e4ef-fe5e-41a8-a157-160aa73e4f80" + mock_session_controller_client_instance.create_session.return_value.result.side_effect = [ + second_session_response + ] + + session2 = ManagedSparkSession.builder.getOrCreate() + try: + pre_run_calls = [ + call + for call in shell.events.register.call_args_list + if call.args[0] == "pre_run_cell" + ] + post_run_calls = [ + call + for call in shell.events.register.call_args_list + if call.args[0] == "post_run_cell" + ] + self.assertEqual(len(pre_run_calls), 1) + self.assertEqual(len(post_run_calls), 1) + finally: + mock_session_controller_client_instance.terminate_session.return_value = ( + mock.Mock() + ) + ManagedSparkSessionBuilderTests.stopSession( + mock_session_controller_client_instance, session2 + ) + + if __name__ == "__main__": unittest.main()