From b6bcc1744dd2d3e6fee9c047dbb231e6c3e9d1ba Mon Sep 17 00:00:00 2001 From: wakonig_k Date: Tue, 28 Jul 2026 13:44:58 +0200 Subject: [PATCH] feat(data_api): initial implementation --- bec_widgets/widgets/plots/heatmap/heatmap.py | 134 ++++++++++++++++--- tests/unit_tests/test_heatmap_widget.py | 104 ++++++++++++++ 2 files changed, 216 insertions(+), 22 deletions(-) diff --git a/bec_widgets/widgets/plots/heatmap/heatmap.py b/bec_widgets/widgets/plots/heatmap/heatmap.py index 70b8a422..80e6dca1 100644 --- a/bec_widgets/widgets/plots/heatmap/heatmap.py +++ b/bec_widgets/widgets/plots/heatmap/heatmap.py @@ -2,11 +2,12 @@ from __future__ import annotations import json from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal +from typing import TYPE_CHECKING, Any, Literal import numpy as np import pyqtgraph as pg from bec_lib import bec_logger, messages +from bec_lib.data_api import DataAPI, DataSubscription from bec_lib.endpoints import MessageEndpoints from bec_lib.utils.import_utils import lazy_import, lazy_import_from from bec_qthemes import material_icon @@ -266,6 +267,7 @@ class Heatmap(ImageBase): sync_signal_update = Signal() heatmap_property_changed = Signal() interpolation_requested = Signal(object, int) + data_api_update = Signal(object, object) def __init__(self, parent=None, config: HeatmapConfig | None = None, **kwargs): if config is None: @@ -294,6 +296,8 @@ class Heatmap(ImageBase): self._interpolation_thread: QThread | None = None self._interpolation_worker: _StepInterpolationWorker | None = None self._pending_interpolation_request: _InterpolationRequest | None = None + self._data_api: DataAPI | None = None + self._data_subscription: DataSubscription | None = None self.heatmap_dialog = None self.scan_history_dialog = None self.scan_history_widget = None @@ -310,6 +314,7 @@ class Heatmap(ImageBase): self.bec_dispatcher.connect_slot(self.on_scan_status, MessageEndpoints.scan_status()) self.bec_dispatcher.connect_slot(self.on_scan_progress, MessageEndpoints.scan_progress()) self.heatmap_property_changed.connect(lambda: self.sync_signal_update.emit()) + self.data_api_update.connect(self.update_plot) self.proxy_update_sync = pg.SignalProxy( self.sync_signal_update, rateLimit=5, slot=self.update_plot @@ -459,6 +464,7 @@ class Heatmap(ImageBase): self._history_scan_id = None self._fetch_running_scan() + self._setup_data_api_subscription() # Also notifies settings widgets and triggers a plot update via sync_signal_update self.heatmap_property_changed.emit() @@ -490,6 +496,8 @@ class Heatmap(ImageBase): if scan_item is None: return + self._cleanup_data_api_subscription() + if scan_id is not None: target_scan_id = scan_id elif hasattr(scan_item, "metadata"): @@ -704,6 +712,91 @@ class Heatmap(ImageBase): self.heatmap_dialog = None self.toolbar.components.get_action("heatmap_settings").action.setChecked(False) + def _cleanup_data_api_subscription(self): + if self._data_subscription is None: + return + try: + self._data_subscription.close() + finally: + self._data_subscription = None + + def _setup_data_api_subscription(self): + self._cleanup_data_api_subscription() + + if self._history_scan_id is not None or self._image_config is None: + return + + if not all( + [ + self._image_config.device_x, + self._image_config.device_y, + self._image_config.device_z, + ] + ): + return + + try: + if self._data_api is None: + self._data_api = DataAPI(self.client) + subscription = self._data_api.create_subscription(live=True, buffered=True) + subscription.set_callback(self.data_api_update.emit) + subscription.add_device( + self._image_config.device_x.device, self._image_config.device_x.signal + ) + subscription.add_device( + self._image_config.device_y.device, self._image_config.device_y.signal + ) + subscription.add_device( + self._image_config.device_z.device, self._image_config.device_z.signal + ) + except Exception as exc: + logger.warning(f"Failed to configure heatmap data-api subscription: {exc}") + self._cleanup_data_api_subscription() + return + + self._data_subscription = subscription + + @staticmethod + def _normalize_series_values(values): + if values is None: + return None + if isinstance(values, np.ndarray): + return values.tolist() + if isinstance(values, list): + return values + if isinstance(values, tuple): + return list(values) + return [values] + + def _extract_buffered_series( + self, data: dict, device: str, signal: str + ) -> list[Any] | None: + signal_buffer = data.get(device, {}).get(signal) + if signal_buffer is None: + return None + if isinstance(signal_buffer, dict): + return self._normalize_series_values(signal_buffer.get("value")) + if not isinstance(signal_buffer, list): + return None + + values = [] + for item in signal_buffer: + if not isinstance(item, dict) or "value" not in item: + return None + values.append(item["value"]) + return values + + def _extract_scan_series(self, data, access_key: str, device: str, signal: str) -> list[Any] | None: + if access_key == "val": + values = data.get(device, {}).get(signal, {}).get(access_key, None) + return self._normalize_series_values(values) + + readback = data.get(device, {}).get(signal, None) + if readback is None: + return None + values = readback.read().get("value", None) + return self._normalize_series_values(values) + @SafeSlot(dict, dict) def on_scan_status(self, msg: dict, meta: dict): """ @@ -727,13 +820,16 @@ class Heatmap(ImageBase): self.scan_id = current_scan_id self.scan_item = self.queue.scan_storage.find_scan_by_ID(self.scan_id) # type: ignore - # First trigger to update the scan curves - self.sync_signal_update.emit() + if self._data_subscription is None: + # First trigger to update the scan curves + self.sync_signal_update.emit() @SafeSlot(dict, dict) def on_scan_progress(self, msg: dict, meta: dict): if self._history_scan_id is not None: return + if self._data_subscription is not None: + return self.sync_signal_update.emit() status = msg.get("done") if status: @@ -741,17 +837,13 @@ class Heatmap(ImageBase): QTimer.singleShot(300, self.update_plot) @SafeSlot(verify_sender=True) - def update_plot(self, _=None) -> None: + def update_plot(self, data: dict | None = None, metadata: dict | None = None) -> None: """ Update the plot with the current data. """ if self.scan_item is None: logger.info("No scan executed so far; skipping update.") return - data, access_key = self._fetch_scan_data_and_access() - if data == "none": - logger.info("No scan executed so far; skipping update.") - return if self._image_config is None: return @@ -765,21 +857,18 @@ class Heatmap(ImageBase): except AttributeError: return - if access_key == "val": - x_data = data.get(device_x, {}).get(signal_x, {}).get(access_key, None) - y_data = data.get(device_y, {}).get(signal_y, {}).get(access_key, None) - z_data = data.get(device_z, {}).get(signal_z, {}).get(access_key, None) + if isinstance(data, dict): + x_data = self._extract_buffered_series(data, device_x, signal_x) + y_data = self._extract_buffered_series(data, device_y, signal_y) + z_data = self._extract_buffered_series(data, device_z, signal_z) else: - x_data = data.get(device_x, {}).get(signal_x, {}).read().get("value", None) - y_data = data.get(device_y, {}).get(signal_y, {}).read().get("value", None) - z_data = data.get(device_z, {}).get(signal_z, {}).read().get("value", None) - - if not isinstance(x_data, list): - x_data = x_data.tolist() if isinstance(x_data, np.ndarray) else None - if not isinstance(y_data, list): - y_data = y_data.tolist() if isinstance(y_data, np.ndarray) else None - if not isinstance(z_data, list): - z_data = z_data.tolist() if isinstance(z_data, np.ndarray) else None + data, access_key = self._fetch_scan_data_and_access() + if data == "none": + logger.info("No scan executed so far; skipping update.") + return + x_data = self._extract_scan_series(data, access_key, device_x, signal_x) + y_data = self._extract_scan_series(data, access_key, device_y, signal_y) + z_data = self._extract_scan_series(data, access_key, device_z, signal_z) if x_data is None or y_data is None or z_data is None: logger.warning("x, y, or z data is None; skipping update.") @@ -1705,6 +1794,7 @@ class Heatmap(ImageBase): def cleanup(self): self._finish_interpolation_thread() + self._cleanup_data_api_subscription() if self.scan_history_dialog is not None: self.scan_history_dialog.reject() self.scan_history_dialog = None diff --git a/tests/unit_tests/test_heatmap_widget.py b/tests/unit_tests/test_heatmap_widget.py index 4d0092fb..488cea8c 100644 --- a/tests/unit_tests/test_heatmap_widget.py +++ b/tests/unit_tests/test_heatmap_widget.py @@ -30,6 +30,35 @@ def heatmap_widget(qtbot, mocked_client): yield widget +class _FakeDataSubscription: + def __init__(self): + self.callback = None + self.devices = [] + self.closed = False + + def set_callback(self, callback): + self.callback = callback + return self + + def add_device(self, device, signal): + self.devices.append((device, signal)) + return self + + def close(self): + self.closed = True + + +class _FakeDataAPI: + def __init__(self, client): + self.client = client + self.create_subscription_calls = [] + self.subscription = _FakeDataSubscription() + + def create_subscription(self, **kwargs): + self.create_subscription_calls.append(kwargs) + return self.subscription + + def test_heatmap_plot(heatmap_widget): heatmap_widget.plot(device_x="samx", device_y="samy", device_z="bpm4i") @@ -38,6 +67,31 @@ def test_heatmap_plot(heatmap_widget): assert heatmap_widget._image_config.device_z.device == "bpm4i" +def test_heatmap_plot_sets_up_live_data_api_subscription(heatmap_widget, monkeypatch): + fake_data_api = _FakeDataAPI(heatmap_widget.client) + monkeypatch.setattr( + "bec_widgets.widgets.plots.heatmap.heatmap.DataAPI", lambda client: fake_data_api + ) + + heatmap_widget.plot( + device_x="samx", + device_y="samy", + device_z="bpm4i", + signal_x="samx", + signal_y="samy", + signal_z="bpm4i", + ) + + assert fake_data_api.create_subscription_calls == [{"live": True, "buffered": True}] + assert fake_data_api.subscription.callback == heatmap_widget.data_api_update.emit + assert fake_data_api.subscription.devices == [ + ("samx", "samx"), + ("samy", "samy"), + ("bpm4i", "bpm4i"), + ] + assert heatmap_widget._data_subscription is fake_data_api.subscription + + def test_heatmap_plot_with_scan_id_uses_history(heatmap_widget): history_scan = mock.MagicMock() history_scan.scan_id = "scan-123" @@ -86,6 +140,17 @@ def test_heatmap_update_with_scan_history_resets_cached_image_state(heatmap_widg assert heatmap_widget.scan_id == "scan-456" +def test_heatmap_update_with_scan_history_closes_live_data_api_subscription(heatmap_widget): + history_scan = mock.MagicMock() + history_scan.scan_id = "scan-456" + heatmap_widget._data_subscription = _FakeDataSubscription() + + with mock.patch.object(heatmap_widget, "get_history_scan_item", return_value=history_scan): + heatmap_widget.update_with_scan_history(scan_id="scan-456") + + assert heatmap_widget._data_subscription is None + + def test_heatmap_on_scan_status_resets_after_history_scan_selection(heatmap_widget): heatmap_widget.scan_id = "scan-123" scan_msg = messages.ScanStatusMessage(scan_id="live-scan", status="open", metadata={}, info={}) @@ -388,6 +453,45 @@ def test_heatmap_update_plot(heatmap_widget): assert img.shape == (10, 10) +def test_heatmap_update_plot_from_buffered_data_api_payload(heatmap_widget): + heatmap_widget._image_config = HeatmapConfig( + parent_id="parent_id", + device_x=HeatmapDeviceSignal(device="samx", signal="samx"), + device_y=HeatmapDeviceSignal(device="samy", signal="samy"), + device_z=HeatmapDeviceSignal(device="bpm4i", signal="bpm4i"), + color_map="viridis", + ) + heatmap_widget.scan_item = create_dummy_scan_item() + x_levels = np.linspace(-5, 5, 10).tolist() + y_levels = np.linspace(-5, 5, 10).tolist() + heatmap_widget.scan_item.status_message = messages.ScanStatusMessage( + scan_id="123", + status="open", + scan_name="grid_scan", + metadata={}, + info={ + "positions": _grid_positions(slow_levels=y_levels, fast_levels=x_levels, snaked=True) + }, + request_inputs={"arg_bundle": ["samx", -5, 5, 10, "samy", -5, 5, 10], "kwargs": {}}, + ) + payload = { + "samx": {"samx": [{"value": value, "timestamp": idx} for idx, value in enumerate(x_levels)]}, + "samy": {"samy": [{"value": value, "timestamp": idx} for idx, value in enumerate(y_levels)]}, + "bpm4i": { + "bpm4i": [{"value": idx, "timestamp": idx} for idx in range(len(x_levels))] + }, + } + + with mock.patch.object(heatmap_widget.main_image, "setImage") as mock_set_image: + heatmap_widget.update_plot( + data=payload, + metadata={"scan_id": "123"}, + _override_slot_params={"verify_sender": False}, + ) + img = mock_set_image.mock_calls[0].args[0] + assert img.shape == (10, 10) + + def test_heatmap_update_plot_without_status_message(heatmap_widget): heatmap_widget._image_config = HeatmapConfig( parent_id="parent_id",