From aa1b61736943112118110a26ba39659848a65495 Mon Sep 17 00:00:00 2001 From: wakonig_k Date: Tue, 15 Sep 2026 12:03:20 +0200 Subject: [PATCH 1/2] fix: adjust controller-based devices to new base class --- csaxs_bec/devices/mcs2/mcs2_controller.py | 50 ++----------- csaxs_bec/devices/omny/galil/galil_ophyd.py | 21 +++--- csaxs_bec/devices/omny/galil/galil_rio.py | 3 + csaxs_bec/devices/sim/sim_galil.py | 9 ++- csaxs_bec/devices/sim/sim_lamni.py | 29 ++++---- csaxs_bec/devices/sim/sim_omny.py | 6 +- csaxs_bec/devices/sim/sim_rt_flomni.py | 36 ++++----- csaxs_bec/devices/sim/sim_socket.py | 2 +- .../devices/smaract/smaract_controller.py | 40 +++------- tests/tests_devices/test_fupr_ophyd.py | 6 +- tests/tests_devices/test_galil.py | 74 ++++++++++++++++--- tests/tests_devices/test_galil_flomni.py | 58 +++++++++++++-- tests/tests_devices/test_mcs2.py | 34 +++++++-- tests/tests_devices/test_rt_flomni.py | 28 +++++-- tests/tests_devices/test_smaract.py | 28 +++---- 15 files changed, 256 insertions(+), 168 deletions(-) diff --git a/csaxs_bec/devices/mcs2/mcs2_controller.py b/csaxs_bec/devices/mcs2/mcs2_controller.py index 1523eaf8..af92e0fa 100644 --- a/csaxs_bec/devices/mcs2/mcs2_controller.py +++ b/csaxs_bec/devices/mcs2/mcs2_controller.py @@ -87,6 +87,10 @@ class Mcs2Controller(Controller): _axes_per_controller = 9 _initialized = False PM_PER_MM = 1e9 + put_trail_char = "\r\n" + put_lead_char = "" + get_trail_sequences = ("\r\n",) + get_trim_sequence = "\r\n" USER_ACCESS = [ "query", @@ -132,50 +136,14 @@ class Mcs2Controller(Controller): ) self._initialized = True - @threadlocked - def socket_put(self, cmd: str) -> None: - """Send a raw command. ``cmd`` must already include its own leading ':' or '*'.""" - self.command_history.append(f"[PUT]: {cmd}") - self.sock.put(f"{cmd}\r\n".encode()) - - @threadlocked - def _read_response(self) -> str: - """Read a single response, waiting for the terminator. Returns an - empty string if no response arrived within the timeout (this is expected for - set/move commands, and can also happen for malformed queries).""" - return_val = "" - max_wait_time = 1 - elapsed_time = 0 - sleep_time = 0.01 - while True: - ret = self.socket_get() - return_val += ret - if ret.endswith("\r\n"): - break - time.sleep(sleep_time) - elapsed_time += sleep_time - if elapsed_time > max_wait_time: - break - return self._remove_trailing_characters(return_val) - - def _remove_trailing_characters(self, var: str) -> str: - if len(var) > 1: - return var.split("\r\n")[0] - return var - - @threadlocked - def _send_and_read(self, cmd: str) -> str: - self.socket_put(cmd) - return self._read_response() - def get_error_count(self) -> int: """Return the number of errors currently in the device's error queue.""" - return int(self._send_and_read(":SYST:ERR:COUN?")) + return int(self.socket_put_and_receive(":SYST:ERR:COUN?")) def get_next_error(self) -> tuple[int, str]: """Pop and return the next (oldest) error from the device's error queue as (code, message). code == 0 means the queue was empty ("No Error").""" - raw = self._send_and_read(":SYST:ERR:NEXT?") + raw = self.socket_put_and_receive(":SYST:ERR:NEXT?") code_str, _, message = raw.partition(",") return int(code_str), message.strip().strip('"') @@ -198,7 +166,7 @@ class Mcs2Controller(Controller): """Send a query (``cmd`` must end with '?') and return the bare value as a string (quotes stripped for string properties). Numeric callers are expected to convert the result themselves, e.g. ``int(controller.query(...))``.""" - raw = self._send_and_read(cmd) + raw = self.socket_put_and_receive(cmd) if raw == "": code, message = self.get_next_error() if code != 0: @@ -283,9 +251,7 @@ class Mcs2Controller(Controller): @retry_once def stop_all_axes(self): - return [ - self.command(f":STOP{ax.axis_Id_numeric}") for ax in self._axis if ax is not None - ] + return [self.command(f":STOP{ax.axis_Id_numeric}") for ax in self._axis if ax is not None] @retry_once @axis_checked diff --git a/csaxs_bec/devices/omny/galil/galil_ophyd.py b/csaxs_bec/devices/omny/galil/galil_ophyd.py index de0ce5ff..4c88592c 100644 --- a/csaxs_bec/devices/omny/galil/galil_ophyd.py +++ b/csaxs_bec/devices/omny/galil/galil_ophyd.py @@ -2,9 +2,9 @@ This module contains the base class for Galil controllers as well as the signals used for Galil devices. """ -import functools +from __future__ import annotations + import time -from typing import Any from bec_lib import bec_logger from ophyd.utils import ReadOnlyError @@ -56,10 +56,10 @@ class GalilController(Controller): FAIL = "\033[91m" ENDC = "\033[0m" - @threadlocked - def socket_put(self, val: str) -> None: - self.command_history.append(f"[PUT]: {val}") - self.sock.put(f"{val}\r".encode()) + get_trail_sequences = (":", "?") + get_trim_sequence = "\r\n:" + put_lead_char = "" + put_trail_char = "\r" @retry_once def socket_put_confirmed(self, val: str) -> None: @@ -114,8 +114,7 @@ class GalilController(Controller): def stop_all_axes(self) -> str: if not self.is_thread_active(1): return self.socket_put_and_receive("XQ#STOP,1") - else: - return ":" + return ":" def hard_abort_and_restore_positioning_mode( self, transfer_thread_id: int = 3, timeout: float = 5.0 @@ -161,8 +160,7 @@ class GalilController(Controller): # this is only valid for omny. consider moving to ogalil voltage = float(self.socket_put_and_receive(f"MG @AN[{axis_Id_numeric+1}]").strip()) voltage2 = float(self.socket_put_and_receive(f"MG @AN[{axis_Id_numeric+1}]").strip()) - if voltage2 < voltage: - voltage = voltage2 + voltage = min(voltage, voltage2) # convert from [-10,10]V to [0,300]degC temperature_degC = round((voltage + 10.0) / 20.0 * 300.0, 1) @@ -239,8 +237,7 @@ class GalilController(Controller): if not limit: raise GalilError(f"Failed to drive axis {axis_Id}/{axis_Id_numeric} to limit.") - else: - print("Limit reached.") + print("Limit reached.") def find_reference(self, axis_Id_numeric: int, verbose=0, raise_error=1) -> None: """ diff --git a/csaxs_bec/devices/omny/galil/galil_rio.py b/csaxs_bec/devices/omny/galil/galil_rio.py index 9be4f1da..8fbba2ec 100644 --- a/csaxs_bec/devices/omny/galil/galil_rio.py +++ b/csaxs_bec/devices/omny/galil/galil_rio.py @@ -36,6 +36,9 @@ logger = bec_logger.logger class GalilRIOController(Controller): """Controller Class for Galil RIO controller communication.""" + get_trail_sequences = (":", "?") + get_trim_sequence = "\r\n:" + @threadlocked def socket_put(self, val: str) -> None: """Socker put method.""" diff --git a/csaxs_bec/devices/sim/sim_galil.py b/csaxs_bec/devices/sim/sim_galil.py index f80e2a6b..adf01171 100644 --- a/csaxs_bec/devices/sim/sim_galil.py +++ b/csaxs_bec/devices/sim/sim_galil.py @@ -327,7 +327,14 @@ class SimGalilSocket(SimSocketBase): state_cls = SimGalilState - def handle_command(self, line: str): # noqa: C901 + def handle_command(self, line: str): + """Return a bare status marker or a value followed by CRLF and a colon.""" + reply = self._handle_command(line) + if reply is None or reply in (":", "?"): + return reply + return f"{reply}\r\n:" + + def _handle_command(self, line: str): # noqa: C901 state: SimGalilState = self.state if line.startswith("MG"): diff --git a/csaxs_bec/devices/sim/sim_lamni.py b/csaxs_bec/devices/sim/sim_lamni.py index 406c237c..2215f582 100644 --- a/csaxs_bec/devices/sim/sim_lamni.py +++ b/csaxs_bec/devices/sim/sim_lamni.py @@ -20,7 +20,6 @@ from __future__ import annotations import random import threading -import time from bec_lib.logger import bec_logger @@ -223,7 +222,7 @@ class SimRtLamniState: for cap in caps: parts += [f"{cap + random.gauss(0.0, 0.001):.5f}", f"{0.001:.5f}"] parts += [f"{angle + random.gauss(0.0, 1e-6):.5f}", f"{1e-6:.5f}"] - return ", ".join(parts) + "\n" + return ", ".join(parts) + "\r\n" class SimRtLamniSocket(SimSocketBase): @@ -244,14 +243,14 @@ class SimRtLamniSocket(SimSocketBase): return None if args.startswith("2"): status = 0 if state.feedback_running else 1 - return f"{status},{state.ssi[0]:.0f},{state.ssi[1]:.0f}\n" + return f"{status},{state.ssi[0]:.0f},{state.ssi[1]:.0f}\r\n" if args.startswith("3"): # start ssi averaging return None if args.startswith("4"): measured = state.measured_positions() - return f"0,{measured[1]:f},{measured[0]:f}\n" + return f"0,{measured[1]:f},{measured[0]:f}\r\n" if args.startswith("7"): - return f"1,{state.angle_rad:f},{state.angle_interferometer_signal:f}\n" + return f"1,{state.angle_rad:f},{state.angle_interferometer_signal:f}\r\n" return None if cmd == "A": @@ -259,12 +258,12 @@ class SimRtLamniSocket(SimSocketBase): return None caps = state.cap_sensors values = ",".join(f"{val + random.gauss(0.0, 0.001):.3f}" for val in caps) - return f"100,{values}\n" + return f"100,{values}\r\n" if cmd == "S": if args.startswith("s"): # reset and start rt sampler return None - return "100,0.0,0.0,0.0,0.0,0.0,0.0\n" + return "100,0.0,0.0,0.0,0.0,0.0,0.0\r\n" if cmd == "V": # velocity limit with state.lock: @@ -273,14 +272,14 @@ class SimRtLamniSocket(SimSocketBase): if cmd == "a": if args.startswith("r"): - return f"{state.angle_rad:f}\n" + return f"{state.angle_rad:f}\r\n" state.angle_rad = float(args) return None if cmd == "p": if args.startswith("r"): measured = state.measured_positions() - return f"{measured[0]:f},{measured[1]:f},{measured[2]:f}\n" + return f"{measured[0]:f},{measured[1]:f},{measured[2]:f}\r\n" if args.startswith("a"): axis_str, value_str = args[1:].split(",") with state.lock: @@ -302,23 +301,23 @@ class SimRtLamniSocket(SimSocketBase): state.scan_current, ) # LAMNI_server: "%.0f,%.0f,%.0f" - integer-formatted, parsed with int() - return f"{mode:.0f},{num_pos:.0f},{current:.0f}\n" + return f"{mode:.0f},{num_pos:.0f},{current:.0f}\r\n" if args.startswith("d"): num_pos = state.start_scan() - return f"Scan mode started with #positions in scan {num_pos:.0f}. Timing: Detector Trigger.\n" + return f"Scan mode started with #positions in scan {num_pos:.0f}. Timing: Detector Trigger.\r\n" if args.startswith("h"): state.clear_scan() - return "0.0, 0.0, 0.0\n" + return "0.0, 0.0, 0.0\r\n" count = state.add_scan_position([float(val) for val in args.split(",")]) - return f"{count:.0f}\n" + return f"{count:.0f}\r\n" if cmd == "r": return state.sample_row(int(args)) if cmd == "o": - return "1\n" + return "1\r\n" if cmd == "t": - return "1\n" + return "1\r\n" if cmd == "d": return None diff --git a/csaxs_bec/devices/sim/sim_omny.py b/csaxs_bec/devices/sim/sim_omny.py index 5aebf832..982bfb53 100644 --- a/csaxs_bec/devices/sim/sim_omny.py +++ b/csaxs_bec/devices/sim/sim_omny.py @@ -146,7 +146,7 @@ class SimRtOmnyState(SimRtFlomniState): # [15..18] sample stage <-> FZP x/y avg/stdev for target_val in (target[0], target[1]): fields += [f"{target_val + random.gauss(0.0, noise):.5f}", f"{noise:.5f}"] - return ", ".join(fields) + "\n" + return ", ".join(fields) + "\r\n" class SimRtOmnySocket(SimRtFlomniSocket): @@ -159,12 +159,12 @@ class SimRtOmnySocket(SimRtFlomniSocket): cmd, args = line[0], line[1:] if cmd == "y": # OMNY: single combined slew-rate-limiter flag (flOMNI returns the sum, 3) - return "1.000000\n" + return "1.000000\r\n" if cmd == "J": if args.startswith("2"): signals = ",".join(f"{val:.0f}" for val in state.ssi_signals) status = 0 if state.feedback_running else 1 - return f"{status},{signals}\n" + return f"{status},{signals}\r\n" # J0 disable / J1 enable with reset / J3 averaging: no reply if args.startswith("0"): state.feedback_running = False diff --git a/csaxs_bec/devices/sim/sim_rt_flomni.py b/csaxs_bec/devices/sim/sim_rt_flomni.py index 873d88aa..6600f141 100644 --- a/csaxs_bec/devices/sim/sim_rt_flomni.py +++ b/csaxs_bec/devices/sim/sim_rt_flomni.py @@ -119,7 +119,7 @@ class SimRtFlomniState: return ( f"{index}, 100, {target[0]:.5f}, {avg_x:.5f}, {abs(random.gauss(noise, noise / 4)):.5f}," f" {target[1]:.5f}, {avg_y:.5f}, {abs(random.gauss(noise, noise / 4)):.5f}," - f" {rotz:.5f}, {abs(random.gauss(0.1, 0.02)):.5f}\n" + f" {rotz:.5f}, {abs(random.gauss(0.1, 0.02)):.5f}\r\n" ) @@ -140,13 +140,13 @@ class SimRtFlomniSocket(SimSocketBase): state.feedback_running = True return None if args.startswith("2"): - return f"{0 if state.feedback_running else 1}\n" + return f"{0 if state.feedback_running else 1}\r\n" return None if cmd == "p": if args.startswith("r"): targets = state.targets - return f"{targets[0]:.5f},{targets[1]:.5f},{targets[2]:.5f}\n" + return f"{targets[0]:.5f},{targets[1]:.5f},{targets[2]:.5f}\r\n" if args.startswith("a"): axis_str, value_str = args[1:].split(",") with state.lock: @@ -163,22 +163,22 @@ class SimRtFlomniSocket(SimSocketBase): return None if args.startswith("r"): mode, num_pos, current = state.scan_status() - return f"{mode:.5f},{num_pos:.5f},{current:.5f}\n" + return f"{mode:.5f},{num_pos:.5f},{current:.5f}\r\n" if args.startswith("d"): num_pos = state.start_scan() - return f"Scan started {num_pos:.0f} positions.\n" + return f"Scan started {num_pos:.0f} positions.\r\n" if args.startswith("h"): state.clear_scan() - return "0.00000, 0.00000, 0.00000\n" + return "0.00000, 0.00000, 0.00000\r\n" count = state.add_scan_position([float(val) for val in args.split(",")]) - return f"{count:.0f}\n" + return f"{count:.0f}\r\n" if cmd == "r": return state.sample_row(int(args)) if cmd == "a": if args.startswith("r"): - return f"{state.angle_rad:f}\n" + return f"{state.angle_rad:f}\r\n" state.angle_rad = float(args) return None @@ -197,31 +197,31 @@ class SimRtFlomniSocket(SimSocketBase): emitter = state.emitter return ( f"{emitter['ty']:.5f},{emitter['tz']:.5f}," - f"{emitter['thr_laser']:.2f},{emitter['psd_low']:.5f}\n" + f"{emitter['thr_laser']:.2f},{emitter['psd_low']:.5f}\r\n" ) if cmd == "g": - return f"{state.pid_x_voltage + random.gauss(0.0, 0.005):f}\n" + return f"{state.pid_x_voltage + random.gauss(0.0, 0.005):f}\r\n" if cmd == "G": - return f"{random.gauss(0.0, 0.005):f}\n" + return f"{random.gauss(0.0, 0.005):f}\r\n" if cmd == "y": - return "3.000000\n" + return "3.000000\r\n" if cmd == "w": - return "32.000000\n" + return "32.000000\r\n" if cmd == "j": - return f"{state.ssi_signal:f}\n" + return f"{state.ssi_signal:f}\r\n" if cmd == "k": axis = int(args) - return f"{state.targets[axis] if axis < 3 else 0.0:f}\n" + return f"{state.targets[axis] if axis < 3 else 0.0:f}\r\n" if cmd == "v": return None if cmd == "d": return None if cmd == "o": # not present in the current real server; provided for API completeness - return "1\n" + return "1\r\n" if cmd == "t": - return "1" + return "1\r\n" logger.warning(f"[sim rt] {self.host}:{self.port} unhandled command '{line}'") return None @@ -236,7 +236,7 @@ class SimRtFlomniSocket(SimSocketBase): return ( f"{target_z:.2f},{target_z:.2f},{intensity:.2f},{threshold:.2f},{5.0:.2f}," f"{target_y:.2f},{target_y:.2f},{intensity:.2f},{threshold:.2f},{5.0:.2f}," - f"{enabled:.0f}\n" + f"{enabled:.0f}\r\n" ) diff --git a/csaxs_bec/devices/sim/sim_socket.py b/csaxs_bec/devices/sim/sim_socket.py index a52adaa3..3a747b68 100644 --- a/csaxs_bec/devices/sim/sim_socket.py +++ b/csaxs_bec/devices/sim/sim_socket.py @@ -99,7 +99,7 @@ class SimSocketBase: if reply is not None: self._recv_buffer.append(reply.encode()) - def receive(self, buffer_length=1024): + def receive(self, buffer_length=1024, timeout: int = 2): with self._lock: if self._recv_buffer: return self._recv_buffer.pop(0) diff --git a/csaxs_bec/devices/smaract/smaract_controller.py b/csaxs_bec/devices/smaract/smaract_controller.py index 5eca450c..0936fdaf 100644 --- a/csaxs_bec/devices/smaract/smaract_controller.py +++ b/csaxs_bec/devices/smaract/smaract_controller.py @@ -70,6 +70,10 @@ class SmaractController(Controller): "is_axis_moving", "print_command_history", ] + get_trail_sequences = ("\n",) + get_trim_sequence = "\n" + put_lead_char = ":" + put_trail_char = "\n" def __init__( self, @@ -98,34 +102,15 @@ class SmaractController(Controller): ) self._sensors = SmaractSensors() - @threadlocked - def socket_put(self, val: str): - self.command_history.append(f"[PUT]: {val}") - self.sock.put(f":{val}\n".encode()) - @threadlocked def socket_put_and_receive( - self, val: str, remove_trailing_chars=True, check_for_errors=True, raise_if_not_status=False + self, val: str, remove_trailing_chars=True, timeout: int = 2, raise_if_not_status=False ) -> str: - self.socket_put(val) - return_val = "" - max_wait_time = 1 - elapsed_time = 0 - sleep_time = 0.01 - while True: - ret = self.socket_get() - return_val += ret - if ret.endswith("\n"): - break - time.sleep(sleep_time) - elapsed_time += sleep_time - if elapsed_time > max_wait_time: - break - if remove_trailing_chars: - return_val = self._remove_trailing_characters(return_val) - logger.debug(f"Sending {val}; Returned {return_val}") - if check_for_errors: - self._check_for_error(return_val, raise_if_not_status=raise_if_not_status) + + return_val = super().socket_put_and_receive( + val, remove_trailing_chars=remove_trailing_chars, timeout=timeout + ) + self._check_for_error(return_val, raise_if_not_status=raise_if_not_status) return return_val @retry_once @@ -492,11 +477,6 @@ class SmaractController(Controller): "Expected error / status message but failed to parse it." ) - def _remove_trailing_characters(self, var: str) -> str: - if len(var) > 1: - return var.split("\n")[0] - return var - def _message_starts_with(self, msg: str, leading_chars: str) -> bool: if msg.startswith(leading_chars): return True diff --git a/tests/tests_devices/test_fupr_ophyd.py b/tests/tests_devices/test_fupr_ophyd.py index 9e568563..3d796f2d 100644 --- a/tests/tests_devices/test_fupr_ophyd.py +++ b/tests/tests_devices/test_fupr_ophyd.py @@ -33,8 +33,8 @@ def fsamroy(dm_with_devices): @pytest.mark.parametrize( "pos,msg_received,msg_put,sign", [ - (-0.5, b" -12800\n\r", [b"TPA\r", b"MG_BGA\r", b"TPA\r"], 1), - (-0.5, b" 12800\n\r", [b"TPA\r", b"MG_BGA\r", b"TPA\r"], -1), + (-0.5, b" -12800\r\n:", [b"TPA\r", b"MG_BGA\r", b"TPA\r"], 1), + (-0.5, b" 12800\r\n:", [b"TPA\r", b"MG_BGA\r", b"TPA\r"], -1), ], ) def test_axis_get(fsamroy, pos, msg_received, msg_put, sign): @@ -54,7 +54,7 @@ def test_axis_get(fsamroy, pos, msg_received, msg_put, sign): ( 0, [b"MG axisref\r", b"PAA=0\r", b"PAA=0\r", b"BGA\r"], - [b"1.00", b"-1", b":", b":", b":", b":", b"-1"], + [b"1.00\r\n:", b"-1\r\n:", b":", b":", b":", b":", b"-1\r\n:"], ) ], ) diff --git a/tests/tests_devices/test_galil.py b/tests/tests_devices/test_galil.py index ac49678e..7c5bb104 100644 --- a/tests/tests_devices/test_galil.py +++ b/tests/tests_devices/test_galil.py @@ -57,7 +57,7 @@ def leyex(dm_with_devices): leyex_motor.controller._reset_controller() -@pytest.mark.parametrize("pos,msg,sign", [(1, b" -12800\n\r", 1), (-1, b" -12800\n\r", -1)]) +@pytest.mark.parametrize("pos,msg,sign", [(1, b" -12800\r\n:", 1), (-1, b" -12800\r\n:", -1)]) def test_axis_get(leyey, pos, msg, sign): leyey.sign = sign leyey.controller.sock.flush_buffer() @@ -81,7 +81,7 @@ def test_axis_get(leyey, pos, msg, sign): b"XQ#NEWPAR\r", b"MG_XQ0\r", ], - [b"1.00", b"-1", b":", b":", b":", b":", b"-1"], + [b"1.00\r\n:", b"-1\r\n:", b":", b":", b":", b":", b"-1\r\n:"], ) ], ) @@ -110,7 +110,18 @@ def test_axis_put(leyey, target_pos, socket_put_messages, socket_get_messages): b"MG_XQ2\r", b"MG _LRA, _LFA\r", ], - [b":", b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"1.000 0.000"], + [ + b":", + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"1.000 0.000\r\n:", + ], ), ( 1, @@ -127,7 +138,18 @@ def test_axis_put(leyey, target_pos, socket_put_messages, socket_get_messages): b"MG_XQ2\r", b"MG _LRB, _LFB\r", ], - [b":", b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"0.000 1.000"], + [ + b":", + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"0.000 1.000\r\n:", + ], ), ], ) @@ -154,7 +176,17 @@ def test_drive_axis_to_limit(leyex, axis_nr, direction, socket_put_messages, soc b"MG_XQ2\r", b"MG axisref[0]\r", ], - [b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"1.00"], + [ + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"1.00\r\n:", + ], ), ( 1, @@ -169,7 +201,17 @@ def test_drive_axis_to_limit(leyex, axis_nr, direction, socket_put_messages, soc b"MG_XQ2\r", b"MG axisref[1]\r", ], - [b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"1.00"], + [ + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"1.00\r\n:", + ], ), ], ) @@ -183,6 +225,20 @@ def test_find_reference(leyex, axis_nr, socket_put_messages, socket_get_messages assert leyex.controller.sock.buffer_put == socket_put_messages +@pytest.mark.parametrize("device_fixture", ["leyey", "galil_rio"]) +@pytest.mark.parametrize( + "chunks,expected", [([b":"], ":"), ([b"?"], "?"), ([b"123\r\n", b":"], "123")] +) +def test_galil_receive_and_trim_sequences(request, device_fixture, chunks, expected): + controller = request.getfixturevalue(device_fixture).controller + controller.sock.flush_buffer() + controller.sock.buffer_recv = chunks.copy() + + assert controller.socket_put_and_receive("MG 123") == expected + assert controller.sock.buffer_recv == [] + assert controller.sock.buffer_put == [b"MG 123\r"] + + def test_wait_for_connection_called(dm_with_devices): """Test that wait_for_connection is called on all motors that have a socket controller.""" dm = dm_with_devices @@ -283,7 +339,7 @@ def test_galil_rio_signal_read(galil_rio): assert galil_rio.analog_in.ch0._readback_timeout == 0.1 # Default read timeout of 100ms # Mock the socket to return specific values - analog_bufffer = b" 1.234 2.345 3.456 4.567 5.678 6.789 7.890 8.901\r\n" + analog_bufffer = b" 1.234 2.345 3.456 4.567 5.678 6.789 7.890 8.901\r\n:" galil_rio.controller.sock.buffer_recv = [] # Clear any existing buffer galil_rio.controller.sock.buffer_recv.append(analog_bufffer) read_values = galil_rio.read() @@ -328,7 +384,7 @@ def test_galil_rio_signal_read(galil_rio): value_callback_buffer.append(readback) galil_rio.analog_in.ch0.subscribe(value_callback, run=False) - galil_rio.controller.sock.buffer_recv = [b" 2.5 2.6 2.7 2.8 2.9 3.0 3.1 3.2"] + galil_rio.controller.sock.buffer_recv = [b" 2.5 2.6 2.7 2.8 2.9 3.0 3.1 3.2\r\n:"] expected_values = [2.5, 2.6, 2.7, 2.8, 2.9, 3.0, 3.1, 3.2] ################## @@ -381,7 +437,7 @@ def test_galil_rio_digital_out_signal(galil_rio): for ii in range(galil_rio.digital_out.ch0._NUM_DIGITAL_OUTPUT_CHANNELS): cmd = f"MG@OUT[{ii}]\r".encode() excepted_put_buffer.append(cmd) - recv = " 1.000".encode() + recv = b" 1.000\r\n:" buffer_receive.append(recv) galil_rio.controller.sock.buffer_recv = buffer_receive # Mock response for readback diff --git a/tests/tests_devices/test_galil_flomni.py b/tests/tests_devices/test_galil_flomni.py index 64b09cb2..ca7d9772 100644 --- a/tests/tests_devices/test_galil_flomni.py +++ b/tests/tests_devices/test_galil_flomni.py @@ -40,7 +40,7 @@ def leyex(dm_with_devices): leyex_motor.controller._reset_controller() -@pytest.mark.parametrize("pos,msg,sign", [(1, b" -12800\n\r", 1), (-1, b" -12800\n\r", -1)]) +@pytest.mark.parametrize("pos,msg,sign", [(1, b" -12800\r\n:", 1), (-1, b" -12800\r\n:", -1)]) def test_axis_get(leyey, pos, msg, sign): leyey.sign = sign leyey.controller.sock.flush_buffer() @@ -64,7 +64,7 @@ def test_axis_get(leyey, pos, msg, sign): b"XQ#NEWPAR\r", b"MG_XQ0\r", ], - [b"1.00", b"-1", b":", b":", b":", b":", b"-1"], + [b"1.00\r\n:", b"-1\r\n:", b":", b":", b":", b":", b"-1\r\n:"], ) ], ) @@ -93,7 +93,18 @@ def test_axis_put(leyey, target_pos, socket_put_messages, socket_get_messages): b"MG _MOA\r", b"MG _LRA, _LFA\r", ], - [b":", b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"1.000 0.000"], + [ + b":", + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"1.000 0.000\r\n:", + ], ), ( 1, @@ -110,7 +121,18 @@ def test_axis_put(leyey, target_pos, socket_put_messages, socket_get_messages): b"MG _MOB\r", b"MG _LRB, _LFB\r", ], - [b":", b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"0.000 1.000"], + [ + b":", + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"0.000 1.000\r\n:", + ], ), ], ) @@ -137,7 +159,17 @@ def test_drive_axis_to_limit(leyex, axis_nr, direction, socket_put_messages, soc b"MG _MOA\r", b"MG axisref[0]\r", ], - [b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"1.00"], + [ + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"1.00\r\n:", + ], ), ( 1, @@ -152,7 +184,17 @@ def test_drive_axis_to_limit(leyex, axis_nr, direction, socket_put_messages, soc b"MG _MOB\r", b"MG axisref[1]\r", ], - [b":", b":", b"-1", b":", b"0", b"0", b"-1", b"-1", b"1.00"], + [ + b":", + b":", + b"-1\r\n:", + b":", + b"0\r\n:", + b"0\r\n:", + b"-1\r\n:", + b"-1\r\n:", + b"1.00\r\n:", + ], ), ], ) @@ -166,8 +208,8 @@ def test_find_reference(leyex, axis_nr, socket_put_messages, socket_get_messages @pytest.mark.parametrize( "axis_Id,socket_put_messages,socket_get_messages,triggered", [ - ("A", [b"MG @IN[14]\r"], [b" 1.0000\n"], True), - ("B", [b"MG @IN[14]\r"], [b" 0.0000\n"], False), + ("A", [b"MG @IN[14]\r"], [b" 1.0000\r\n:"], True), + ("B", [b"MG @IN[14]\r"], [b" 0.0000\r\n:"], False), ], ) def test_fosaz_light_curtain_is_triggered( diff --git a/tests/tests_devices/test_mcs2.py b/tests/tests_devices/test_mcs2.py index 8f8ac461..96c32e2d 100644 --- a/tests/tests_devices/test_mcs2.py +++ b/tests/tests_devices/test_mcs2.py @@ -2,6 +2,7 @@ from unittest import mock import pytest from ophyd_devices.tests.utils import SocketMock +from ophyd_devices.utils.controller import ControllerCommunicationError from csaxs_bec.devices.mcs2 import Mcs2Controller from csaxs_bec.devices.mcs2.mcs2_controller import Mcs2ChannelType @@ -17,7 +18,10 @@ NO_ERROR = b'0,"No Error"\r\n' def controller(dm_with_devices): Mcs2Controller._reset_controller() controller = Mcs2Controller( - socket_cls=SocketMock, socket_host="dummy", socket_port=55551, device_manager=dm_with_devices + socket_cls=SocketMock, + socket_host="dummy", + socket_port=55551, + device_manager=dm_with_devices, ) controller.on() controller.sock.flush_buffer() @@ -118,19 +122,30 @@ def test_command_no_error_does_not_raise(controller): assert controller.sock.buffer_put == [b":STOP0\r\n", b":SYST:ERR:NEXT?\r\n"] +def test_command_does_not_retry_error_queue_pop(controller): + error = TimeoutError("Error reply timed out") + with mock.patch.object(controller.sock, "receive", side_effect=[error, NO_ERROR]) as receive: + with pytest.raises(ControllerCommunicationError) as exc_info: + controller.command(":STOP0") + + assert exc_info.value.__cause__ is error + assert receive.call_count == 1 + assert controller.sock.buffer_put == [b":STOP0\r\n", b":SYST:ERR:NEXT?\r\n"] + + def test_query_timeout_without_error_raises_communication_error(controller): # _send_and_read is patched directly (rather than driving the real 1s socket # timeout loop via buffer_recv) to keep this test fast and deterministic: the # first call simulates the original query timing out (empty response), the # second simulates the subsequent error-queue poll finding nothing queued. - with mock.patch.object(controller, "_send_and_read", side_effect=["", '0,"No Error"']): + with mock.patch.object(controller, "socket_put_and_receive", side_effect=["", '0,"No Error"']): with pytest.raises(Mcs2CommunicationError): controller.query(":CHAN0:POS?") def test_query_timeout_with_queued_error_raises_error_code(controller): with mock.patch.object( - controller, "_send_and_read", side_effect=["", '-113,"Undefined header"'] + controller, "socket_put_and_receive", side_effect=["", '-113,"Undefined header"'] ): with pytest.raises(Mcs2ErrorCode) as exc_info: controller.query(":CHAN0:BOGUS?") @@ -142,7 +157,12 @@ def test_query_timeout_with_queued_error_raises_error_code(controller): [ (50, 0, None, [b":CHAN0:MMOD 0\r\n", b":CHAN0:HOLD 1000\r\n", b":MOVE0 50000000000\r\n"]), (0, 0, 800, [b":CHAN0:MMOD 0\r\n", b":CHAN0:HOLD 800\r\n", b":MOVE0 0\r\n"]), - (20.23, 1, None, [b":CHAN1:MMOD 0\r\n", b":CHAN1:HOLD 1000\r\n", b":MOVE1 20230000000\r\n"]), + ( + 20.23, + 1, + None, + [b":CHAN1:MMOD 0\r\n", b":CHAN1:HOLD 1000\r\n", b":MOVE1 20230000000\r\n"], + ), ], ) def test_move_axis_to_absolute_position(controller, pos, axis, hold_time, get_msg): @@ -197,7 +217,11 @@ def test_find_reference_mark(controller): @pytest.mark.parametrize( "move_speed,axis,get_msg", - [(50, 0, b":CHAN0:VEL 50000000000\r\n"), (0, 0, b":CHAN0:VEL 0\r\n"), (20.23, 1, b":CHAN1:VEL 20230000000\r\n")], + [ + (50, 0, b":CHAN0:VEL 50000000000\r\n"), + (0, 0, b":CHAN0:VEL 0\r\n"), + (20.23, 1, b":CHAN1:VEL 20230000000\r\n"), + ], ) def test_set_closed_loop_move_speed(controller, move_speed, axis, get_msg): controller.sock.buffer_recv = NO_ERROR diff --git a/tests/tests_devices/test_rt_flomni.py b/tests/tests_devices/test_rt_flomni.py index 811d6571..01a0b307 100644 --- a/tests/tests_devices/test_rt_flomni.py +++ b/tests/tests_devices/test_rt_flomni.py @@ -17,6 +17,7 @@ def rt_flomni(): socket_port=8081, device_manager=mock.MagicMock(), ) + rt_flomni.on() with mock.patch.object(rt_flomni, "sock"): rtx = mock.MagicMock(spec=RtFlomniMotor) rtx.name = "rtx" @@ -34,6 +35,7 @@ def rt_flomni(): rt_flomni.set_axis(axis=rty, axis_nr=1) rt_flomni.set_axis(axis=rtz, axis_nr=2) yield rt_flomni + rt_flomni.off(update_config=False) RtFlomniController._reset_controller() @@ -46,7 +48,7 @@ def test_rt_flomni_move_to_zero(rt_flomni): ] -@pytest.mark.parametrize("return_value,is_running", [(b"1.00\n", False), (b"0.00\n", True)]) +@pytest.mark.parametrize("return_value,is_running", [(b"1.00\r\n", False), (b"0.00\r\n", True)]) def test_rt_flomni_feedback_is_running(rt_flomni, return_value, is_running): rt_flomni.sock.receive.return_value = return_value assert rt_flomni.feedback_is_running() == is_running @@ -54,7 +56,8 @@ def test_rt_flomni_feedback_is_running(rt_flomni, return_value, is_running): def test_feedback_enable_with_reset(rt_flomni): - + # Cyclic error compensation status (w0, w1), followed by PID voltage (g). + rt_flomni.sock.receive.side_effect = [b"32\r\n", b"32\r\n", b"0.05\r\n"] device_manager = rt_flomni.device_manager device_manager.devices.fsamx.user_parameter.get.return_value = 0.05 device_manager.devices.fsamx.obj.readback.get.return_value = 0.05 @@ -69,13 +72,24 @@ def test_feedback_enable_with_reset(rt_flomni): rt_flomni.feedback_enable_with_reset() laser_tracker_on.assert_called_once() + assert rt_flomni.rt_pid_voltage == 0.05 + device_manager.devices.rtx.update_user_parameter.assert_called_once_with( + {"rt_pid_voltage": 0.05} + ) + assert rt_flomni.sock.put.call_args_list[-3:] == [ + mock.call(b"w0\n"), + mock.call(b"w1\n"), + mock.call(b"g\n"), + ] + def test_move_samx_to_scan_region(rt_flomni): - device_manager = rt_flomni.device_manager - device_manager.devices.rtx.user_parameter.get.return_value = 1 - rt_flomni.move_samx_to_scan_region(20, 2) - assert mock.call(b"v0\n") not in rt_flomni.sock.put.mock_calls - assert mock.call(b"v1\n") in rt_flomni.sock.put.mock_calls + with mock.patch.object(rt_flomni, "get_pid_x", return_value=1): + device_manager = rt_flomni.device_manager + device_manager.devices.rtx.user_parameter.get.return_value = 1 + rt_flomni.move_samx_to_scan_region(20, 2) + assert mock.call(b"v0\n") not in rt_flomni.sock.put.mock_calls + assert mock.call(b"v1\n") in rt_flomni.sock.put.mock_calls def test_feedback_enable_without_reset(rt_flomni): diff --git a/tests/tests_devices/test_smaract.py b/tests/tests_devices/test_smaract.py index 3e77ca03..c77a1e62 100644 --- a/tests/tests_devices/test_smaract.py +++ b/tests/tests_devices/test_smaract.py @@ -78,24 +78,24 @@ def test_axis_is_referenced(controller, axis, is_referenced, get_message, return "return_msg,exception,raised", [ (b"false\n", SmaractCommunicationError, False), - (b":E0,1", SmaractErrorCode, True), - (b":E,1", SmaractCommunicationError, True), - (b":E,-1", SmaractCommunicationError, True), + (b":E0,1\n", SmaractErrorCode, True), + (b":E,1\n", SmaractCommunicationError, True), + (b":E,-1\n", SmaractCommunicationError, True), ], ) def test_socket_put_and_receive_raises_exception(controller, return_msg, exception, raised): controller.sock.buffer_recv = return_msg with pytest.raises(exception): - controller.socket_put_and_receive(b"test", raise_if_not_status=True) + controller.socket_put_and_receive("test", raise_if_not_status=True) controller.sock.flush_buffer() controller.sock.buffer_recv = return_msg if raised: with pytest.raises(exception): - controller.socket_put_and_receive(b"test") + controller.socket_put_and_receive("test") else: - assert controller.socket_put_and_receive(b"test") == return_msg.split(b"\n")[0].decode() + assert controller.socket_put_and_receive("test") == return_msg.split(b"\n")[0].decode() @pytest.mark.parametrize( @@ -120,7 +120,7 @@ def test_communication_mode(controller, mode, get_message, return_msg): (0, b":GS0\n", b":S0,6\n"), (1, b":GS0\n", b":S0,7\n"), (0, b":GS0\n", b":S0,9\n"), - (0, [b":GS0\n", b":GS0\n"], [b":E0,0\n", b":S0,9"]), + (0, [b":GS0\n", b":GS0\n"], [b":E0,0\n", b":S0,9\n"]), ], ) def test_axis_is_moving(controller, is_moving, get_message, return_msg): @@ -149,9 +149,9 @@ def test_get_sensor_definition(controller, sensor_id, axis, get_msg, return_msg) @pytest.mark.parametrize( "move_speed,axis,get_msg,return_msg", [ - (50, 0, b":SCLS0,50000000\n", b":E-1,0"), - (0, 0, b":SCLS0,0\n", b":E-1,0"), - (20.23, 1, b":SCLS1,20230000\n", b":E-1,0"), + (50, 0, b":SCLS0,50000000\n", b":E-1,0\n"), + (0, 0, b":SCLS0,0\n", b":E-1,0\n"), + (20.23, 1, b":SCLS1,20230000\n", b":E-1,0\n"), ], ) def test_set_move_speed(controller, move_speed, axis, get_msg, return_msg): @@ -163,9 +163,9 @@ def test_set_move_speed(controller, move_speed, axis, get_msg, return_msg): @pytest.mark.parametrize( "pos,axis,hold_time,get_msg,return_msg", [ - (50, 0, None, b":MPA0,50000000,1000\n", b":E0,0"), - (0, 0, 800, b":MPA0,0,800\n", b":E0,0"), - (20.23, 1, None, b":MPA1,20230000,1000\n", b":E0,0"), + (50, 0, None, b":MPA0,50000000,1000\n", b":E0,0\n"), + (0, 0, 800, b":MPA0,0,800\n", b":E0,0\n"), + (20.23, 1, None, b":MPA1,20230000,1000\n", b":E0,0\n"), ], ) def test_move_axis_to_absolute_position(controller, pos, axis, hold_time, get_msg, return_msg): @@ -209,7 +209,7 @@ def test_move_axis(lsmarA, pos, get_msg, return_msg): assert controller.sock.buffer_put == get_msg -@pytest.mark.parametrize("num_axes,get_msg,return_msg", [(1, [b":S0\n"], [b":E0,0"])]) +@pytest.mark.parametrize("num_axes,get_msg,return_msg", [(1, [b":S0\n"], [b":E0,0\n"])]) def test_stop_axis(lsmarA, num_axes, get_msg, return_msg): controller = lsmarA.controller controller.sock.buffer_recv = return_msg -- 2.54.0 From 35b1428f3c04c86024f5a64fe276e533b1bf3993 Mon Sep 17 00:00:00 2001 From: wakonig_k Date: Tue, 15 Sep 2026 12:03:38 +0200 Subject: [PATCH 2/2] fix(npoint): cleanup and fix for new controller base class --- csaxs_bec/devices/npoint/npoint.py | 301 +++++------------------ tests/tests_devices/test_npoint_piezo.py | 99 +++++--- 2 files changed, 129 insertions(+), 271 deletions(-) diff --git a/csaxs_bec/devices/npoint/npoint.py b/csaxs_bec/devices/npoint/npoint.py index 0ab35893..e570d358 100644 --- a/csaxs_bec/devices/npoint/npoint.py +++ b/csaxs_bec/devices/npoint/npoint.py @@ -1,4 +1,3 @@ -import functools import threading import time @@ -8,25 +7,13 @@ from ophyd import Component as Cpt from ophyd import Device, PositionerBase, Signal, SignalRO from ophyd.status import wait as status_wait from ophyd.utils import LimitError, ReadOnlyError -from ophyd_devices.utils.controller import Controller, threadlocked +from ophyd_devices.utils.controller import Controller, ControllerCommunicationError, threadlocked from ophyd_devices.utils.socket import SocketIO, SocketSignal, raise_if_disconnected from prettytable import PrettyTable logger = bec_logger.logger -def channel_checked(fcn): - """Decorator to catch attempted access to channels that are not available.""" - - @functools.wraps(fcn) - def wrapper(self, *args, **kwargs): - # pylint: disable=protected-access - self._check_channel(args[0]) - return fcn(self, *args, **kwargs) - - return wrapper - - class NpointError(Exception): """ Base class for Npoint errors. @@ -40,11 +27,17 @@ class NPointController(Controller): """ _axes_per_controller = 3 - _read_single_loc_bit = "A0" - _write_single_loc_bit = "A2" - _trailing_bit = "55" - _range_offset = "78" - _channel_base = ["11", "83"] + READ = b"\xa0" + WRITE = b"\xa2" + END = b"\x55" + + CHANNEL_STRIDE = 0x1000 + RANGE = 0x11831078 + POSITION = 0x11831334 + TARGET = 0x11831218 + SERVO = 0x11831084 + + COUNTS_PER_UM = 1_048_574 / 100 def show_all(self) -> None: """Display current status of all channels @@ -62,244 +55,70 @@ class NPointController(Controller): t.add_row([ii, self._get_range(ii), self.get_current_pos(ii), self.get_target_pos(ii)]) print(t) - @channel_checked - def _get_range(self, channel: int) -> int: - """Get the range of the specified channel axis. - - Args: - channel (int): Channel for which the range should be requested. - - Raises: - RuntimeError: Raised if the received message doesn't have the expected number of bytes (10). - - Returns: - int: Range - """ - - # for first channel: 0x11 83 10 78 - addr = self._channel_base.copy() - addr.extend([f"{16 + 16 * channel:x}", self._range_offset]) - send_buffer = self.__read_single_location_buffer(addr) - - recvd = self._put_and_receive(send_buffer) - if len(recvd) != 10: - raise RuntimeError( - f"Received buffer is corrupted. Expected 10 bytes and instead got {len(recvd)}" - ) - device_range = self._hex_list_to_int(recvd[5:-1], signed=False) - return device_range - - @channel_checked def get_current_pos(self, channel: int) -> float: - # for first channel: 0x11 83 13 34 - addr = self._channel_base.copy() - addr.extend([f"{19 + 16 * channel:x}", "34"]) - send_buffer = self.__read_single_location_buffer(addr) + """ + Return the current channel position in micrometres. - recvd = self._put_and_receive(send_buffer) + Args: + channel (int): The channel number (0, 1, or 2). + """ + return self._read_register(self.POSITION, channel) / self.COUNTS_PER_UM - pos_buffer = recvd[5:-1] - pos = self._hex_list_to_int(pos_buffer) / 1048574 * 100 - return pos - - @channel_checked - def set_target_pos(self, channel: int, pos: float) -> None: - # for first channel: 0x11 83 12 18 00 00 00 00 - addr = self._channel_base.copy() - addr.extend([f"{18 + channel * 16:x}", "18"]) - - target = int(round(1048574 / 100 * pos)) - data = [f"{m:02x}" for m in target.to_bytes(4, byteorder="big", signed=True)] - - send_buffer = self.__write_single_location_buffer(addr, data) - self._put(send_buffer) - - @channel_checked def get_target_pos(self, channel: int) -> float: - # for first channel: 0x11 83 12 18 - addr = self._channel_base.copy() - addr.extend([f"{18 + channel * 16:x}", "18"]) - send_buffer = self.__read_single_location_buffer(addr) - - recvd = self._put_and_receive(send_buffer) - pos_buffer = recvd[5:-1] - pos = self._hex_list_to_int(pos_buffer) / 1048574 * 100 - return pos - - @channel_checked - def _set_servo(self, channel: int, enable: bool) -> None: - print("Not tested") - return - # # for first channel: 0x11 83 10 84 00 00 00 00 - # addr = self._channel_base.copy() - # addr.extend([f"{16 + channel * 16:x}", "84"]) - - # if enable: - # data = ["00"] * 3 + ["01"] - # else: - # data = ["00"] * 4 - # send_buffer = self.__write_single_location_buffer(addr, data) - - # self._put(send_buffer) - - @channel_checked - def _get_servo(self, channel: int) -> int: - # for first channel: 0x11 83 10 84 00 00 00 00 - addr = self._channel_base.copy() - addr.extend([f"{16 + channel * 16:x}", "84"]) - send_buffer = self.__read_single_location_buffer(addr) - - recvd = self._put_and_receive(send_buffer) - buffer = recvd[5:-1] - status = self._hex_list_to_int(buffer) - return status - - @threadlocked - def _put(self, buffer: list) -> None: - """Translates a list of hex values to bytes and sends them to the socket. + """ + Return the target channel position in micrometres. Args: - buffer (list): List of hex values without leading 0x - - Returns: - None + channel (int): The channel number (0, 1, or 2). """ + return self._read_register(self.TARGET, channel) / self.COUNTS_PER_UM - buffer = b"".join([bytes.fromhex(m) for m in buffer]) - self.sock.put(buffer) - - @threadlocked - def _put_and_receive(self, msg_hex_list: list) -> list: - """Send msg to socket and wait for a reply. + def set_target_pos(self, channel: int, pos: float) -> None: + """ + Set the target channel position in micrometres. Args: - msg_hex_list (list): List of hex values without leading 0x. - - Returns: - list: Received message as a list of hex values + channel (int): The channel number (0, 1, or 2). + pos (float): The target position in micrometres. """ - - buffer = b"".join([bytes.fromhex(m) for m in msg_hex_list]) - self.sock.put(buffer) - recv_msg = self.sock.receive() - recv_hex_list = [hex(m) for m in recv_msg] - self._verify_received_msg(msg_hex_list, recv_hex_list) - return recv_hex_list - - def _verify_received_msg(self, in_list: list, out_list: list) -> None: - """Ensure that the first address bits of sent and received messages are the same. - - Args: - in_list (list): list containing the sent message - out_list (list): list containing the received message - - Raises: - RuntimeError: Raised if first two address bits of 'in' and 'out' are not identical - - Returns: - None - """ - - # first, translate hex (str) values to int - in_list_int = [int(val, 16) for val in in_list] - out_list_int = [int(val, 16) for val in out_list] - - # first ints of the reply should be the same. Otherwise something went wrong - if not in_list_int[:2] == out_list_int[:2]: - raise RuntimeError("Connection failure. Please restart the controller.") + self._write_register(self.TARGET, channel, round(pos * self.COUNTS_PER_UM)) def _check_channel(self, channel: int) -> None: - if channel >= self._axes_per_controller: - raise ValueError( - f"Channel {channel+1} exceeds the available number of channels ({self._axes_per_controller})" + if not 0 <= channel < self._axes_per_controller: + raise ValueError(f"Invalid channel: {channel}") + + def _register_address(self, register: int, channel: int) -> bytes: + self._check_channel(channel) + return (register + channel * self.CHANNEL_STRIDE).to_bytes(4, "little") + + def _read_register(self, register: int, channel: int, *, signed: bool = True) -> int: + """Read a 32-bit register, validating the opcode, address and end marker.""" + header = self.READ + self._register_address(register, channel) + response = self.socket_put_and_receive(header + self.END, response_length=10) + if response[:5] != header or response[-1:] != self.END: + raise NpointError( + f"Invalid response for {header.hex(' ')}: {response.hex(' ')}. " + "Connection failure. Please restart the controller." ) + return int.from_bytes(response[5:9], "little", signed=signed) - @staticmethod - def _hex_list_to_int(in_buffer: list, byteorder="little", signed=True) -> int: - """Translate hex list to int. + def _write_register(self, register: int, channel: int, value: int) -> None: + """Write a signed 32-bit register. The controller sends no reply.""" + address = self._register_address(register, channel) + payload = value.to_bytes(4, "little", signed=True) + self.socket_put(self.WRITE + address + payload + self.END) - Args: - in_buffer (list): Input buffer; received as list of hex values - byteorder (str, optional): Byteorder of in_buffer. Defaults to "little". - signed (bool, optional): Whether the hex list represents a signed int. Defaults to True. + def _get_range(self, channel: int) -> int: + """Return the channel range as an unsigned register value.""" + return self._read_register(self.RANGE, channel, signed=False) - Returns: - int: Translated integer. - """ - if byteorder == "little": - in_buffer.reverse() + def _set_servo(self, channel: int, enable: bool) -> None: + """Servo writes remain disabled until the command has been verified.""" + self._check_channel(channel) + print("Not tested") - # make sure that all hex strings have the same format ("FF") - val_hex = [f"{int(m, 16):02x}" for m in in_buffer] - - val_bytes = [bytes.fromhex(m) for m in val_hex] - val = int.from_bytes(b"".join(val_bytes), byteorder="big", signed=signed) - return val - - @staticmethod - def __read_single_location_buffer(addr) -> list: - """Prepare buffer for reading from a single memory location (hex address). - Number of bytes: 6 - Format: 0xA0 [addr] 0x55 - Return Value: 0xA0 [addr] [data] 0x55 - Sample Hex Transmission from PC to LC.400: A0 18 12 83 11 55 - Sample Hex Return Transmission from LC.400 to PC: A0 18 12 83 11 64 00 00 00 55 - - Args: - addr (list): Hex address to read from - - Returns: - list: List of hex values representing the read instruction. - """ - buffer = [] - buffer.append(NPointController._read_single_loc_bit) - if isinstance(addr, list): - addr.reverse() - buffer.extend(addr) - else: - buffer.append(addr) - buffer.append(NPointController._trailing_bit) - - return buffer - - @staticmethod - def __write_single_location_buffer(addr: list, data: list) -> list: - """Prepare buffer for writing to a single memory location (hex address). - Number of bytes: 10 - Format: 0xA2 [addr] [data] 0x55 - Return Value: none - Sample Hex Transmission from PC to C.400: A2 18 12 83 11 E8 03 00 00 55 - - Args: - addr (list): List of hex values representing the address to write to. - data (list): List of hex values representing the data that should be written. - - Returns: - list: List of hex values representing the write instruction. - """ - buffer = [] - buffer.append(NPointController._write_single_loc_bit) - if isinstance(addr, list): - addr.reverse() - buffer.extend(addr) - else: - buffer.append(addr) - - if isinstance(data, list): - data.reverse() - buffer.extend(data) - else: - buffer.append(data) - buffer.append(NPointController._trailing_bit) - return buffer - - @staticmethod - def __read_array(): - raise NotImplementedError - - @staticmethod - def __write_next_command(): - raise NotImplementedError + def _get_servo(self, channel: int) -> int: + return self._read_register(self.SERVO, channel) def __del__(self): if self.connected: @@ -452,7 +271,7 @@ class NPointAxis(Device, PositionerBase): try: self.controller.on(timeout=timeout) self._update_setpoint_from_readback() - except TimeoutError: + except (ControllerCommunicationError, TimeoutError): self.controller.off(update_config=False) time.sleep(1) else: @@ -463,7 +282,7 @@ class NPointAxis(Device, PositionerBase): f"Try to reload the config and if the problem persists, check the connection to the nPoint controller " f"and ensure that it is powered on and accessible at {self.controller._socket_host}:{self.controller._socket_port}." ) - + def _update_setpoint_from_readback(self): """ The setpoint is only stored locally. After a restart, diff --git a/tests/tests_devices/test_npoint_piezo.py b/tests/tests_devices/test_npoint_piezo.py index 85e6e545..f9682d8c 100644 --- a/tests/tests_devices/test_npoint_piezo.py +++ b/tests/tests_devices/test_npoint_piezo.py @@ -1,9 +1,9 @@ -import copy from unittest import mock import pytest from csaxs_bec.devices.npoint import NPointAxis, NPointController +from csaxs_bec.devices.npoint.npoint import NpointError # pylint: disable=protected-access # pylint: disable=redefined-outer-name @@ -101,8 +101,8 @@ def test_axis_get_out(npointx, pos, msg_in, msg_out): "axis, msg_in, msg_out", [ (0, b"\xa04\x13\x83\x11U", b"\xa0\x34\x13\x83\x11\xcd\xcc\x00\x00U"), - (1, b"\xa04#\x83\x11U", b"\xa0\x34\x13\x83\x11\x00\x00\x00\x00U"), - (2, b"\xa043\x83\x11U", b"\xa0\x34\x13\x83\x1133\xff\xffU"), + (1, b"\xa04#\x83\x11U", b"\xa0\x34\x23\x83\x11\x00\x00\x00\x00U"), + (2, b"\xa043\x83\x11U", b"\xa0\x34\x33\x83\x1133\xff\xffU"), ], ) def test_axis_get_in(npointx, axis, msg_in, msg_out): @@ -130,48 +130,32 @@ def test_axis_out_of_range(dm_with_devices): ) -def test_get_axis_out_of_range(controller): +@pytest.mark.parametrize("channel", [-1, 3]) +def test_get_axis_out_of_range(controller, channel): """ Test that an error is raised when trying to get the current position of an invalid axis. """ with pytest.raises(ValueError): - controller.get_current_pos(3) + controller.get_current_pos(channel=channel) + controller.sock.put.assert_not_called() -def test_set_axis_out_of_range(controller): +@pytest.mark.parametrize("channel", [-1, 3]) +def test_set_axis_out_of_range(controller, channel): """ Test that an error is raised when trying to set the target position of an invalid axis. """ with pytest.raises(ValueError): - controller.set_target_pos(3, 5) - - -@pytest.mark.parametrize( - "in_buffer, byteorder, signed, val", - [ - (["0x0", "0x0", "0xcc", "0xcd"], "big", True, 52429), - (["0xcd", "0xcc", "0x0", "0x0"], "little", True, 52429), - (["cd", "cc", "00", "00"], "little", True, 52429), - ], -) -def test_hex_list_to_int(in_buffer, byteorder, signed, val): - """ - Test that the hex list is correctly converted to an integer - """ - assert ( - NPointController._hex_list_to_int( - copy.deepcopy(in_buffer), byteorder=byteorder, signed=signed - ) - == val - ) + controller.set_target_pos(channel=channel, pos=5) + controller.sock.put.assert_not_called() @pytest.mark.parametrize( "axis, msg_in, msg_out", [ - (0, b"\xa0x\x10\x83\x11U", b"\xa0\x78\x13\x83\x11\x64\x00\x00\x00U"), - (1, b"\xa0x \x83\x11U", b"\xa0\x78\x13\x83\x11\x64\x00\x00\x00U"), - (2, b"\xa0x0\x83\x11U", b"\xa0\x78\x13\x83\x11\x64\x00\x00\x00U"), + (0, b"\xa0x\x10\x83\x11U", b"\xa0\x78\x10\x83\x11\x64\x00\x00\x00U"), + (1, b"\xa0x \x83\x11U", b"\xa0\x78\x20\x83\x11\x64\x00\x00\x00U"), + (2, b"\xa0x0\x83\x11U", b"\xa0\x78\x30\x83\x11\x64\x00\x00\x00U"), ], ) def test_get_range(npointx, axis, msg_in, msg_out): @@ -183,3 +167,58 @@ def test_get_range(npointx, axis, msg_in, msg_out): val = npointx.controller._get_range(axis) npointx.controller.sock.put.assert_called_once_with(msg_in) assert val == 100 + + +@pytest.mark.parametrize( + "response", + [ + b"\xa2\x34\x13\x83\x11\x00\x00\x00\x00U", # Wrong opcode + b"\xa0\x34\x23\x83\x11\x00\x00\x00\x00U", # Wrong channel + b"\xa0\x34\x13\x82\x11\x00\x00\x00\x00U", # Wrong address + b"\xa0\x34\x13\x83\x11\x00\x00\x00\x00X", # Wrong end marker + ], +) +def test_read_rejects_invalid_frame(controller, response): + controller.sock.receive.side_effect = [response] + with pytest.raises(NpointError, match="Invalid response"): + controller.get_current_pos(0) + + +def test_read_fragmented_payload_with_end_marker(controller): + controller.sock.receive.side_effect = [b"\xa0\x34\x13\x83\x11U", b"\x00\x00", b"\x00U"] + assert controller.get_current_pos(0) == pytest.approx(85 / 1048574 * 100) + assert [call.kwargs["buffer_length"] for call in controller.sock.receive.call_args_list] == [ + 10, + 4, + 2, + ] + + +@pytest.mark.parametrize( + "method,command,response,expected", + [ + ( + "get_target_pos", + b"\xa0\x18\x22\x83\x11U", + b"\xa0\x18\x22\x83\x1133\xff\xffU", + -52429 / 1048574 * 100, + ), + ("_get_servo", b"\xa0\x84\x20\x83\x11U", b"\xa0\x84\x20\x83\x11\x01\x00\x00\x00U", 1), + ( + "_get_range", + b"\xa0\x78\x20\x83\x11U", + b"\xa0\x78\x20\x83\x11\xff\xff\xff\xffU", + 4294967295, + ), + ], +) +def test_register_read_types(controller, method, command, response, expected): + controller.sock.receive.side_effect = [response] + assert getattr(controller, method)(channel=1) == pytest.approx(expected) + controller.sock.put.assert_called_once_with(command) + + +def test_target_write_channel_offset(controller): + controller.set_target_pos(channel=2, pos=-5) + controller.sock.put.assert_called_once_with(b"\xa2\x18\x32\x83\x1133\xff\xffU") + controller.sock.receive.assert_not_called() -- 2.54.0