diff --git a/debye_bec/devices/mo1_bragg/acs.py b/debye_bec/devices/mo1_bragg/acs.py index f874d41..d423d03 100644 --- a/debye_bec/devices/mo1_bragg/acs.py +++ b/debye_bec/devices/mo1_bragg/acs.py @@ -13,29 +13,16 @@ Protocol: from __future__ import annotations -from typing import TYPE_CHECKING +from enum import Enum import numpy as np -from ophyd import Component as Cpt -from ophyd import Kind -from ophyd_devices.interfaces.base_classes.psi_device_base import PSIDeviceBase -from ophyd_devices.utils.controller import Controller, threadlocked -from ophyd_devices.utils.socket import SocketIO, SocketSignal - -if TYPE_CHECKING: - from bec_lib.devicemanager import DeviceManagerBase - from bec_lib.logger import bec_logger +from ophyd_devices.utils.controller import Controller, threadlocked +from ophyd_devices.utils.socket import SocketSignal -# Initialise logger logger = bec_logger.logger -# --------------------------------------------------------------------------- -# Shared communicator -# --------------------------------------------------------------------------- - - class ACSController(Controller): """ Shared TCP/IP communicator for one ACS controller. @@ -47,32 +34,57 @@ class ACSController(Controller): _axes_per_controller = 0 # not used for plain variables, no motion axes + def __init__(self, *, socket_cls, socket_host, socket_port, device_manager): + socket_cls.socket_timeout = 5 + super().__init__( + socket_cls=socket_cls, + socket_host=socket_host, + socket_port=socket_port, + device_manager=device_manager, + term="\r", + trail=["\r:\r", ":\r"], + ) + @threadlocked - def get_var(self, tag: int, precision=3) -> float: + def get_var(self, tag: int, prec: int, idx: int | None = None) -> float: if self.sock is None: self.on() - i = 0 - reply = self.socket_put_and_receive(f"?{{%0.{precision:0.0f}f}}GETVAR({tag},{i})") + + idx = f",{idx:0.0f}" if idx is not None else "" + reply = self.socket_put_and_receive(f"?{{%0.{prec:0.0f}f}}GETVAR({tag}{idx})") + + if reply.startswith("?"): + error = self._query_error(reply) + raise RuntimeError(f"ACS error {reply}: {error}") + return float(reply) @threadlocked - def set_var(self, tag: int, value, precision=3) -> None: + def _query_error(self, reply: str) -> str: + # reply is like "?2002" + return self.socket_put_and_receive(f"?{reply}") + + @threadlocked + def set_var(self, tag: int, value, prec: int, idx: int | None = None) -> None: if self.sock is None: self.on() - i = 0 - self.socket_put_and_receive(f"SETVAR({np.round(value, precision)},{tag},{i})") + idx = f",{idx:0.0f}" if idx is not None else "" + # logger.info(f"Send request: SETVAR({np.round(value, prec)},{tag}{idx})") + reply = self.socket_put_and_receive(f"SETVAR({np.round(value, prec)},{tag}{idx})") + + if reply.startswith("?"): + error = self._query_error(reply) + raise RuntimeError(f"ACS error {reply}: {error}") -# --------------------------------------------------------------------------- -# Signal talking through the shared controller -# --------------------------------------------------------------------------- - - -class ACSVariableSignal(SocketSignal): +class AcsSignal(SocketSignal): """Read/write ACS controller variable, identified by its tag number.""" - def __init__(self, *args, tag: int, **kwargs): + def __init__(self, *args, tag: int, prec: int, num_el: int = 1, enum: Enum = None, **kwargs): self.tag = tag + self.prec = prec + self.num_el = num_el + self.enum = enum super().__init__(*args, **kwargs) @property @@ -80,40 +92,37 @@ class ACSVariableSignal(SocketSignal): return self.root.controller def _socket_get(self): - logger.info(self.controller) - logger.info(self.controller.sock) - return self.controller.get_var(self.tag) + def convert(val): + return self.enum(val).name if self.enum is not None else val + + if self.num_el <= 1: + return convert(self.controller.get_var(self.tag, self.prec)) + return np.array( + [convert(self.controller.get_var(self.tag, self.prec, i)) for i in range(self.num_el)] + ) def _socket_set(self, val): - self.controller.set_var(self.tag, val) + def convert(v): + if self.enum is None: + return v + if isinstance(v, str): + return self.enum[v].value # e.g. "SI111" -> 0 + return self.enum(v).value # e.g. 0 or Xtal.SI111 -> 0 + + if self.num_el <= 1: + self.controller.set_var(self.tag, convert(val), self.prec) + else: + if len(val) != self.num_el: + raise ValueError( + f"Length of val ({len(val)}) must be equal to specified length of variable ({self.num_el})" + ) + + for i, v in enumerate(val): + self.controller.set_var(self.tag, convert(v), self.prec, i) -# --------------------------------------------------------------------------- -# The device -# --------------------------------------------------------------------------- +class AcsSignalRO(AcsSignal): + """Readonly ACS controller variable, identified by its tag number.""" - -class ACSVariables(PSIDeviceBase): - """Three read/write variables on an ACS controller, sharing one connection.""" - - var1 = Cpt(ACSVariableSignal, tag=1, kind=Kind.normal) - var2 = Cpt(ACSVariableSignal, tag=2, kind=Kind.normal) - var3 = Cpt(ACSVariableSignal, tag=3, kind=Kind.normal) - - def __init__( - self, - name: str, - host: str, - port: int = 701, - device_manager: "DeviceManagerBase" | None = None, - **kwargs, - ): - # controller must exist before super().__init__() builds the Cpt signals - self.controller = ACSController( - socket_cls=SocketIO, socket_host=host, socket_port=port, device_manager=device_manager - ) - super().__init__(name=name, device_manager=device_manager, **kwargs) - - def on_connected(self): - # Idempotent: safe even if other devices already opened this controller. - self.controller.on() + def _socket_set(self, val): + return diff --git a/debye_bec/devices/mo1_bragg/acscontroller.py b/debye_bec/devices/mo1_bragg/acscontroller.py deleted file mode 100644 index faf4900..0000000 --- a/debye_bec/devices/mo1_bragg/acscontroller.py +++ /dev/null @@ -1,62 +0,0 @@ -import socket -import threading - -import numpy as np -from ophyd.signal import Signal - - -class ACSSignal(Signal): - - def __init__(self, controller, tag, **kwargs): - self.controller = controller - self.tag = tag - super().__init__(**kwargs) - - def get(self): - value = self.controller.read_tag(self.tag) - self._readback = value - return value - - def put(self, value, **kwargs): - self.controller.write_tag(self.tag, value) - self._readback = value - return super().put(value, **kwargs) - - -class ACSController: - - def __init__(self, host: str, port: int, timeout: float = 2.0): - self.host = host - self.port = port - self.timeout = timeout - self._socket = None - self._lock = threading.Lock() - - self.precision = 12 - - def connect(self): - if self._socket is not None: - return - self._socket = socket.create_connection((self.host, self.port), timeout=self.timeout) - - def disconnect(self): - if self._socket is not None: - self._socket.close() - self._socket = None - - def _query(self, command: str) -> str: - with self._lock: - self.connect() - msg = (command + "\r\n").encode() - self._socket.sendall(msg) - reply = self._socket.recv(4096) - return reply.decode().strip() - - def read_tag(self, tag: int): - i = 0 # For arrays, to be implemented - response = self._query(f"?{{%0.{self.precision:0.0f}f}}GETVAR({tag},{i}") - return float(response) - - def write_tag(self, tag: int, value): - i = 0 # For arrays, to be implemented - self._query(f"SETVAR({np.round(value, self.precision)},{tag},{i}")