From 9a98e3fdf344cdccaa4b0b9625143341bd15e9fe Mon Sep 17 00:00:00 2001 From: x01da Date: Tue, 18 Aug 2026 14:57:45 +0200 Subject: [PATCH] fix(AcsSignals): Various bugfixes for edge cases --- debye_bec/devices/mo1_bragg/acs.py | 110 ++++++++++++++++++++++------- 1 file changed, 84 insertions(+), 26 deletions(-) diff --git a/debye_bec/devices/mo1_bragg/acs.py b/debye_bec/devices/mo1_bragg/acs.py index c8ac44a..7575d6c 100644 --- a/debye_bec/devices/mo1_bragg/acs.py +++ b/debye_bec/devices/mo1_bragg/acs.py @@ -24,6 +24,7 @@ from enum import Enum import numpy as np from bec_lib.logger import bec_logger from ophyd_devices.utils.controller import threadlocked +from ophyd_devices.utils.psi_device_base_utils import CompareStatus, SubscriptionStatus from ophyd_devices.utils.socket import SocketSignal from debye_bec.devices.utils.term_trail_controller import TermTrailController @@ -104,7 +105,9 @@ class AcsSignal(SocketSignal): self.last_get = time.time() self._poll_time = max_poll self._poll_thread = None - self._stop_event = threading.Event() + self._poll_stop_event = None + self._lifecycle_lock = threading.Lock() + self._sub_count = 0 super().__init__(*args, **kwargs) @property @@ -123,7 +126,9 @@ class AcsSignal(SocketSignal): if self.num_el <= 1: val = convert(self.controller.get_var(self.tag, self.prec)) - # logger.info(f"Got signal with tag {self.tag} and value {val}, time to last get: {interval*1e3} ms") + # logger.info( + # f"Got signal with tag {self.tag} and value {val}, time to last get: {interval*1e3} ms" + # ) return val return np.array( [convert(self.controller.get_var(self.tag, self.prec, i)) for i in range(self.num_el)] @@ -148,40 +153,81 @@ class AcsSignal(SocketSignal): for i, v in enumerate(val): self.controller.set_var(self.tag, convert(v), self.prec, i) + def set(self, value, *, timeout=None, settle_time=None, **kwargs): + """ + Write value to the controller and return a Status that finishes once + the readback confirms the new value (via CompareStatus + polling). + """ + if self.num_el > 1: + status = self.array_compare_status(value, atol=10 ** (-self.prec), timeout=timeout) + else: + status = CompareStatus(self, value, timeout=timeout, settle_time=settle_time or 0) + try: + self.put(value, **kwargs) + except Exception as exc: + status.set_exception(exc) + return status + + def array_compare_status(self, value, *, atol=None, rtol=None, timeout=None, settle_time=0): + """CompareStatus equivalent for array-valued (num_el > 1) signals.""" + target = np.asarray(value) + + def _compare(value, **kwargs): + current = np.asarray(value) + if atol is not None or rtol is not None: + return bool(np.allclose(current, target, atol=atol or 0, rtol=rtol or 0)) + return bool(np.array_equal(current, target)) + + return SubscriptionStatus(self, _compare, timeout=timeout, settle_time=settle_time) + def subscribe(self, callback, event_type=None, run=True): - self._ensure_polling() + with self._lifecycle_lock: + self._sub_count += 1 + self._ensure_polling_locked() if run: self._force_fresh_read() return super().subscribe(callback, event_type=event_type, run=run) - def _force_fresh_read(self): - old_value = self._readback - try: - new_value = self._socket_get() - except Exception as e: - logger.warning(f"Fresh read failed for {self.name} during subscribe() with {e}") - return - self._readback = new_value - self._run_subs( - sub_type=self.SUB_VALUE, old_value=old_value, value=new_value, timestamp=time.time() - ) - def clear_sub(self, cb, event_type=None): + before = sum(len(d) for d in self._callbacks.values()) super().clear_sub(cb, event_type=event_type) - if not any(self._callbacks.values()): - self._stop_polling() + removed = before - sum(len(d) for d in self._callbacks.values()) + if not removed: + return + with self._lifecycle_lock: + self._sub_count = max(0, self._sub_count - removed) + if self._sub_count == 0: + self._stop_polling_locked() - def _ensure_polling(self): - if self._poll_thread is None or not self._poll_thread.is_alive(): - self._stop_event.clear() - self._poll_thread = threading.Thread(target=self._poll_loop, daemon=True) - self._poll_thread.start() + def _ensure_polling_locked(self): + # Called with _lifecycle_lock held. + # A thread only "covers us" if it's alive AND its own stop + # event hasn't already been set — an alive-but-doomed thread + # (mid-sleep, about to notice a stop request) must NOT block a + # fresh thread from starting. + if ( + self._poll_thread is not None + and self._poll_thread.is_alive() + and self._poll_stop_event is not None + and not self._poll_stop_event.is_set() + ): + return + stop_event = threading.Event() + self._poll_stop_event = stop_event + self._poll_thread = threading.Thread( + target=self._poll_loop, args=(stop_event,), daemon=True + ) + # logger.info(f"Start poll thread for tag {self.tag}") + self._poll_thread.start() - def _stop_polling(self): - self._stop_event.set() + def _stop_polling_locked(self): + # Called with _lifecycle_lock held. + if self._poll_stop_event is not None: + self._poll_stop_event.set() + # logger.info(f"Stop poll thread for tag {self.tag}") - def _poll_loop(self): - while not self._stop_event.is_set(): + def _poll_loop(self, stop_event): + while not stop_event.is_set(): old_value = self._readback try: new_value = self._socket_get() @@ -194,6 +240,18 @@ class AcsSignal(SocketSignal): ) time.sleep(self._poll_time) + def _force_fresh_read(self): + old_value = self._readback + try: + new_value = self._socket_get() + except Exception: + logger.exception(f"Fresh read failed for {self.name} during subscribe()") + return + self._readback = new_value + self._run_subs( + sub_type=self.SUB_VALUE, old_value=old_value, value=new_value, timestamp=time.time() + ) + class AcsSignalRO(AcsSignal): """Readonly ACS controller variable, identified by its tag number."""