mirror of
https://github.com/bec-project/bec_widgets.git
synced 2026-09-06 16:40:56 +02:00
feat(data_api): initial implementation
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user