fix(AcsSignals): Various bugfixes for edge cases
This commit is contained in:
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user