From 395daa119fe4e7d1ac34dfa3c18e101bcf931a4f Mon Sep 17 00:00:00 2001 From: Andrew Ma <136692+ajma@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:48:21 -0700 Subject: [PATCH 1/2] feat: measure where a notebook cell's time goes in Managed Spark Adds the timing primitive: module level state accumulating what a cell spends in Managed Spark, reported against the cell's own wall time. Each round trip is kept with its kind and duration, alongside session creation and the bytes pulled over the websocket bridge, so the line distinguishes a slow query from a cold session, a schema lookup, or a large result dragging back through the tunnel. Shared state rather than locals, because the timestamps are taken in different call stacks. Round trips start and end inside gRPC wrappers on threads that are not necessarily the main one, transport is counted on the bridge's forwarding threads, and the reporting happens in an IPython cell callback. Nothing sees all of it, so the totals sit behind a lock. Analyze is a kind of round trip rather than a concept of its own, which is why record takes a kind and there is no separate recorder for it. Anything recorded while no cell is in progress is dropped, keeping this inert and bounded outside IPython, where nothing ever calls start_cell. Round trips are clipped to the cell window, so a generator built in an earlier cell and drained in this one cannot report time from before the cell began. Both the Managed Spark total and the transport blocked time are clamped to the cell's own duration: concurrent round trips and concurrent forwarding threads each accumulate into one counter, so either could otherwise sum past the cell that contains it. register_cell_timing starts a cell when none is in progress. On the first cell nothing has registered a pre_run_cell handler yet, so without this the cell that creates the session, and pays the largest cost there is, would report nothing at all. --- .../managed_spark_connect/execution_timer.py | 269 +++++++ tests/unit/test_execution_timer.py | 739 ++++++++++++++++++ 2 files changed, 1008 insertions(+) create mode 100644 google/cloud/managed_spark_connect/execution_timer.py create mode 100644 tests/unit/test_execution_timer.py 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/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() From 1cdaae6243d97ade2243bca88344bca1b55b89a9 Mon Sep 17 00:00:00 2001 From: Andrew Ma <136692+ajma@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:48:31 -0700 Subject: [PATCH 2/2] feat: report per-cell Managed Spark timing from the session Feeds the timer from the Spark Connect client and the websocket bridge, and registers the IPython hooks, so a cell that touches the cluster logs what it cost: Cell took 61.79s, 61.79s in Managed Spark (session creation 61.79s) Cell took 1.67s, 1.67s in Managed Spark across 1 round trip (analyze 1.67s) 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) Three call sites do the work. _analyze was previously untimed, so a cell that only read a schema reported nothing despite costing 1.67s against a live cluster. Session creation was already measured but logged at debug, which basicConfig(INFO) suppresses, hiding the single largest cost in a first cell; it is now recorded against that cell and logged at info so batch runs see it too. The bridge's recv and send count payload bytes and blocked time, which is how the 2.7 MiB/s above became visible at all. register_cell_timing is also called at the top of getOrCreate, before any creation work, so provisioning lands inside a live cell rather than being dropped. It is idempotent, so the existing call on the reuse path stays as it is. _execute is blocking, so try/finally around it measures the round trip. _execute_and_fetch_as_iterator is a generator function, so the wrapper returns a wrapping generator; timing the call itself would record zero seconds for every collect. All paths record in a finally, because a round trip that failed or was abandoned still spent the time. Lazy SELECT plan registrations are round trips too, so they are timed and counted. Only their link display stays suppressed, as before. --- .../managed_spark_connect/client/proxy.py | 25 +- google/cloud/managed_spark_connect/session.py | 55 +- tests/unit/test_proxy.py | 75 +- tests/unit/test_session.py | 721 ++++++++++++++++++ 4 files changed, 866 insertions(+), 10 deletions(-) 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/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_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()