diff --git a/scripts/camera_stat_thread.py b/scripts/camera_stat_thread.py index 5f444731..0120b739 100644 --- a/scripts/camera_stat_thread.py +++ b/scripts/camera_stat_thread.py @@ -1,7 +1,7 @@ import json import threading import time -from typing import Any, Dict +from typing import Any import cv2 import numpy as np @@ -33,7 +33,7 @@ class ImageStatsReceiver: self.connection_attempts = 0 self.last_message_time = None - def calculate_projections(self, image: np.ndarray) -> Dict[str, Any]: + def calculate_projections(self, image: np.ndarray) -> dict[str, Any]: """Calculate X and Y projections of the image""" # Convert to grayscale if color image if len(image.shape) == 3: @@ -99,7 +99,7 @@ class ImageStatsReceiver: "y_coords": list(range(len(y_projection))), } - def calculate_radial_integration(self, image: np.ndarray, num_bins: int = 50) -> Dict[str, Any]: + def calculate_radial_integration(self, image: np.ndarray, num_bins: int = 50) -> dict[str, Any]: """Calculate radial integration of the image""" # Convert to grayscale if color image if len(image.shape) == 3: @@ -153,7 +153,7 @@ class ImageStatsReceiver: "num_bins": num_bins, } - def calculate_image_stats(self, image: np.ndarray) -> Dict[str, Any]: + def calculate_image_stats(self, image: np.ndarray) -> dict[str, Any]: """Calculate comprehensive statistics for an image""" image_float = image.astype(np.float64) @@ -250,7 +250,7 @@ class ImageStatsReceiver: else: print(f"[{time.strftime('%H:%M:%S')}] Waiting for image data...") - def print_formatted_stats(self, stats: Dict[str, Any]): + def print_formatted_stats(self, stats: dict[str, Any]): """Print formatted statistics in a compact terminal format""" timestamp = time.strftime("%H:%M:%S", time.localtime(stats["timestamp"])) @@ -295,7 +295,7 @@ class ImageStatsReceiver: ) # Optional: Print per-channel stats if available - if any(k.startswith("mean_ch") for k in stats.keys()): + if any(k.startswith("mean_ch") for k in stats): channels = [] i = 0 while f"mean_ch{i}" in stats: @@ -327,7 +327,7 @@ class ImageStatsReceiver: } return None - def save_radial_profile_to_file(self, filename: str = None): + def save_radial_profile_to_file(self, filename: str | None = None): """Save the current radial profile to a file""" if filename is None: filename = f"radial_profile_{int(time.time())}.txt" @@ -343,13 +343,10 @@ class ImageStatsReceiver: f.write(f"# Center: {radial['center']}\n") f.write("# Radius(pixels)\tMean_Intensity\tStd_Intensity\tPixel_Count\n") - for i in range(len(radial["r_centers"])): - f.write( - f"{radial['r_centers'][i]:.2f}\t" + f.writelines(f"{radial['r_centers'][i]:.2f}\t" f"{radial['radial_profile'][i]:.2f}\t" f"{radial['radial_std'][i]:.2f}\t" - f"{radial['pixel_counts'][i]}\n" - ) + f"{radial['pixel_counts'][i]}\n" for i in range(len(radial["r_centers"]))) print(f"Radial profile saved to {filename}") return filename @@ -425,7 +422,7 @@ def get_radial_profile(): return None -def save_current_radial_profile(filename: str = None): +def save_current_radial_profile(filename: str | None = None): """Save the current radial profile to a file""" global stats_receiver diff --git a/scripts/gui_desginer.py b/scripts/gui_desginer.py index dacc544d..3e606353 100644 --- a/scripts/gui_desginer.py +++ b/scripts/gui_desginer.py @@ -1,19 +1,19 @@ +import math +import sys + +from PySide6.QtCore import QPointF, QRectF, Qt +from PySide6.QtGui import QColor, QFont, QFontMetrics, QLinearGradient, QPainter, QPen from PySide6.QtWidgets import ( QApplication, - QWidget, - QVBoxLayout, + QFrame, QHBoxLayout, QLabel, - QFrame, QPushButton, QScrollArea, QStackedWidget, + QVBoxLayout, + QWidget, ) -from PySide6.QtCore import Qt, QPointF, QRectF -from PySide6.QtGui import QPainter, QColor, QPen, QLinearGradient, QFont, QFontMetrics -import sys -import math - # --------------------------------------------------------------------------- # Colour palette diff --git a/scripts/redis_dump.py b/scripts/redis_dump.py index 8db28f8f..84148556 100644 --- a/scripts/redis_dump.py +++ b/scripts/redis_dump.py @@ -1,9 +1,10 @@ # python -import sys import argparse -from typing import Any, Dict, Tuple -import yaml +import sys +from typing import Any + import redis +import yaml def decode_bulk(value: Any): @@ -19,7 +20,7 @@ def decode_bulk(value: Any): return value -def fetch_key(r: redis.Redis, key: bytes) -> Tuple[str, Any]: +def fetch_key(r: redis.Redis, key: bytes) -> tuple[str, Any]: decode_bulk(key) t = r.type(key) if isinstance(t, bytes): @@ -66,7 +67,7 @@ def dump_redis_to_yaml( ): r = redis.Redis(host=host, port=port, db=db, password=password) - dump: Dict[str, Any] = {"meta": {"host": host, "port": port, "db": db}, "data": {}} + dump: dict[str, Any] = {"meta": {"host": host, "port": port, "db": db}, "data": {}} cursor = 0 while True: diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py index 2ceeec9a..fd119bb6 100644 --- a/src/aare/daq/aaredb.py +++ b/src/aare/daq/aaredb.py @@ -2,7 +2,6 @@ import datetime import io import json import os -from typing import List, Optional import aareDB import cv2 @@ -76,7 +75,7 @@ class AareWrapper: self._key_file = configuration.key_file @log_timing(logger, "AareDB call") - def set_pucks_beamline(self, input_list: List[PuckLoadedInfo]): + def set_pucks_beamline(self, input_list: list[PuckLoadedInfo]): o = [] for i in input_list: @@ -103,7 +102,7 @@ class AareWrapper: @log_timing(logger, "AareDB call") def send_sample_event( - self, sample_id: StrictInt, event_type: SampleEventType, comment: Optional[str] = None + self, sample_id: StrictInt, event_type: SampleEventType, comment: str | None = None ) -> None: if sample_id is None or sample_id < 0: if sample_id is None: @@ -123,7 +122,7 @@ class AareWrapper: @log_timing(logger, "AareDB call") def upload_image( - self, sample_id: int, filename: str, bgr_image: np.ndarray, message: Optional[str] = None + self, sample_id: int, filename: str, bgr_image: np.ndarray, message: str | None = None ): _, buffer = cv2.imencode(".jpg", bgr_image) jpeg_bytes = io.BytesIO(buffer) @@ -145,7 +144,7 @@ class AareWrapper: logger.debug(f"Response status code: {response.status_code}") @log_timing(logger, "AareDB call") - def upload_jpg(self, sample_id: int, filename: str, jpg_image, message: Optional[str] = None): + def upload_jpg(self, sample_id: int, filename: str, jpg_image, message: str | None = None): logger.debug(f"jppg_image of type: {type(jpg_image)}") url = f"{self._host}/protected_router/sample_runner/{sample_id}/upload-images" headers = { @@ -165,7 +164,7 @@ class AareWrapper: @log_timing(logger, "AareDB call") def create_rotation_run( - self, s: Optional[SampleShortInfo], r: RotationScanRequest, d: DAQStatusModel + self, s: SampleShortInfo | None, r: RotationScanRequest, d: DAQStatusModel ): if s is None: return @@ -238,7 +237,7 @@ class AareWrapper: @log_timing(logger, "AareDB call") def create_gridscan_run( - self, s: Optional[SampleShortInfo], r: RasterGridRequest, d: DAQStatusModel + self, s: SampleShortInfo | None, r: RasterGridRequest, d: DAQStatusModel ): if s is None: return @@ -301,11 +300,11 @@ class AareWrapper: @log_timing(logger, "AareDB call") def ingest_gridscan( self, - sample: Optional[SampleShortInfo], + sample: SampleShortInfo | None, raster_result: ScanResult, raster_request: RasterGridRequest, geom: SampleGeometryModel, - com: Optional[CenterOfMassModel], + com: CenterOfMassModel | None, beam_mark_pxl: tuple[float, float], ): @@ -339,11 +338,11 @@ class AareWrapper: def format_gridscan_payload( self, - sample: Optional[SampleShortInfo], + sample: SampleShortInfo | None, raster_result: ScanResult, raster_request: RasterGridRequest, geom: SampleGeometryModel, - com: Optional[CenterOfMassModel], + com: CenterOfMassModel | None, beam_mark_pxl: tuple[float, float], ) -> RasterPayloadModel | None: @@ -398,12 +397,12 @@ class AareWrapper: except Exception as e: logger.error(e) - raise e + raise @log_timing(logger, "AareDB call") def ingest_scan( self, - sample: Optional[SampleShortInfo], + sample: SampleShortInfo | None, result: ScanResult, geom: SampleGeometryModel, beam_mark_pxl: tuple[float, float], @@ -437,7 +436,7 @@ class AareWrapper: def format_scan_payload( self, - sample: Optional[SampleShortInfo], + sample: SampleShortInfo | None, result: ScanResult, geom: SampleGeometryModel, beam_mark_pxl: tuple[float, float], @@ -455,4 +454,4 @@ class AareWrapper: except Exception as e: logger.error(e) - raise e + raise diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py index 927478b2..92c61938 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -5,7 +5,6 @@ import pwd import time import uuid from datetime import UTC, datetime, timedelta -from typing import List import jwt from aarecommon.errors.exception_handler import ( @@ -38,7 +37,7 @@ oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token") class TokenData(BaseModel): sub: str # Username - pgroups: List[str] + pgroups: list[str] session: int staff: bool = False diff --git a/src/aare/daq/beamcenterfit.py b/src/aare/daq/beamcenterfit.py index de0f815c..280a2160 100644 --- a/src/aare/daq/beamcenterfit.py +++ b/src/aare/daq/beamcenterfit.py @@ -1,5 +1,5 @@ -import numpy as np import cv2 +import numpy as np from scipy.optimize import curve_fit diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py index 47d71310..24ca2c9b 100644 --- a/src/aare/daq/config.py +++ b/src/aare/daq/config.py @@ -4,7 +4,6 @@ import json import time from dataclasses import asdict, is_dataclass from datetime import datetime -from typing import List, Tuple import numpy as np import redis @@ -263,7 +262,7 @@ class BeamlineConfig: holder_session = holder.session if holder is not None else None sessions: list[OpenGuiSessionInfo] = [] - for session_id in sorted((int(s) for s in session_ids)): + for session_id in sorted(int(s) for s in session_ids): payload = self._read_gui_session(session_id) if payload is not None: payload.holds_baton = payload.session == holder_session @@ -522,9 +521,7 @@ class BeamlineConfig: @property def commissioning_mode(self) -> bool: tmp = self._client.get(f"{self._bl}:commissioning_mode") - if tmp is None: - return False - return True + return tmp is not None @commissioning_mode.setter def commissioning_mode(self, commisioning_mode: bool) -> None: @@ -635,7 +632,7 @@ class BeamlineConfig: return float(np.log(lens_factor / (b * target_pixel_in_mm)) / a) @property - def beam_center(self) -> Tuple[float, float]: + def beam_center(self) -> tuple[float, float]: tmp_x = self._client.get(f"{self._bl}:beam_center_x") tmp_y = self._client.get(f"{self._bl}:beam_center_y") if tmp_x: @@ -649,7 +646,7 @@ class BeamlineConfig: return val_x, val_y @beam_center.setter - def beam_center(self, data: Tuple[float, float]): + def beam_center(self, data: tuple[float, float]): self._client.set(f"{self._bl}:beam_center_x", data[0]) self._client.set(f"{self._bl}:beam_center_y", data[1]) @@ -721,7 +718,7 @@ class BeamlineConfig: data_dict = json.loads(tmp) return SampleShortInfoList(**data_dict) - def spreadsheet_pgroup(self, pgroups: List[str]) -> SampleShortInfoList: + def spreadsheet_pgroup(self, pgroups: list[str]) -> SampleShortInfoList: sample = self.spreadsheet sample.s = list(filter(lambda x: x.user in pgroups, sample.s)) return sample diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 402b6bd9..0c6ee565 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -2,10 +2,10 @@ import copy import json import secrets import time -from datetime import datetime, timezone +from collections.abc import Callable +from datetime import UTC, datetime from math import ceil from pathlib import Path -from typing import Callable, List, Optional, Tuple from aarecommon.config.beamline import cfg_get from aarecommon.config.logger import setup_logger @@ -611,9 +611,8 @@ class AareDAQ: if ( tell_state.last_event_class == "Motion Sync" and tell_state.last_event_value == "Sample put on Puck" - ): - if state_ts is not None and state_ts >= started_at: - return True + ) and state_ts is not None and state_ts >= started_at: + return True for event in reversed(self._get_tell_events_from_redis()): if event.get("class") == "Motion Sync" and event.get("event") == "Sample put on Puck": @@ -943,10 +942,10 @@ class AareDAQ: def _handle_operation_error( self, operation: DAQOperation, - sample: Optional[SampleShortInfo], + sample: SampleShortInfo | None, error: Exception, event_type: SampleEventType = SampleEventType.FAILED, - additional_comment: Optional[str] = None, + additional_comment: str | None = None, ) -> None: """ Centralized databse maessage error handling for all operations. @@ -1006,7 +1005,7 @@ class AareDAQ: True if successful, False otherwise """ previous_sample = None - mount_started_at = datetime.now(timezone.utc) + mount_started_at = datetime.now(UTC) self._last_mount_error_message = "" try: @@ -1169,7 +1168,7 @@ class AareDAQ: step_size: int = 15, face_min_ratio: float = 0.3, report_error: bool = True, - sample: Optional[SampleShortInfo] = None, + sample: SampleShortInfo | None = None, ) -> FaceDetectionResult: """ Execute face detection sequence through the face detection service. @@ -1583,7 +1582,7 @@ class AareDAQ: # name = clean_filename(name) return name - def spreadsheet_params(self) -> tuple[Optional[SimpleScanParameters], str | None]: + def spreadsheet_params(self) -> tuple[SimpleScanParameters | None, str | None]: file_prefix = None if self.status.sample is None: @@ -1622,8 +1621,7 @@ class AareDAQ: new_res = 1 / ((1 / res) + 0.1) corrected_dtz = self.diffraction_geometry.calc_dtz_mm(new_res) logger.debug(f"corrected dtz: {corrected_dtz}") - if corrected_dtz < 108: - corrected_dtz = 108 + corrected_dtz = max(corrected_dtz, 108) params.dtz = round(corrected_dtz) osc = getattr(aaredb_params, "oscillation", None) @@ -1900,7 +1898,7 @@ class AareDAQ: logger.debug(f"Failed to change mounted sample: {e}") raise - def list_loaded_pucks(self) -> List[PuckLoadedInfo]: + def list_loaded_pucks(self) -> list[PuckLoadedInfo]: return [] def _auto_focus(self, settings: AutofocusSettings, settle_time_s: float = 1.0) -> float: @@ -2014,7 +2012,6 @@ class AareDAQ: ) self._devs.smargon_wait(timeout=180) - return def _build_fake_rotation_result(self, request: RotationScanRequest) -> CompletedRotationScan: start_angle = 0.0 @@ -2222,7 +2219,6 @@ class AareDAQ: except Exception: self._cfg.state_busy = False raise - pass def mark_beam(self, x_pxl: float, y_pxl: float): self._cfg.set_busy(BeamlineStateEnum.BeamLocation) @@ -2255,11 +2251,11 @@ class AareDAQ: return sample_geom @property - def beam_center(self) -> Tuple[float, float]: + def beam_center(self) -> tuple[float, float]: return self._cfg.beam_center @beam_center.setter - def beam_center(self, val: Tuple[float, float]): + def beam_center(self, val: tuple[float, float]): self._cfg.beam_center = val @property @@ -2329,7 +2325,7 @@ class AareDAQ: self._cfg.state_busy = False @log_timing(logger, "Auto loop center") - def auto_loop_center(self, sample: Optional[SampleShortInfo] = None) -> float: + def auto_loop_center(self, sample: SampleShortInfo | None = None) -> float: """ Automatically center the loop using ML-based detection. This performs a multi-step sequence including rotation and centering. @@ -2416,7 +2412,7 @@ class AareDAQ: ) def send_message_db( - self, db_id: int, event_type: SampleEventType, comment: Optional[str] = None + self, db_id: int, event_type: SampleEventType, comment: str | None = None ): self._aare.send_sample_event(db_id, event_type, comment) @@ -2524,7 +2520,7 @@ class AareDAQ: return params def get_collection_params(self, prefer_smart: bool = False) -> tuple[SimpleScanParameters, str]: - spreadsheet_params, file_prefix = self.spreadsheet_params() + spreadsheet_params, _file_prefix = self.spreadsheet_params() logger.debug(f"spreadsheet_params: {spreadsheet_params}") smart_params = self._cfg.auto_params default_params = SimpleScanParameters(exp_time_s=0.04, dtz=110, incr_omega_deg=0.2) @@ -2544,7 +2540,7 @@ class AareDAQ: def _end_operation( self, start: float, - operation: Optional[DAQOperation] = DAQOperation.AUTOMATION, + operation: DAQOperation | None = DAQOperation.AUTOMATION, error: bool = False, ) -> float: """ @@ -2603,9 +2599,7 @@ class AareDAQ: self._emit_automation_progress(progress) formatted_date = datetime.now().strftime("%Y%m%d") - sample_prefix = "{}/{}/{:02d}/{}".format( - formatted_date, sample.puck_name, sample.pin, sample.sample_name - ) + sample_prefix = f"{formatted_date}/{sample.puck_name}/{sample.pin:02d}/{sample.sample_name}" try: self._validate_automation_state(context="automation start") @@ -2697,9 +2691,7 @@ class AareDAQ: if raster_params.filename is not None: logger.info(f"Using filename {raster_params.filename}") - sample_prefix = "{filename}/{prefix}".format( - filename=raster_params.filename, prefix=sample.sample_name - ) + sample_prefix = f"{raster_params.filename}/{sample.sample_name}" raster_grid = RasterGridRequest( exp_time_s=raster_params.exp_time_s, @@ -2777,9 +2769,7 @@ class AareDAQ: if params.filename is not None: logger.info(f"Using filename {params.filename}") - sample_prefix = "{filename}/{prefix}".format( - filename=params.filename, prefix=sample.sample_name - ) + sample_prefix = f"{params.filename}/{sample.sample_name}" rotation_request = RotationScanRequest( start_omega_deg=start_omega, @@ -2883,7 +2873,7 @@ class AareDAQ: logger.error(f"Failed to mount sample: {e}") if e.critical: self._end_operation(start, operation=DAQOperation.AUTOMATION, error=True) - raise e + raise else: pass @@ -3324,7 +3314,7 @@ class AareDAQ: return self.beamline_status except Exception as e: # TODO add error message to send to GUI to say problem - logger.error(f"Failed to retrieve beamline status: {str(e)}") + logger.error(f"Failed to retrieve beamline status: {e!s}") raise def _safe_tell_state(self) -> TellStateModel | None: @@ -3352,7 +3342,7 @@ class AareDAQ: try: return self.diffraction_geometry except Exception as e: - logger.warning(f"Failed to retrieve diffraction geomtrey: {str(e)}") + logger.warning(f"Failed to retrieve diffraction geomtrey: {e!s}") # Must satisfy pydantic constraints in DiffractionGeometry return DiffractionGeometry( energy_keV=12.4, diff --git a/src/aare/daq/mlbox.py b/src/aare/daq/mlbox.py index 7fdb30fc..6a70ed08 100644 --- a/src/aare/daq/mlbox.py +++ b/src/aare/daq/mlbox.py @@ -1,6 +1,6 @@ import time +from collections.abc import Iterable from dataclasses import dataclass, field -from typing import Iterable, Optional import cv2 import numpy as np @@ -308,15 +308,13 @@ class MlBox: def box_filter_overlap( self, model: MLBoxModel, pin: MLBoxModel, overlap_parameter: float = 0.5 ) -> bool: - if self._check_overlap(model, pin) >= overlap_parameter: - return True - return False + return self._check_overlap(model, pin) >= overlap_parameter def _filter_predictions( self, predictions: MLOutputModel, - overlap_with_pin: Optional[float] = None, - confidence_min: Optional[float] = None, + overlap_with_pin: float | None = None, + confidence_min: float | None = None, ): pin = predictions.get_best_for_class(MLBoxType.PIN) @@ -337,7 +335,7 @@ class MlBox: predictions.boxes.pop(k, None) @staticmethod - def _best_by_class(results) -> Optional[MLOutputModel]: + def _best_by_class(results) -> MLOutputModel | None: out = MLOutputModel() for pred in results: try: @@ -375,7 +373,7 @@ class MlBox: @staticmethod def _all_from_prediction_model( prediction: LatestPredictionModel | None, - ) -> Optional[MLOutputModel]: + ) -> MLOutputModel | None: if prediction is None or not getattr(prediction, "boxes", None): return None @@ -406,9 +404,9 @@ class MlBox: @staticmethod def get_preferred_class_box_with_confidence_threshold( boxes: MLOutputModel, - preferred_class: Optional[Iterable[int] | int | MLBoxType] = None, + preferred_class: Iterable[int] | int | MLBoxType | None = None, loop_preference_margin: float = 0.1, - ) -> Optional[MLBoxModel]: + ) -> MLBoxModel | None: """ Get best box, preferring loops over pin even if pin has higher confidence, unless pin's confidence exceeds loops by the margin. @@ -466,8 +464,8 @@ class MlBox: @staticmethod def get_preferred_class_box( - boxes: MLOutputModel, preferred_class: Optional[Iterable[int] | int | MLBoxType] = None - ) -> Optional[MLBoxModel]: + boxes: MLOutputModel, preferred_class: Iterable[int] | int | MLBoxType | None = None + ) -> MLBoxModel | None: if preferred_class is None: order = (MLBoxType.CRYSTAL, MLBoxType.LOOP_FACE, MLBoxType.LOOP_ALL, MLBoxType.PIN) else: @@ -525,8 +523,8 @@ class MlBox: def predict_best_no_filter( self, return_image: bool = False, return_bundle_meta: bool = False ) -> ( - Optional[MLOutputModel] - | tuple[Optional[MLOutputModel], np.ndarray | None] + MLOutputModel | None + | tuple[MLOutputModel | None, np.ndarray | None] | MLBoxPredictionsResult ): ml_bundle = self._collect_best_bundle() @@ -553,8 +551,8 @@ class MlBox: return_image: bool = False, return_bundle_meta: bool = False, ) -> ( - Optional[MLOutputModel] - | tuple[Optional[MLOutputModel], np.ndarray | None] + MLOutputModel | None + | tuple[MLOutputModel | None, np.ndarray | None] | MLBoxPredictionsResult ): ml_bundle = self._collect_best_bundle() @@ -630,7 +628,7 @@ class MlBox: overlap_with_pin: float | None = None, confidence_min: float | None = None, return_image: bool = False, - ) -> Optional[MLBoxModel] | tuple[Optional[MLBoxModel], np.ndarray | None]: + ) -> MLBoxModel | None | tuple[MLBoxModel | None, np.ndarray | None]: """ Request bounding boxes for the next N bundle fetches and return the best available box. """ diff --git a/src/aare/daq/operations/common/ml_bounding_box.py b/src/aare/daq/operations/common/ml_bounding_box.py index be364a20..77188e2d 100644 --- a/src/aare/daq/operations/common/ml_bounding_box.py +++ b/src/aare/daq/operations/common/ml_bounding_box.py @@ -1,7 +1,7 @@ import time +from collections.abc import Callable from dataclasses import dataclass from math import ceil, floor -from typing import Callable from aarecommon.config.beamline import cfg_get from aarecommon.config.logger_events import ( @@ -73,8 +73,8 @@ def scale_auto_raster_grid( physical_size_y_mm = n_y * grid_size.y scale = (image_count / max_images) ** 0.5 - scaled_n_x = max(1, int(floor(n_x / scale))) - scaled_n_y = max(1, int(floor(n_y / scale))) + scaled_n_x = max(1, floor(n_x / scale)) + scaled_n_y = max(1, floor(n_y / scale)) while scaled_n_x * scaled_n_y > max_images: if scaled_n_x >= scaled_n_y and scaled_n_x > 1: @@ -201,9 +201,9 @@ def _box_to_raster_request( frac_x = float(cfg_get("daq.auto_raster.grid_padding_fraction_x", 0.15)) frac_y_top = float(cfg_get("daq.auto_raster.grid_padding_fraction_y", 0.15)) frac_y_bottom = float(cfg_get("daq.auto_raster.grid_padding_fraction_y_bottom", frac_y_top)) - pad_x = max(1, int(ceil(frac_x * n_x))) - pad_y_top = max(1, int(ceil(frac_y_top * n_y))) - pad_y_bottom = max(1, int(ceil(frac_y_bottom * n_y))) + pad_x = max(1, ceil(frac_x * n_x)) + pad_y_top = max(1, ceil(frac_y_top * n_y)) + pad_y_bottom = max(1, ceil(frac_y_bottom * n_y)) x1 = x1 - pad_x * grid_size.x / geom.pixel_in_mm y1 = y1 - pad_y_top * grid_size.y / geom.pixel_in_mm n_x = n_x + 2 * pad_x diff --git a/src/aare/daq/operations/face_detection/utils.py b/src/aare/daq/operations/face_detection/utils.py index fe127170..998da418 100644 --- a/src/aare/daq/operations/face_detection/utils.py +++ b/src/aare/daq/operations/face_detection/utils.py @@ -3,7 +3,6 @@ import math import statistics import time import warnings -from typing import Dict, List, Optional, Tuple import numpy as np from aarecommon.config.logger import setup_logger @@ -13,7 +12,7 @@ logger = setup_logger("aareDAQ") def box_height_from_tuple(box: tuple[float, float, float, float]) -> float: - x1, y1, x2, y2 = box + _x1, y1, _x2, y2 = box return abs(y2 - y1) @@ -24,7 +23,7 @@ def box_area_from_tuple(box: tuple[float, float, float, float]) -> float: def prepare_samples( boxes_by_angle: dict[int, tuple[float, float, float, float]], area=False -) -> List[Tuple[float, float]]: +) -> list[tuple[float, float]]: # angles in degrees -> (theta_rad, height) samples = [] for deg, box in boxes_by_angle.items(): @@ -37,7 +36,7 @@ def cos_model(theta_deg: float | np.ndarray, A: float, B: float, phi_rad: float, return A + B * np.cos(C * np.deg2rad(theta_deg) - phi_rad) -def mad_filter(samples: List[Tuple[float, float]], k: float = 3.5) -> List[Tuple[float, float]]: +def mad_filter(samples: list[tuple[float, float]], k: float = 3.5) -> list[tuple[float, float]]: if not samples: return samples ys = [y for _, y in samples] @@ -54,7 +53,6 @@ def samples_to_json(samples): } with open("cos_test.json", "w") as f: json.dump(output_data, f, indent=2) - return def fit_metrics(y_true: np.ndarray, y_pred: np.ndarray) -> tuple[float, float, float]: @@ -67,7 +65,7 @@ def fit_metrics(y_true: np.ndarray, y_pred: np.ndarray) -> tuple[float, float, f return rmse, mae, r2 -def fit_cosine(samples: List[Tuple[float, float]]) -> dict: +def fit_cosine(samples: list[tuple[float, float]]) -> dict: samples = mad_filter(samples, k=3.5) if len(samples) < 3: A = sum(y for _, y in samples) / max(1, len(samples)) @@ -136,10 +134,10 @@ def get_samples_out(boxes): def choose_best_fit( - fits_by_name: Dict[str, Dict], -) -> Tuple[Optional[float], Optional[Dict], Optional[str]]: + fits_by_name: dict[str, dict], +) -> tuple[float | None, dict | None, str | None]: - def key(entry: Dict): + def key(entry: dict): params = entry.get("params") or {} rmse = params.get("rmse") mae = params.get("mae") diff --git a/src/aare/daq/operations/loop_centering/__init__.py b/src/aare/daq/operations/loop_centering/__init__.py index 2026c242..f6d0c2be 100644 --- a/src/aare/daq/operations/loop_centering/__init__.py +++ b/src/aare/daq/operations/loop_centering/__init__.py @@ -1,17 +1,17 @@ +from aare.daq.operations.loop_centering.analyzer import LoopCenteringAnalyzer from aare.daq.operations.loop_centering.models import ( - LoopCenteringContext, AngleAnalysis, AttemptSummary, + LoopCenteringContext, LoopCenteringSettings, ) from aare.daq.operations.loop_centering.service import LoopCenteringService -from aare.daq.operations.loop_centering.analyzer import LoopCenteringAnalyzer __all__ = [ - "LoopCenteringContext", "AngleAnalysis", "AttemptSummary", - "LoopCenteringSettings", - "LoopCenteringService", "LoopCenteringAnalyzer", + "LoopCenteringContext", + "LoopCenteringService", + "LoopCenteringSettings", ] diff --git a/src/aare/daq/operations/raster/__init__.py b/src/aare/daq/operations/raster/__init__.py index 5f87f03c..06414f52 100644 --- a/src/aare/daq/operations/raster/__init__.py +++ b/src/aare/daq/operations/raster/__init__.py @@ -1,4 +1,4 @@ -from aare.daq.operations.raster.models import RasterContext, RasterBoundingBoxResult +from aare.daq.operations.raster.models import RasterBoundingBoxResult, RasterContext from aare.daq.operations.raster.service import RasterService __all__ = ["RasterBoundingBoxResult", "RasterContext", "RasterService"] diff --git a/src/aare/daq/operations/raster/service.py b/src/aare/daq/operations/raster/service.py index 3d4e7c56..8d296b61 100644 --- a/src/aare/daq/operations/raster/service.py +++ b/src/aare/daq/operations/raster/service.py @@ -325,7 +325,7 @@ class RasterService: padded_height_mm = ( box_height_pxl * geom.pixel_in_mm * (1.0 + 2.0 * y_padding_fraction_each_side) ) - n_y = max(1, int(ceil(padded_height_mm / grid_size_mm.y))) + n_y = max(1, ceil(padded_height_mm / grid_size_mm.y)) self.logger.info( "Computed second auto-center raster y size from ML box height", extra=merge_log_context( diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 9a305cdf..d8e756d4 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -4,8 +4,8 @@ import importlib import json import os import time +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager -from typing import AsyncGenerator, Optional import uvicorn from aarecommon.config.beamline import mx_beamline @@ -2487,7 +2487,7 @@ async def send_screenshot_db( async def send_message_db( db_id: int, event_type: SampleEventType, - comment: Optional[str] = None, + comment: str | None = None, token: str = Depends(oauth2_scheme), ): data = auth.parse_token(token) diff --git a/src/aare/daq/server_exception_handler.py b/src/aare/daq/server_exception_handler.py index 4aed5d8e..b78eed76 100644 --- a/src/aare/daq/server_exception_handler.py +++ b/src/aare/daq/server_exception_handler.py @@ -173,11 +173,10 @@ def register_exception_handlers(app) -> None: ) body = _error_body(exc, code_override=code_override) # TODO: migrate baton to its own error type - if isinstance(exc, UserRightsException): - if "do not hold the baton" in exc.message: - return JSONResponse( - status_code=status, content=body, headers=getattr(exc, "headers", None) - ) + if isinstance(exc, UserRightsException) and "do not hold the baton" in exc.message: + return JSONResponse( + status_code=status, content=body, headers=getattr(exc, "headers", None) + ) logger.warning( "AareAuthError: %s", body["message"], diff --git a/src/aare/daq/tell_state_machine.py b/src/aare/daq/tell_state_machine.py index 2bbe432d..53f8a97f 100644 --- a/src/aare/daq/tell_state_machine.py +++ b/src/aare/daq/tell_state_machine.py @@ -1,7 +1,7 @@ from __future__ import annotations import re -from datetime import datetime, timezone +from datetime import UTC, datetime from aarecommon.models.tell import TellActivityEnum, TellPhaseEnum, TellStateModel @@ -11,7 +11,7 @@ UNMOUNT_STATUS_RE = re.compile(r"^unmount:\s*") def _utc_now_iso() -> str: - return datetime.now(timezone.utc).isoformat() + return datetime.now(UTC).isoformat() def initial_tell_state() -> TellStateModel: diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index d459d203..d0ce3c06 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -4,7 +4,7 @@ import ssl import threading import time from collections import deque -from datetime import datetime, timezone +from datetime import UTC, datetime from typing import Any, cast import requests @@ -180,7 +180,7 @@ def set_tell_state_in_redis(state: TellStateModel) -> None: def record_tell_event(event_name, event_value): event_record = TellEventRecord( - timestamp=datetime.now(timezone.utc).isoformat(), class_=event_name, event=event_value + timestamp=datetime.now(UTC).isoformat(), class_=event_name, event=event_value ) latest_tell_events[event_name] = event_value diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py index 5970f5ad..10732d2c 100644 --- a/src/aare/devices/aerotech.py +++ b/src/aare/devices/aerotech.py @@ -1,4 +1,3 @@ -from typing import Optional, Union from aarecommon.config.beamline import cfg_get, mx_beamline from aarecommon.config.logger import setup_logger @@ -23,7 +22,7 @@ AEROTECH_HOME = AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0) logger = setup_logger("aareDAQ") -class AerotechController(object): +class AerotechController: def __init__(self, bl: MXBeamline): if bl == MXBeamline.X06DA: self._simulated = False @@ -185,9 +184,9 @@ class AerotechController(object): def rotation_scan( self, - rotation_deg: float | int, - time_sec: float | int, - start_pos_deg: float | int, + rotation_deg: float, + time_sec: float, + start_pos_deg: float, run_async: bool = False, ): payload = RotationRequest( @@ -212,11 +211,11 @@ class AerotechController(object): def grid_scan( self, grid_elem_count_y: int, - grid_elem_size_y_um: int | float, - time_sec: int | float, - grid_elem_size_x_um: Optional[Union[float, int]] = None, - grid_elem_count_x: Optional[int] = None, - run_async: Optional[bool] = False, + grid_elem_size_y_um: float, + time_sec: float, + grid_elem_size_x_um: float | None = None, + grid_elem_count_x: int | None = None, + run_async: bool | None = False, ): payload = GridRequest( grid_elem_count_x=grid_elem_count_x, @@ -240,9 +239,9 @@ class AerotechController(object): def screening_scan( self, - rotation_deg: float | int, - wedge_deg: float | int, - time_sec: float | int, + rotation_deg: float, + wedge_deg: float, + time_sec: float, steps: int, run_async: bool = False, ): diff --git a/src/aare/devices/area_detector.py b/src/aare/devices/area_detector.py index b5104ed2..51427404 100644 --- a/src/aare/devices/area_detector.py +++ b/src/aare/devices/area_detector.py @@ -21,7 +21,7 @@ class AutoExposureSettings: exp_max: float = 30000.000 -class epicsAD(object): +class epicsAD: def __init__(self, prefix, cam="cam1:", image="image1:"): self.img = None self.monitored = False @@ -63,7 +63,6 @@ class epicsAD(object): epics.ca.pend_io() except Exception as e: print(f"EPICS error: {e}") - pass def _init_auto_exp(self, settings: AutoExposureSettings = AutoExposureSettings()): self.acquire.put(0) diff --git a/src/aare/devices/bec_worker.py b/src/aare/devices/bec_worker.py index f6649bbe..945fe909 100644 --- a/src/aare/devices/bec_worker.py +++ b/src/aare/devices/bec_worker.py @@ -1,6 +1,5 @@ import time from enum import Enum -from typing import List, Optional from aarecommon.config.beamline import cfg_get, mx_beamline from aarecommon.config.logger import setup_logger @@ -118,7 +117,7 @@ class BECClientWorker: raise Exception(f"Error initialising BEC devices: {e}") def _raise_bec_error( - self, exc: Exception, *, operation: str, tags: Optional[List[str]] = None + self, exc: Exception, *, operation: str, tags: list[str] | None = None ) -> None: message = f"BEC operation '{operation}' failed: {type(exc).__name__}: {exc}" # logger.exception(message) @@ -156,7 +155,7 @@ class BECClientWorker: raise BECCommunicationError(message, operation=operation, exception=exc) from exc - def _set_scilog_tags(self, tags: Optional[List[str]] = None): + def _set_scilog_tags(self, tags: list[str] | None = None): try: if tags: self.client.messaging.scilog.set_default_tags(tags) @@ -171,13 +170,13 @@ class BECClientWorker: message: str, error: bool = False, warning: bool = False, - error_message: Optional[str] = None, + error_message: str | None = None, attachments=None, bold: bool = False, italic: bool = False, - color: Optional[str] = None, - additonal_text: Optional[List[str]] = None, - tags: Optional[List[str]] = None, + color: str | None = None, + additonal_text: list[str] | None = None, + tags: list[str] | None = None, ): if color and color not in ["red", "green", "yellow", "blue", "pink"]: logger.warning("specified color not in allowed list,using default") @@ -375,7 +374,7 @@ class BECClientWorker: energy_kev = energy_ev / 1000 return energy_kev - def change_energy(self, value: float | int, plot: bool = False): + def change_energy(self, value: float, plot: bool = False): current_energy = self.check_current_energy() logger.info(f"Current energy: {current_energy:.1f} eV") logger.info(f"Change energy requested: from {current_energy:.1f} to {value:.1f} eV") diff --git a/src/aare/devices/enum_pv.py b/src/aare/devices/enum_pv.py index f8219696..fd725e17 100755 --- a/src/aare/devices/enum_pv.py +++ b/src/aare/devices/enum_pv.py @@ -1,7 +1,7 @@ from enum import Enum from typing import Any -from aare.devices.set_get_pv import SetGetPV, MoveResult +from aare.devices.set_get_pv import MoveResult, SetGetPV class EnumPV(SetGetPV): diff --git a/src/aare/devices/fluorimeter.py b/src/aare/devices/fluorimeter.py index 13faa3e5..12c9f0a5 100644 --- a/src/aare/devices/fluorimeter.py +++ b/src/aare/devices/fluorimeter.py @@ -7,7 +7,7 @@ from epics import PV, poll logger = setup_logger("aareaDAQ") -class Fluorimeter(object): +class Fluorimeter: def __init__(self, beamline: MXBeamline, **kwargs): BEAMLINE = beamline.value.upper() diff --git a/src/aare/devices/jfjoch.py b/src/aare/devices/jfjoch.py index d4ecf250..46da0749 100644 --- a/src/aare/devices/jfjoch.py +++ b/src/aare/devices/jfjoch.py @@ -176,23 +176,22 @@ class JFJochWrapper: ) dataset_settings.xray_fluorescence_spectrum = xrf - if s.sample.aaredb_params: - if s.sample.aaredb_params.unitcell: - unit_cell_db = s.sample.aaredb_params.unitcell - unit_cell_split = unit_cell_db.replace(",", " ").split() - unit_cell_floats = [float(x) for x in unit_cell_split] - unit_cell = jfjoch_client.UnitCell( - a=unit_cell_floats[0], - b=unit_cell_floats[1], - c=unit_cell_floats[2], - alpha=unit_cell_floats[3], - beta=unit_cell_floats[4], - gamma=unit_cell_floats[5], - ) - dataset_settings.unit_cell = unit_cell - if s.sample.aaredb_params.spacegroupnumber: - space_group_number = s.sample.aaredb_params.spacegroupnumber - dataset_settings.space_group_number = space_group_number + if s.sample.aaredb_params and s.sample.aaredb_params.unitcell: + unit_cell_db = s.sample.aaredb_params.unitcell + unit_cell_split = unit_cell_db.replace(",", " ").split() + unit_cell_floats = [float(x) for x in unit_cell_split] + unit_cell = jfjoch_client.UnitCell( + a=unit_cell_floats[0], + b=unit_cell_floats[1], + c=unit_cell_floats[2], + alpha=unit_cell_floats[3], + beta=unit_cell_floats[4], + gamma=unit_cell_floats[5], + ) + dataset_settings.unit_cell = unit_cell + if s.sample.aaredb_params.spacegroupnumber: + space_group_number = s.sample.aaredb_params.spacegroupnumber + dataset_settings.space_group_number = space_group_number return dataset_settings @@ -230,7 +229,7 @@ class JFJochWrapper: def measure_raster(self, r: RasterGridRequest, s: DAQStatusModel, async_start: bool = True): self._start_scan(ScanTypeEnum.RASTER, r, s, async_start=async_start) - def wait_till_running(self, timeout: int | float = 60): + def wait_till_running(self, timeout: float = 60): if self._simulated: return None try: @@ -244,7 +243,7 @@ class JFJochWrapper: endpoint="wait_until_running_post", ) - def wait_till_done(self, timeout: int | float) -> jfjoch_client.models.ScanResult | None: + def wait_till_done(self, timeout: float) -> jfjoch_client.models.ScanResult | None: if self._simulated: return None try: diff --git a/src/aare/devices/mx_lib.py b/src/aare/devices/mx_lib.py index 09bf2eab..1a28678a 100644 --- a/src/aare/devices/mx_lib.py +++ b/src/aare/devices/mx_lib.py @@ -1,6 +1,7 @@ import re import time -from typing import Callable, Union, Any +from collections.abc import Callable +from typing import Any from epics import PV, Motor, poll @@ -21,12 +22,12 @@ def wait_for_movement_to_finish(*motors): longest = 0.0 for motor in motors: time_to_target = motor.readback / motor.slew_speed - longest = time_to_target if time_to_target > longest else longest + longest = max(longest, time_to_target) timeout = time.time() + 1.5 * longest done = False while not done and time.time() < timeout: - done = all([m.done_moving for m in motors]) + done = all(m.done_moving for m in motors) if time.time() > timeout: print("TIMEOUT waiting for motors to be done moving; current motor positions:") @@ -102,7 +103,7 @@ def is_epics_type(pv: PV, pv_type: str) -> bool: def wait_string_condition( - pv: PV, target: Union[str, re.Pattern], *, timeout: float = 60.0, polling: float = 0.1 + pv: PV, target: str | re.Pattern, *, timeout: float = 60.0, polling: float = 0.1 ): """wait until an epics.PV of type string reaches target :pv: epics.PV @@ -230,7 +231,7 @@ def wait_motor_position( def wait_enum_condition( - pv: PV, value: Union[str, int, re.Pattern], *, timeout: float = 60.0, polling=0.1 + pv: PV, value: str | int | re.Pattern, *, timeout: float = 60.0, polling=0.1 ): """wait until an epics.PV enum reaches value pv: epics.PV @@ -249,16 +250,16 @@ def wait_enum_condition( if not (isinstance(pv, PV) and pv.type.lower().endswith("enum")): raise AttributeError("argument 'pv' must be an epics.PV of type enum") - if not (isinstance(value, str) or isinstance(value, int) or isinstance(value, re.Pattern)): + if not (isinstance(value, (str, int, re.Pattern))): raise AttributeError("argument 'value' must be either an int, str, or re.Pattern") if type(value) is int: - tester = lambda pv: value == pv.get() # noqa: E731 + tester = lambda pv: value == pv.get() elif type(value) is str: - tester = lambda pv: str(value) == pv.get(as_string=True).lower() # noqa: E731 + tester = lambda pv: str(value) == pv.get(as_string=True).lower() value = str(value).lower() # it's already a str :-/ elif isinstance(value, re.Pattern): - tester = lambda pv: value.match(pv.get(as_string=True)) # noqa: E731 + tester = lambda pv: value.match(pv.get(as_string=True)) else: raise AttributeError("argument 'value' must be either an int, str, or re.Pattern") diff --git a/src/aare/devices/set_get_pv.py b/src/aare/devices/set_get_pv.py index 79fea92d..409b4038 100644 --- a/src/aare/devices/set_get_pv.py +++ b/src/aare/devices/set_get_pv.py @@ -1,9 +1,11 @@ from __future__ import annotations +from collections.abc import Callable, Mapping from dataclasses import dataclass -from typing import Any, Mapping, Optional, Union, Callable +from typing import Any, Union from epics import PV + from aare.devices.mx_lib import pv_wait RawValue = Union[str, float, int] @@ -16,7 +18,7 @@ ResolverValue = Union[ @dataclass class MoveResult: target: RawValue - name: Optional[str] = None + name: str | None = None class SetGetPV: diff --git a/src/aare/devices/smargon.py b/src/aare/devices/smargon.py index e870e277..9977b347 100644 --- a/src/aare/devices/smargon.py +++ b/src/aare/devices/smargon.py @@ -15,7 +15,7 @@ class SmargonMode(Enum): ERROR = 99 -class Smargon(object): +class Smargon: SMARGON_HOME = SmargonCoordinate(sh_mm=Coordinate(x=0, y=0, z=18), phi_deg=0, chi_deg=0) AERO_HOME = AerotechCoordinate(x=0, y=0, z=0, omega=0) @@ -175,13 +175,11 @@ class Smargon(object): target_string = "" if coord.sh_mm is not None: - target_string += "&SHX={:.5f}&SHY={:.5f}&SHZ={:.5f}".format( - coord.sh_mm.x, coord.sh_mm.y, coord.sh_mm.z - ) + target_string += f"&SHX={coord.sh_mm.x:.5f}&SHY={coord.sh_mm.y:.5f}&SHZ={coord.sh_mm.z:.5f}" if coord.chi_deg is not None: - target_string += "&CHI={:.5f}".format(coord.chi_deg) + target_string += f"&CHI={coord.chi_deg:.5f}" if coord.phi_deg is not None: - target_string += "&PHI={:.5f}".format(coord.phi_deg) + target_string += f"&PHI={coord.phi_deg:.5f}" if target_string: self.gonput(f"targetSCS?{target_string}") @@ -203,11 +201,9 @@ class Smargon(object): target_string = "" if coord.at_mm is not None: - target_string += "&GMX={:.5f}&GMY={:.5f}&GMZ={:.5f}".format( - coord.at_mm.x, coord.at_mm.y, coord.at_mm.z - ) + target_string += f"&GMX={coord.at_mm.x:.5f}&GMY={coord.at_mm.y:.5f}&GMZ={coord.at_mm.z:.5f}" if coord.omega_deg is not None: - target_string += "&GMU={:.5f}".format(coord.omega_deg) + target_string += f"&GMU={coord.omega_deg:.5f}" if target_string: self.gonput(f"targetAEROTECH?{target_string}") diff --git a/src/aare/devices/tell_backend.py b/src/aare/devices/tell_backend.py index 1ce95ff6..5a638305 100644 --- a/src/aare/devices/tell_backend.py +++ b/src/aare/devices/tell_backend.py @@ -1,7 +1,8 @@ import json import re import time -from typing import Any, Callable, Protocol +from collections.abc import Callable +from typing import Any, Protocol from urllib.parse import urlparse import requests diff --git a/src/aare/devices/tell_client.py b/src/aare/devices/tell_client.py index cf9ad4d8..c8590579 100644 --- a/src/aare/devices/tell_client.py +++ b/src/aare/devices/tell_client.py @@ -1,8 +1,8 @@ import ast import json import re +from datetime import UTC from enum import Enum -from typing import List from aarecommon.config.logger import setup_logger from aarecommon.errors.exception_handler import ( @@ -164,7 +164,7 @@ class TellClient: "Mount can't start: TELL doors are open", operation="mount_precheck" ) - def set_samples_info(self, info: List[PuckWithTellPosition]): + def set_samples_info(self, info: list[PuckWithTellPosition]): """sets the samples in the robot dewar based on the given list of PuckWithTellPosition objects and runs set_sample_info in the background""" @@ -205,9 +205,8 @@ class TellClient: result = self.get_result(self._last_cmd_id) logger.debug(f"getting result for command {self._last_cmd_id}: {result}") status = result["status"] - if "completed" != status: - if "removed" != status: - raise MountingFailed(f"{msg} {result}") + if "completed" != status and "removed" != status: + raise MountingFailed(f"{msg} {result}") return f"{msg} {result}" def estimate_mounting_time(self, segment) -> int: @@ -452,7 +451,7 @@ class TellClient: # return eval(status) return ast.literal_eval(status) - def get_detected_pucks(self) -> List[PuckLoadedInfo]: + def get_detected_pucks(self) -> list[PuckLoadedInfo]: j = json.loads(self.backend.eval("get_pucks_info()&")) output = [] @@ -482,7 +481,7 @@ class TellClient: return float(current) def set_current(self, current: float) -> float: - self.backend.eval("smart_magnet.set_current({:.1f})&".format(current)) + self.backend.eval(f"smart_magnet.set_current({current:.1f})&") current = self.backend.eval("smart_magnet.get_current_rb()&") return float(current) @@ -549,7 +548,7 @@ class TellClient: raise SmartMagnetFaultException except Exception as e: logger.error(f"check_smart_magnet_mounted failed: {e}") - raise e + raise def make_tell_client(bl: MXBeamline) -> TellClient: @@ -561,7 +560,7 @@ def make_tell_client(bl: MXBeamline) -> TellClient: if __name__ == "__main__": - from datetime import datetime, timezone + from datetime import datetime from aarecommon.config.beamline import mx_beamline @@ -574,8 +573,8 @@ if __name__ == "__main__": ts = float(tell_client.get_setting("dry_timestamp")) print("dry timestape: ", ts) - past = datetime.fromtimestamp(ts, tz=timezone.utc) - now = datetime.now(timezone.utc) + past = datetime.fromtimestamp(ts, tz=UTC) + now = datetime.now(UTC) seconds_ago = int((now - past).total_seconds()) print(seconds_ago) print("door closer :", tell_client.backend.eval("is_door_closed()&")) diff --git a/src/aare/devices/workflow_tools.py b/src/aare/devices/workflow_tools.py index 88cb8377..54675670 100755 --- a/src/aare/devices/workflow_tools.py +++ b/src/aare/devices/workflow_tools.py @@ -7,10 +7,10 @@ def wait_position(motor, target, tolerance=None, timeout=60.0): position = motor.readback if isinstance(position, float) and tolerance is None: tst = "%f == %s" - ltst = lambda x, y, z: x == y # noqa: E731 + ltst = lambda x, y, z: x == y elif isinstance(position, float) and tolerance is not None: tst = "abs(%f - %f) < %f" - ltst = lambda x, y, z: abs(x - y) < z # noqa: E731 + ltst = lambda x, y, z: abs(x - y) < z elif isinstance(position, (bytes, str)): tst = "'%s' == '%s'" if isinstance(position, bytes): @@ -38,18 +38,12 @@ def wait_position(motor, target, tolerance=None, timeout=60.0): if n > 20: n = 0 print( - "waiting_position test: %s (%s, %s, %s)" - % (tst, str(motor.readback), str(target), str(tolerance)) + f"waiting_position test: {tst} ({motor.readback!s}, {target!s}, {tolerance!s})" ) timeisup = timeout < time.time() condition = ltst(motor.readback, target, tolerance) if timeisup: - msg = "Timeout when waiting for %s to reach %s with tolerance %s. Device was at: %s" % ( - motor, - str(target), - str(tolerance), - str(motor.readback), - ) + msg = f"Timeout when waiting for {motor} to reach {target!s} with tolerance {tolerance!s}. Device was at: {motor.readback!s}" print(msg) raise RuntimeError(msg) diff --git a/src/aare/devices/zmq_client.py b/src/aare/devices/zmq_client.py index 2b4db0d9..15e52c9e 100755 --- a/src/aare/devices/zmq_client.py +++ b/src/aare/devices/zmq_client.py @@ -5,7 +5,6 @@ with fallback to area_detector if ZMQ is unavailable. """ import json -from typing import Optional import cv2 import numpy as np @@ -46,9 +45,9 @@ class ZMQCameraClient: raise Exception("unknown beamline") self._timeout_ms = timeout_ms - self._context: Optional[zmq.Context] = None - self._socket: Optional[zmq.Socket] = None - self._last_image: Optional[np.ndarray] = None + self._context: zmq.Context | None = None + self._socket: zmq.Socket | None = None + self._last_image: np.ndarray | None = None self._last_fetch_time: float = 0.0 self._connected = False @@ -81,7 +80,7 @@ class ZMQCameraClient: self._connected = False return False - def get_image(self, gray: bool = False) -> Optional[np.ndarray]: + def get_image(self, gray: bool = False) -> np.ndarray | None: """ Fetch the latest image from the ZMQ stream. diff --git a/src/aare/gui/about.py b/src/aare/gui/about.py index 0119fa04..c7d29e42 100644 --- a/src/aare/gui/about.py +++ b/src/aare/gui/about.py @@ -8,7 +8,7 @@ if TYPE_CHECKING: from aare.gui.threads.daq_worker import DAQWorker -def about_text(client: "DAQWorker"): +def about_text(client: DAQWorker): from aare.gui import gui try: diff --git a/src/aare/gui/gui.py b/src/aare/gui/gui.py index daaacbad..bd153f1a 100644 --- a/src/aare/gui/gui.py +++ b/src/aare/gui/gui.py @@ -173,7 +173,7 @@ def main(): QMessageBox.critical( None, "Authentication Error", - f"Cannot connect to AareDAQ server:\n{str(e)}\n\n Please check the server is running and your network connection.", + f"Cannot connect to AareDAQ server:\n{e!s}\n\n Please check the server is running and your network connection.", ) sys.exit(1) @@ -204,7 +204,7 @@ def main(): "Fatal Error", f"An error occurred during startup. See console for details." f"\nPlease check the server is running and your network connection." - f"\n\n{str(e)}\n\n", + f"\n\n{e!s}\n\n", ) sys.exit(1) diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index d6825979..c6cb96e6 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -2454,16 +2454,14 @@ class MainWindow(QMainWindow): status = self._latest_daq_status if status is not None and bool(status.busy): return False - if self._is_automation_active(): - return False - return True + return not self._is_automation_active() @Slot() def _check_remote_close_deadline(self) -> None: if self._remote_close_deadline_ts is None: return - remaining = int(round(self._remote_close_deadline_ts - time.time())) + remaining = round(self._remote_close_deadline_ts - time.time()) if remaining > 0: if self._remote_close_banner_active: self.alert_banner_secondary.show_message( diff --git a/src/aare/gui/models/gui_state_manager.py b/src/aare/gui/models/gui_state_manager.py index 2cdaa90f..91608a2b 100644 --- a/src/aare/gui/models/gui_state_manager.py +++ b/src/aare/gui/models/gui_state_manager.py @@ -1,4 +1,5 @@ import json + from PySide6.QtCore import QSettings diff --git a/src/aare/gui/models/sample_queue_model.py b/src/aare/gui/models/sample_queue_model.py index 3ceedc1a..1e5d8115 100644 --- a/src/aare/gui/models/sample_queue_model.py +++ b/src/aare/gui/models/sample_queue_model.py @@ -78,9 +78,7 @@ class SampleQueueSpreadsheet(QAbstractTableModel): return ["text/plain"] def canDropMimeData(self, data, action, row, column, parent): - if data.hasText(): - return True - return False + return bool(data.hasText()) def dropMimeData(self, data, action, row, column, parent): if not self.canDropMimeData(data, action, row, column, parent): diff --git a/src/aare/gui/models/user_sample_model.py b/src/aare/gui/models/user_sample_model.py index 9e410552..55e513cc 100644 --- a/src/aare/gui/models/user_sample_model.py +++ b/src/aare/gui/models/user_sample_model.py @@ -182,7 +182,7 @@ class UserSampleSpreadsheet(QAbstractTableModel): sample_data = SampleShortInfoList(s=[]) - for i in sorted(set(index.row() for index in indexes)): + for i in sorted({index.row() for index in indexes}): sample_data.s.append(self._sorted_samples[i]) mime_data.setText(sample_data.model_dump_json()) diff --git a/src/aare/gui/panels/automation_panel.py b/src/aare/gui/panels/automation_panel.py index e5bcf9cf..e9edb4c9 100644 --- a/src/aare/gui/panels/automation_panel.py +++ b/src/aare/gui/panels/automation_panel.py @@ -147,8 +147,8 @@ class AutomationProgressWidget(QWidget): if seconds is None or seconds <= 0: return "0m 00s" if seconds < 60: - return f"0m {int(round(seconds)):02d}s" - minutes, secs = divmod(int(round(seconds)), 60) + return f"0m {round(seconds):02d}s" + minutes, secs = divmod(round(seconds), 60) if minutes < 60: return f"{minutes}m {secs:02d}s" hours, minutes = divmod(minutes, 60) diff --git a/src/aare/gui/panels/axis_video_panel.py b/src/aare/gui/panels/axis_video_panel.py index bd8153a5..2bc5e5e8 100644 --- a/src/aare/gui/panels/axis_video_panel.py +++ b/src/aare/gui/panels/axis_video_panel.py @@ -1,5 +1,5 @@ from PySide6.QtCore import Qt, Signal -from PySide6.QtWidgets import QWidget, QVBoxLayout, QHBoxLayout, QLabel, QPushButton +from PySide6.QtWidgets import QHBoxLayout, QLabel, QPushButton, QVBoxLayout, QWidget from aare.gui.widgets.busy_overlay import BusyOverlayStyle from aare.gui.widgets.video_image import VideoGraphicsView diff --git a/src/aare/gui/panels/developer_help_dialog.py b/src/aare/gui/panels/developer_help_dialog.py index 905b9c13..aefd0b36 100644 --- a/src/aare/gui/panels/developer_help_dialog.py +++ b/src/aare/gui/panels/developer_help_dialog.py @@ -3,7 +3,6 @@ from __future__ import annotations import json import logging import time -from typing import Dict from aarecommon.config.logger import attach_to_logger, find_existing_formatter from aarecommon.errors.codes import error_code_help @@ -40,7 +39,7 @@ class DeveloperHelpDialog(QDialog): self._daq = daq self._is_staff = bool(is_staff) - self._codes: Dict[str, str] = {} + self._codes: dict[str, str] = {} self._last_payload: dict = {} self._freeze_payload: bool = False self._always_highlight_last_error: bool = True diff --git a/src/aare/gui/panels/loop_centering_panel.py b/src/aare/gui/panels/loop_centering_panel.py index d3b4b947..15d9a179 100644 --- a/src/aare/gui/panels/loop_centering_panel.py +++ b/src/aare/gui/panels/loop_centering_panel.py @@ -1,4 +1,4 @@ -from PySide6.QtWidgets import QWidget, QGridLayout, QPushButton +from PySide6.QtWidgets import QGridLayout, QPushButton, QWidget from aare.gui.widgets.title_label import TitleLabel diff --git a/src/aare/gui/panels/raster_data_collection.py b/src/aare/gui/panels/raster_data_collection.py index 3250f149..bbe4a16e 100644 --- a/src/aare/gui/panels/raster_data_collection.py +++ b/src/aare/gui/panels/raster_data_collection.py @@ -221,7 +221,7 @@ class RasterDataCollectionPanel(ScanSettingsPanel): def update_total_time_label(self): mins = int(self._total_time // 60) - secs = int(round(self._total_time % 60)) + secs = round(self._total_time % 60) if secs == 60: mins += 1 secs = 0 diff --git a/src/aare/gui/panels/reference_tools_panel.py b/src/aare/gui/panels/reference_tools_panel.py index 8fbf4882..711c582c 100644 --- a/src/aare/gui/panels/reference_tools_panel.py +++ b/src/aare/gui/panels/reference_tools_panel.py @@ -1,5 +1,4 @@ # reference_tools_panel.py -from typing import Optional from aarecommon.config.logger import setup_logger from aarecommon.models.models import DAQStatusModel, SampleShortInfo, SampleShortInfoList @@ -41,7 +40,7 @@ def get_entry(sample: SampleShortInfo, column: int): class ReferenceToolsModel(QAbstractTableModel): def __init__( self, - rows: Optional[list[SampleShortInfo]] | None = None, + rows: list[SampleShortInfo] | None = None, parent=None, current_reference: int | None = None, ): @@ -147,7 +146,7 @@ class ReferenceToolsModel(QAbstractTableModel): reverse=(self._sort_order == Qt.SortOrder.DescendingOrder), ) - def get_item(self, row: int) -> Optional[SampleShortInfo]: + def get_item(self, row: int) -> SampleShortInfo | None: """Get sample at the given row index.""" if 0 <= row < len(self._sorted_samples): return self._sorted_samples[row] @@ -227,7 +226,7 @@ class ReferenceToolsPanel(QFrame): self.table_view.setSelectionBehavior(QAbstractItemView.SelectionBehavior.SelectRows) self.table_view.setSelectionMode(QTableView.SelectionMode.SingleSelection) - def _selected_item(self) -> Optional[SampleShortInfo]: + def _selected_item(self) -> SampleShortInfo | None: idx = self.table_view.currentIndex() if not idx.isValid(): return None diff --git a/src/aare/gui/panels/rotation_data_collection.py b/src/aare/gui/panels/rotation_data_collection.py index c01538eb..0e8dacf4 100644 --- a/src/aare/gui/panels/rotation_data_collection.py +++ b/src/aare/gui/panels/rotation_data_collection.py @@ -244,7 +244,7 @@ class RotationDataCollectionPanel(ScanSettingsPanel): def update_total_time_label(self): mins = int(self._total_time // 60) - secs = int(round(self._total_time % 60)) + secs = round(self._total_time % 60) if secs == 60: mins += 1 secs = 0 diff --git a/src/aare/gui/panels/scan_settings_panel.py b/src/aare/gui/panels/scan_settings_panel.py index fcddb5ae..39285963 100644 --- a/src/aare/gui/panels/scan_settings_panel.py +++ b/src/aare/gui/panels/scan_settings_panel.py @@ -218,7 +218,7 @@ class ScanSettingsPanel(QWidget): def _res_to_dtz(self, res: float) -> float: dtz = self._diffraction.calc_dtz_mm(res) - return self.MIN_DTZ if dtz < self.MIN_DTZ else dtz + return max(dtz, self.MIN_DTZ) @Slot(float) def _on_dtz_value_changed(self, v: float): diff --git a/src/aare/gui/panels/smargon_trace_panel.py b/src/aare/gui/panels/smargon_trace_panel.py index e17c6d10..cd2bf95a 100644 --- a/src/aare/gui/panels/smargon_trace_panel.py +++ b/src/aare/gui/panels/smargon_trace_panel.py @@ -7,7 +7,7 @@ from pathlib import Path import numpy as np from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg as FigureCanvas from matplotlib.figure import Figure -from PySide6.QtCore import QSettings, QTimer, Qt +from PySide6.QtCore import QSettings, Qt, QTimer from PySide6.QtWidgets import ( QCheckBox, QComboBox, diff --git a/src/aare/gui/panels/smart_rotation_panel.py b/src/aare/gui/panels/smart_rotation_panel.py index 9f850093..1a49a873 100644 --- a/src/aare/gui/panels/smart_rotation_panel.py +++ b/src/aare/gui/panels/smart_rotation_panel.py @@ -278,7 +278,7 @@ class SimpleRotationSettingsPanel(QWidget): def update_total_time_label(self): mins = int(self.total_time_s // 60) - secs = int(round(self.total_time_s % 60)) + secs = round(self.total_time_s % 60) if secs == 60: mins += 1 secs = 0 @@ -349,8 +349,7 @@ class SimpleRotationSettingsPanel(QWidget): self.transmission = 1.0 self.dtz = self._d.diffraction.calc_dtz_mm(d_tar) - if self.dtz < 108: - self.dtz = 108 + self.dtz = max(self.dtz, 108) self.transmission_label.setText(f"{self.transmission * 100:.1f}") self.image_time_label.setText(f"{self.image_time_s:.4f}") @@ -364,7 +363,7 @@ class SimpleRotationSettingsPanel(QWidget): self.dtz_label.setText(f"{self.dtz:.2f}") self.parameters = SimpleScanParameters( - dtz=int(round(self.dtz)), + dtz=round(self.dtz), exp_time_s=self.image_time_s, start_omega_deg=self.start_angle_enter.value, incr_omega_deg=image_angle, diff --git a/src/aare/gui/panels/target_stability_panel.py b/src/aare/gui/panels/target_stability_panel.py index 7bfdf029..c5bea872 100644 --- a/src/aare/gui/panels/target_stability_panel.py +++ b/src/aare/gui/panels/target_stability_panel.py @@ -1402,9 +1402,8 @@ class TargetStabilityPanel(QWidget): @staticmethod def _coerce_target_point(raw) -> tuple[float, float] | None: try: - if isinstance(raw, dict): - if "x" in raw and "y" in raw: - return float(raw["x"]), float(raw["y"]) + if isinstance(raw, dict) and "x" in raw and "y" in raw: + return float(raw["x"]), float(raw["y"]) if isinstance(raw, (list, tuple)) and len(raw) >= 2: return float(raw[0]), float(raw[1]) except Exception as e: diff --git a/src/aare/gui/scan_logic/raster_grid_manager.py b/src/aare/gui/scan_logic/raster_grid_manager.py index 64b292c1..5016fdf8 100644 --- a/src/aare/gui/scan_logic/raster_grid_manager.py +++ b/src/aare/gui/scan_logic/raster_grid_manager.py @@ -1,6 +1,5 @@ import math from enum import Enum -from typing import List, Tuple import numpy as np from aarecommon.config.beamline import cfg_get, get_jfjoch_url, mx_beamline @@ -148,29 +147,20 @@ class RasterGridManager(QObject): ), dtz=cfg_get("daq.data_collection_settings.default_raster_settings.dtz", 200.0), ) - self._completed_grids: List[CompletedRasterGridElem] = [] + self._completed_grids: list[CompletedRasterGridElem] = [] # Cache of pre-rendered heatmap bitmaps, keyed by (id(grid_elem), metric). # Each entry is (QImage, backing ndarray); the ndarray must be kept alive # because QImage shares its buffer without copying. Rebuilt only when the # data or metric changes, not on every repaint (sample move / zoom). - self._heatmap_cache: dict[tuple[int, "RasterGridMetric"], tuple[QImage, np.ndarray]] = {} + self._heatmap_cache: dict[tuple[int, RasterGridMetric], tuple[QImage, np.ndarray]] = {} @property def active_grid(self) -> RasterGridRequest: return self._active_grid def _is_grid_visible(self, grid: RasterGridRequest): - if ( - grid.visible - and abs(normalize_angle(grid.omega_deg - self._geom.omega_deg)) < 0.2 - and abs(grid.smargon_top_left.phi_deg - self._geom.smargon.phi_deg) < 0.2 - and abs(grid.smargon_top_left.chi_deg - self._geom.smargon.chi_deg) < 0.2 - and grid.n_x > 0 - and grid.n_y > 0 - ): - return True - return False + return bool(grid.visible and abs(normalize_angle(grid.omega_deg - self._geom.omega_deg)) < 0.2 and abs(grid.smargon_top_left.phi_deg - self._geom.smargon.phi_deg) < 0.2 and abs(grid.smargon_top_left.chi_deg - self._geom.smargon.chi_deg) < 0.2 and grid.n_x > 0 and grid.n_y > 0) def _grid_pixel_geometry( self, grid: RasterGridRequest @@ -210,10 +200,10 @@ class RasterGridManager(QObject): if visible_rect is None or visible_rect.isEmpty(): return 0, grid.n_x, 0, grid.n_y - min_x = max(0, int(math.floor((visible_rect.left() - start_x) / cell_w)) - 1) - max_x = min(grid.n_x, int(math.ceil((visible_rect.right() - start_x) / cell_w)) + 1) - min_y = max(0, int(math.floor((visible_rect.top() - start_y) / cell_h)) - 1) - max_y = min(grid.n_y, int(math.ceil((visible_rect.bottom() - start_y) / cell_h)) + 1) + min_x = max(0, math.floor((visible_rect.left() - start_x) / cell_w) - 1) + max_x = min(grid.n_x, math.ceil((visible_rect.right() - start_x) / cell_w) + 1) + min_y = max(0, math.floor((visible_rect.top() - start_y) / cell_h) - 1) + max_y = min(grid.n_y, math.ceil((visible_rect.bottom() - start_y) / cell_h) + 1) return min_x, max_x, min_y, max_y @@ -300,7 +290,7 @@ class RasterGridManager(QObject): 0, 0, self._active_grid.grid_size_mm.x, self._active_grid.grid_size_mm.y ) - def get_grid_coord(self, grid: RasterGridRequest, point: QPointF) -> Tuple[int, int]: + def get_grid_coord(self, grid: RasterGridRequest, point: QPointF) -> tuple[int, int]: point_bl = self._geom.picture_to_sample(Coordinate(x=point.x(), y=point.y())) delta = point_bl - self._geom.smargon_to_beamline(grid.smargon_top_left.sh_mm) @@ -542,7 +532,7 @@ class RasterGridManager(QObject): self._heatmap_cache.clear() def _heatmap_image( - self, cache_key: tuple, grid: RasterGridRequest, values: List[float] | List[int] + self, cache_key: tuple, grid: RasterGridRequest, values: list[float] | list[int] ) -> QImage | None: """Return a cached n_x*n_y heatmap bitmap for this grid, building it once on a cache miss. One pixel per cell; colours baked at full opacity with @@ -560,7 +550,7 @@ class RasterGridManager(QObject): return built[0] def _build_heatmap_image( - self, n_x: int, n_y: int, values: List[float] | List[int] + self, n_x: int, n_y: int, values: list[float] | list[int] ) -> tuple[QImage, np.ndarray] | None: if n_x <= 0 or n_y <= 0: return None @@ -597,7 +587,7 @@ class RasterGridManager(QObject): self, painter: QPainter, elem: CompletedRasterGridElem, - values: List[float] | List[int], + values: list[float] | list[int], alpha: int, visible_rect: QRectF | None, cache_key: tuple, @@ -648,7 +638,7 @@ class RasterGridManager(QObject): self, painter: QPainter, grid: RasterGridRequest, - values: List[float] | List[int] | None = None, + values: list[float] | list[int] | None = None, alpha: int = 127, visible_rect: QRectF | None = None, ): @@ -657,9 +647,8 @@ class RasterGridManager(QObject): if alpha < 0 or alpha > 255: return - if values is None: - if self._draw_active_grid_fast(painter, grid, visible_rect): - return + if values is None and self._draw_active_grid_fast(painter, grid, visible_rect): + return geo = self._grid_pixel_geometry(grid) if geo is None: @@ -759,8 +748,8 @@ class RasterGridManager(QObject): painter.drawRect(bounds) min_spacing_px = 4.0 - stride_x = max(1, int(math.ceil(min_spacing_px / max(cell_w, 1e-9)))) - stride_y = max(1, int(math.ceil(min_spacing_px / max(cell_h, 1e-9)))) + stride_x = max(1, math.ceil(min_spacing_px / max(cell_w, 1e-9))) + stride_y = max(1, math.ceil(min_spacing_px / max(cell_h, 1e-9))) min_ix, max_ix, min_iy, max_iy = self._visible_index_range( grid, visible_rect, start_x, start_y, cell_w, cell_h @@ -793,7 +782,7 @@ class RasterGridManager(QObject): self, painter: QPainter, grid: RasterGridRequest, - values: List[float] | List[int], + values: list[float] | list[int], alpha: int, visible_rect: QRectF | None, ) -> bool: @@ -823,8 +812,8 @@ class RasterGridManager(QObject): ) min_fill_px = 3.0 - stride_x = max(1, int(math.ceil(min_fill_px / max(cell_w, 1e-9)))) - stride_y = max(1, int(math.ceil(min_fill_px / max(cell_h, 1e-9)))) + stride_x = max(1, math.ceil(min_fill_px / max(cell_w, 1e-9))) + stride_y = max(1, math.ceil(min_fill_px / max(cell_h, 1e-9))) painter.save() painter.setRenderHint(QPainter.RenderHint.Antialiasing, False) @@ -855,7 +844,7 @@ class RasterGridManager(QObject): return True def _block_value( - self, values: List[float] | List[int], grid_nx: int, x0: int, x1: int, y0: int, y1: int + self, values: list[float] | list[int], grid_nx: int, x0: int, x1: int, y0: int, y1: int ) -> float | None: best = None for y in range(y0, y1): @@ -967,5 +956,5 @@ class RasterGridManager(QObject): ) self.completed_grid_updated.emit() - def get_completed_grids(self) -> List[CompletedRasterGridElem]: + def get_completed_grids(self) -> list[CompletedRasterGridElem]: return self._completed_grids diff --git a/src/aare/gui/threads/axis_video_thread.py b/src/aare/gui/threads/axis_video_thread.py index 26f868ea..7a27ba8d 100644 --- a/src/aare/gui/threads/axis_video_thread.py +++ b/src/aare/gui/threads/axis_video_thread.py @@ -1,6 +1,6 @@ import cv2 -import requests import numpy as np +import requests from PySide6.QtCore import QThread, Signal from PySide6.QtGui import QImage @@ -76,9 +76,9 @@ class VideoThread(QThread): buffer = buffer[last_boundary:] except requests.exceptions.RequestException as e: - self.error_occurred.emit(f"Connection error: {str(e)}") + self.error_occurred.emit(f"Connection error: {e!s}") except Exception as e: - self.error_occurred.emit(f"Unexpected error: {str(e)}") + self.error_occurred.emit(f"Unexpected error: {e!s}") finally: self.running = False if self.session: @@ -119,7 +119,7 @@ class VideoThread(QThread): rgb_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) # Convert to QImage - height, width, channel = rgb_frame.shape + height, width, _channel = rgb_frame.shape bytes_per_line = 3 * width qt_image = QImage( rgb_frame.data, width, height, bytes_per_line, QImage.Format.Format_RGB888 diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index f08d8213..4683948f 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -352,7 +352,7 @@ class DAQWorker(QObject): self._last_status_request_ts = now request = QNetworkRequest(QUrl(f"{self._base_url}/status")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self.handle_status_response(reply)) @@ -739,7 +739,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/{url}")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) if str: request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -757,7 +757,7 @@ class DAQWorker(QObject): logger.info(f"PUT /{url}: {body}") return request = QNetworkRequest(QUrl(f"{self._base_url}/{url}")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) if str: request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.put(request, QByteArray(body.encode("utf-8"))) @@ -775,7 +775,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/{url}")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.deleteResource(request) reply.finished.connect(lambda: self.handle_req_response(reply)) @@ -874,7 +874,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/state/free_beamline")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = json.dumps({"confirmation_code": confirmation_code}) reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -889,7 +889,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/access/take_over_beamline")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = json.dumps({"confirmation_code": confirmation_code}) reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -904,7 +904,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/recovery/recover_beamline")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = json.dumps({"confirmation_code": confirmation_code}) reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -921,7 +921,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/recovery/unmount_sample")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = json.dumps({"confirmation_code": confirmation_code}) reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -964,7 +964,7 @@ class DAQWorker(QObject): self.staff_pgroups_loaded.emit([]) return request = QNetworkRequest(QUrl(f"{self._base_url}/access/all_pgroups")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.put(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_all_pgroups_response(reply)) @@ -1033,7 +1033,7 @@ class DAQWorker(QObject): self.run_number_incremented.emit() request = QNetworkRequest(QUrl(f"{self._base_url}/scan/rotation")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = r.model_dump_json() reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -1114,7 +1114,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/scan/raster?auto_center=false")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = r.model_dump_json() reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -1156,7 +1156,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/scan/raster?auto_center=true")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = r.model_dump_json() reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -1169,7 +1169,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/sample/spreadsheet")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self.handle_spreadsheet_response(reply)) @@ -1179,7 +1179,7 @@ class DAQWorker(QObject): logger.info("GET /sample/reference_tools") return request = QNetworkRequest(QUrl(f"{self._base_url}/sample/reference_tools")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self.handle_reference_tools_response(reply)) @@ -1314,14 +1314,7 @@ class DAQWorker(QObject): ): return True - if ( - "daq state error" in text - or "must be idle to start measurement" in text - or "must be idle" in text - ): - return True - - return False + return bool("daq state error" in text or "must be idle to start measurement" in text or "must be idle" in text) def handle_auto_scan_response(self, reply, sample_id: int): if reply.error() == QNetworkReply.NetworkError.NoError: @@ -1377,7 +1370,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/scan/auto")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = s.model_dump_json() reply = self._net_manager.post(request, QByteArray(body.encode("utf-8"))) @@ -1425,7 +1418,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/sample/resync")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_sample_resync_response(reply)) @@ -1437,7 +1430,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/resync/detector_metadata")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_detector_metadata_resync_response(reply)) @@ -1499,7 +1492,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/simulation_state")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_local_contact_simulation_state_response(reply)) @@ -1528,7 +1521,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/device_state")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_local_contact_device_state_response(reply)) @@ -1557,7 +1550,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/links")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_local_contact_links_response(reply)) @@ -1589,7 +1582,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/config")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_local_contact_config_response(reply)) @@ -1615,7 +1608,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/config")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") body = QByteArray(json.dumps(payload).encode("utf-8")) reply = self._net_manager.put(request, body) @@ -1657,7 +1650,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/bec/user_macros")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_bec_user_macros_response(reply)) @@ -1680,7 +1673,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/bec/devices")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_bec_devices_response(reply)) @@ -1768,7 +1761,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/alc/ml_bounding_box")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self.handle_ml_box_response(reply)) @@ -2006,7 +1999,7 @@ class DAQWorker(QObject): self._automation_progress_buffer = "" request = QNetworkRequest(QUrl(f"{self._base_url}/sse/automation_progress")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.readyRead.connect(lambda: self._read_automation_progress_stream(reply)) reply.finished.connect(self._restart_automation_progress_stream) @@ -2023,7 +2016,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/sse/face_detection")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.readyRead.connect(lambda: self._read_face_detection_stream(reply)) reply.finished.connect(self._restart_face_detection_stream) @@ -2045,7 +2038,7 @@ class DAQWorker(QObject): request = QNetworkRequest( QUrl(f"{self._base_url}/face_detection/run?steps={steps}&step_size={step_size}") ) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_face_detection_response(reply)) @@ -2071,7 +2064,7 @@ class DAQWorker(QObject): if emit_status: # fetch status and bkg in parallel (simple sequential here) status_req = QNetworkRequest(QUrl(f"{self._base_url}/fluorimeter/status")) - status_req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + status_req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) status_reply = self._net_manager.get(status_req) status_reply.finished.connect( lambda: self._handle_fluorimeter_status_and_emit(data, status_reply) @@ -2087,7 +2080,7 @@ class DAQWorker(QObject): s_payload = self.handle_response(status_reply) s = int(s_payload) if s_payload not in ("", "null") else -1 b_req = QNetworkRequest(QUrl(f"{self._base_url}/fluorimeter/background")) - b_req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + b_req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) b_reply = self._net_manager.get(b_req) b_reply.finished.connect(lambda: self._emit_fluorimeter_with_bkg(data, s, b_reply)) except Exception as e: @@ -2110,7 +2103,7 @@ class DAQWorker(QObject): if self._base_url is None: return req = QNetworkRequest(QUrl(f"{self._base_url}/fluorimeter/spectrum")) - req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) req.setRawHeader(b"Content-Type", b"application/json") body = f.model_dump_json() reply = self._net_manager.post(req, QByteArray(body.encode("utf-8"))) @@ -2130,7 +2123,7 @@ class DAQWorker(QObject): if self._base_url is None: return req = QNetworkRequest(QUrl(f"{self._base_url}/fluorimeter/data")) - req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + req.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(req) reply.finished.connect(lambda: self._handle_fluorimeter_data(reply, emit_status=True)) @@ -2139,7 +2132,7 @@ class DAQWorker(QObject): if self._base_url is None: return request = QNetworkRequest(QUrl(f"{self._base_url}/sse/fluorimeter")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.readyRead.connect(lambda: self._read_fluorimeter_stream(reply)) reply.finished.connect(lambda: reply.deleteLater()) @@ -2173,7 +2166,7 @@ class DAQWorker(QObject): if isinstance(v, dict): group = str(k) for kk, vv in v.items(): - out[f"{group}.{str(kk)}"] = str(vv) + out[f"{group}.{kk!s}"] = str(vv) else: out[str(k)] = str(v) return out @@ -2191,13 +2184,13 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/meta/error-codes")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_error_codes_response(reply)) def _retry_error_codes_legacy(self) -> None: request = QNetworkRequest(QUrl(f"{self._base_url}/meta/error-codes/flat")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_error_codes_response(reply)) @@ -2255,7 +2248,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/sse/baton")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.readyRead.connect(lambda: self._read_baton_stream(reply)) reply.finished.connect(self._restart_baton_stream) @@ -2312,7 +2305,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/baton/request")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_baton_request_response(reply)) @@ -2353,7 +2346,7 @@ class DAQWorker(QObject): request = QNetworkRequest( QUrl(f"{self._base_url}/baton/respond?accept={str(accept).lower()}") ) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_baton_response_result(reply)) @@ -2390,7 +2383,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/baton/check_timeout")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_baton_timeout_response(reply)) @@ -2453,7 +2446,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/admin/gui_sessions")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.get(request) reply.finished.connect(lambda: self._handle_gui_sessions_response(reply)) @@ -2470,7 +2463,7 @@ class DAQWorker(QObject): f"{self._base_url}/admin/gui_sessions/{session_id}/request_close?grace_seconds={grace_seconds}" ) ) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: self._handle_gui_session_mutation_response(reply)) @@ -2482,7 +2475,7 @@ class DAQWorker(QObject): return request = QNetworkRequest(QUrl(f"{self._base_url}/admin/gui_sessions/{session_id}")) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) reply = self._net_manager.deleteResource(request) reply.finished.connect(lambda: self._handle_gui_session_mutation_response(reply)) @@ -2508,7 +2501,7 @@ class DAQWorker(QObject): request = QNetworkRequest( QUrl(f"{self._base_url}/admin/gui_sessions/{session_id}/interaction") ) - request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode("utf-8")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Content-Type", b"application/json") reply = self._net_manager.post(request, QByteArray(b"")) reply.finished.connect(lambda: reply.deleteLater()) diff --git a/src/aare/gui/threads/sse_client.py b/src/aare/gui/threads/sse_client.py index 7264b955..63f32295 100644 --- a/src/aare/gui/threads/sse_client.py +++ b/src/aare/gui/threads/sse_client.py @@ -1,6 +1,6 @@ -from typing import Optional, Dict -from PySide6.QtCore import QObject, Signal, Slot, QUrl, QTimer, QByteArray -from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply, QSslError + +from PySide6.QtCore import QByteArray, QObject, QTimer, QUrl, Signal, Slot +from PySide6.QtNetwork import QNetworkAccessManager, QNetworkReply, QNetworkRequest, QSslError class SSEClient(QObject): @@ -15,13 +15,13 @@ class SSEClient(QObject): super().__init__(parent) self._network_manager = QNetworkAccessManager(self) self._network_manager.sslErrors.connect(self._handle_ssl_errors) - self._reply: Optional[QNetworkReply] = None + self._reply: QNetworkReply | None = None self._reconnect_timer = QTimer(self) self._reconnect_timer.setSingleShot(True) self._reconnect_timer.timeout.connect(self._attempt_reconnect) self._url = QUrl() - self._headers: Dict[str, str] = {} + self._headers: dict[str, str] = {} self._buffer = QByteArray() self._connected = False self._reconnect_delay = 1000 # Start with 1 second @@ -32,7 +32,7 @@ class SSEClient(QObject): self._current_data = "" self._current_id = "" - def connect_to_sse(self, url: str, headers: Optional[Dict[str, str]] = None): + def connect_to_sse(self, url: str, headers: dict[str, str] | None = None): """Connect to SSE endpoint""" self._url = QUrl(url) self._headers = headers or {} @@ -155,8 +155,7 @@ class SSEClient(QObject): value = line[colon_index + 1 :] # Remove leading space from value - if value.startswith(" "): - value = value[1:] + value = value.removeprefix(" ") if field == "data": if self._current_data: diff --git a/src/aare/gui/tutorials/controls_help_dialog.py b/src/aare/gui/tutorials/controls_help_dialog.py index 30ef8f59..0883e97a 100644 --- a/src/aare/gui/tutorials/controls_help_dialog.py +++ b/src/aare/gui/tutorials/controls_help_dialog.py @@ -1,5 +1,5 @@ from PySide6.QtCore import Qt -from PySide6.QtWidgets import QDialog, QVBoxLayout, QTextEdit, QDialogButtonBox, QTabWidget, QWidget +from PySide6.QtWidgets import QDialog, QDialogButtonBox, QTabWidget, QTextEdit, QVBoxLayout, QWidget class ControlsHelpDialog(QDialog): diff --git a/src/aare/gui/tutorials/tutorial_manager.py b/src/aare/gui/tutorials/tutorial_manager.py index 005a4dc5..c9fa6ae8 100644 --- a/src/aare/gui/tutorials/tutorial_manager.py +++ b/src/aare/gui/tutorials/tutorial_manager.py @@ -5,10 +5,10 @@ from dataclasses import dataclass from typing import Any from PySide6.QtCore import ( + Property, QEasingCurve, QObject, QPropertyAnimation, - Property, QRect, Qt, QTimer, @@ -460,8 +460,7 @@ class TutorialManager(QObject): self.tutorial_stopped.emit(scenario_id) return - if new_index < 0: - new_index = 0 + new_index = max(new_index, 0) self.runtime_state.active_step_index = new_index step = self.current_scenario.steps[new_index] @@ -688,9 +687,8 @@ class TutorialManager(QObject): step.skippable and self.current_scenario is not None and self.current_scenario.allow_skip - ): - if self.runtime_state is not None: - self._advance_to_index(self.runtime_state.active_step_index + 1) - return + ) and self.runtime_state is not None: + self._advance_to_index(self.runtime_state.active_step_index + 1) + return self.overlay.next_button.setEnabled(True) diff --git a/src/aare/gui/tutorials/tutorial_models.py b/src/aare/gui/tutorials/tutorial_models.py index d0ab205c..fe2bae5c 100644 --- a/src/aare/gui/tutorials/tutorial_models.py +++ b/src/aare/gui/tutorials/tutorial_models.py @@ -51,7 +51,7 @@ class CompletionKind(str, Enum): class CompletionRule: kind: CompletionKind value: Any = None - children: list["CompletionRule"] = field(default_factory=list) + children: list[CompletionRule] = field(default_factory=list) description: TutorialTextRef | None = None @@ -121,7 +121,7 @@ class TutorialScenario: title: TutorialTextRef description: TutorialTextRef mode: TutorialMode - steps: list["TutorialStepDefinition"] + steps: list[TutorialStepDefinition] version: str = "1.0" tags: set[str] = field(default_factory=set) @@ -146,6 +146,6 @@ class TutorialContext: completed_step_ids: list[str] = field(default_factory=list) state: dict[str, Any] = field(default_factory=dict) - event_log: list["TutorialEvent"] = field(default_factory=list) + event_log: list[TutorialEvent] = field(default_factory=list) metadata: dict[str, Any] = field(default_factory=dict) diff --git a/src/aare/gui/widgets/automation_progress.py b/src/aare/gui/widgets/automation_progress.py index 80e0663a..4c88baab 100644 --- a/src/aare/gui/widgets/automation_progress.py +++ b/src/aare/gui/widgets/automation_progress.py @@ -112,7 +112,7 @@ class CompactAutomationProgressStrip(QFrame): def _format_duration(seconds: float | None) -> str: if seconds is None or seconds <= 0: return "0m 00s" - total = int(round(seconds)) + total = round(seconds) minutes, secs = divmod(total, 60) if minutes < 60: return f"{minutes}m {secs:02d}s" diff --git a/src/aare/gui/widgets/baton_request_dialog.py b/src/aare/gui/widgets/baton_request_dialog.py index 16043c8b..4ec6ae4d 100644 --- a/src/aare/gui/widgets/baton_request_dialog.py +++ b/src/aare/gui/widgets/baton_request_dialog.py @@ -1,14 +1,14 @@ -from PySide6.QtCore import Qt, Signal, QTimer +from PySide6.QtCore import Qt, QTimer, Signal +from PySide6.QtGui import QFont from PySide6.QtWidgets import ( QDialog, - QVBoxLayout, + QFrame, QHBoxLayout, QLabel, - QPushButton, QProgressBar, - QFrame, + QPushButton, + QVBoxLayout, ) -from PySide6.QtGui import QFont class BatonRequestDialog(QDialog): @@ -317,8 +317,7 @@ class BatonPendingDialog(QDialog): def _tick(self): self._remaining -= 1 - if self._remaining < 0: - self._remaining = 0 + self._remaining = max(self._remaining, 0) self.progress.setValue(self._remaining) self.time_label.setText(f"{self._remaining} seconds remaining") diff --git a/src/aare/gui/widgets/camera_image.py b/src/aare/gui/widgets/camera_image.py index 7d94d217..6cae2c24 100644 --- a/src/aare/gui/widgets/camera_image.py +++ b/src/aare/gui/widgets/camera_image.py @@ -405,12 +405,11 @@ class SampleCameraImageLabel(QGraphicsView): self._state = SampleCameraImageState.RESIZE_RASTER_GRID else: self._state = SampleCameraImageState.DRAWING_RASTER_GRID - elif event.button() == Qt.MouseButton.LeftButton: - if ( - self._raster_mgr.is_part_of_active_grid(self.start_point) - and not ctrl_override_move - ): - self._state = SampleCameraImageState.MOVING_RASTER_GRID + elif event.button() == Qt.MouseButton.LeftButton and ( + self._raster_mgr.is_part_of_active_grid(self.start_point) + and not ctrl_override_move + ): + self._state = SampleCameraImageState.MOVING_RASTER_GRID def _update_grid(self): now = time.monotonic() @@ -613,8 +612,7 @@ class SampleCameraImageLabel(QGraphicsView): ratio_h = self.viewport().size().height() / self.pixmap_item.boundingRect().height() ratio = min(ratio_w, ratio_h) - if ratio < 0.1: - ratio = 0.1 + ratio = max(ratio, 0.1) if ratio >= 1.0: # Don't enable scaling when gain in ratio is < 5% (to avoid back-and-forth) diff --git a/src/aare/gui/widgets/clickable_label.py b/src/aare/gui/widgets/clickable_label.py index 97d34aa8..eb3b2aac 100644 --- a/src/aare/gui/widgets/clickable_label.py +++ b/src/aare/gui/widgets/clickable_label.py @@ -1,4 +1,4 @@ -from PySide6.QtCore import Signal, Qt +from PySide6.QtCore import Qt, Signal from PySide6.QtWidgets import QLabel diff --git a/src/aare/gui/widgets/local_contact_status_widget.py b/src/aare/gui/widgets/local_contact_status_widget.py index e9bae8db..c8253b2d 100644 --- a/src/aare/gui/widgets/local_contact_status_widget.py +++ b/src/aare/gui/widgets/local_contact_status_widget.py @@ -351,7 +351,7 @@ class LocalContactStatusWidget(QFrame): def _refresh(self) -> None: if self._last_status is None: - for _key, (_title, value) in self._row_widgets.items(): + for (_title, value) in self._row_widgets.values(): value.setText(self._badge("WAITING", tone="neutral")) return diff --git a/src/aare/gui/widgets/number_line_edit.py b/src/aare/gui/widgets/number_line_edit.py index fd79f088..168d68b4 100644 --- a/src/aare/gui/widgets/number_line_edit.py +++ b/src/aare/gui/widgets/number_line_edit.py @@ -1,6 +1,6 @@ -from PySide6.QtCore import Signal, Slot, Qt +from PySide6.QtCore import Qt, Signal, Slot from PySide6.QtGui import QDoubleValidator -from PySide6.QtWidgets import QLineEdit, QWidget, QCheckBox, QHBoxLayout +from PySide6.QtWidgets import QCheckBox, QHBoxLayout, QLineEdit, QWidget class NumberLineEdit(QLineEdit): @@ -23,7 +23,7 @@ class NumberLineEdit(QLineEdit): self.setValidator(self.validator) self.setAlignment(Qt.AlignmentFlag.AlignRight) self.setToolTip( - "Minimum: {:s}\nMaximum: {:s}".format(self.to_string(min_val), self.to_string(max_val)) + f"Minimum: {self.to_string(min_val):s}\nMaximum: {self.to_string(max_val):s}" ) # Connect the textChanged signal to a custom slot to check validity @@ -75,7 +75,7 @@ class NumberLineEdit(QLineEdit): def update_limits(self, min_val: float, max_val: float): self.validator.setRange(min_val, max_val, self.decimal_count) self.setToolTip( - "Minimum: {:s}\nMaximum: {:s}".format(self.to_string(min_val), self.to_string(max_val)) + f"Minimum: {self.to_string(min_val):s}\nMaximum: {self.to_string(max_val):s}" ) def validate(self, text) -> bool: diff --git a/src/aare/gui/widgets/pgroup_dialog.py b/src/aare/gui/widgets/pgroup_dialog.py index 11e9b0ea..6d467abf 100644 --- a/src/aare/gui/widgets/pgroup_dialog.py +++ b/src/aare/gui/widgets/pgroup_dialog.py @@ -1,12 +1,12 @@ from PySide6.QtCore import Qt from PySide6.QtWidgets import ( - QDialog, - QVBoxLayout, - QPushButton, - QLabel, QComboBox, QCompleter, + QDialog, + QLabel, QMessageBox, + QPushButton, + QVBoxLayout, ) diff --git a/src/aare/gui/widgets/splash_screen.py b/src/aare/gui/widgets/splash_screen.py index 2c7da0d4..91761223 100644 --- a/src/aare/gui/widgets/splash_screen.py +++ b/src/aare/gui/widgets/splash_screen.py @@ -1,5 +1,5 @@ from PySide6.QtCore import Qt -from PySide6.QtWidgets import QSplashScreen, QProgressBar, QApplication +from PySide6.QtWidgets import QApplication, QProgressBar, QSplashScreen class LoadingSplashScreen(QSplashScreen): diff --git a/src/aare/gui/widgets/status_bar.py b/src/aare/gui/widgets/status_bar.py index 34506549..b182b531 100644 --- a/src/aare/gui/widgets/status_bar.py +++ b/src/aare/gui/widgets/status_bar.py @@ -538,7 +538,6 @@ class StatusBar(QStatusBar): self.staff_pgroups_loaded.disconnect(_on_loaded) except Exception as e: logger.debug(f"Error disconnecting: {e}") - pass self.staff_pgroups_loaded.connect(_on_loaded) self._list_staff_pgroups() @@ -586,7 +585,6 @@ class StatusBar(QStatusBar): def _list_staff_pgroups(self): self.get_all_pgroups.emit() - return def _generate_pgroup_dialogue(self, curr: str | None = None, pgroups: list | None = None): logger.info(pgroups) diff --git a/src/aare/gui/widgets/video_image.py b/src/aare/gui/widgets/video_image.py index bd03c132..78d9134d 100644 --- a/src/aare/gui/widgets/video_image.py +++ b/src/aare/gui/widgets/video_image.py @@ -1,6 +1,6 @@ -from PySide6.QtCore import Qt, Slot, QRectF -from PySide6.QtGui import QPainter, QPixmap, QImage, QFont, QColor, QPen, QFontMetrics -from PySide6.QtWidgets import QGraphicsView, QGraphicsScene, QGraphicsPixmapItem +from PySide6.QtCore import QRectF, Qt, Slot +from PySide6.QtGui import QColor, QFont, QFontMetrics, QImage, QPainter, QPen, QPixmap +from PySide6.QtWidgets import QGraphicsPixmapItem, QGraphicsScene, QGraphicsView from aare.gui.widgets.busy_overlay import BusyOverlayStyle diff --git a/tests/conftest.py b/tests/conftest.py index c73d1cbc..b43e2a01 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -50,7 +50,7 @@ def sample_info(): # ------------------------- @pytest.fixture(scope="session") def server_module(): - import aare.daq.server as server + from aare.daq import server return server diff --git a/tests/integration/daq/test_daq_server.py b/tests/integration/daq/test_daq_server.py index c9f26964..0fc50dca 100644 --- a/tests/integration/daq/test_daq_server.py +++ b/tests/integration/daq/test_daq_server.py @@ -1,5 +1,6 @@ -import pytest import os + +import pytest from fastapi.testclient import TestClient # We need to set the environment variable before importing the app diff --git a/tests/unit/daq/operations/face_detection/test_face_detection_logic.py b/tests/unit/daq/operations/face_detection/test_face_detection_logic.py index 4f430cfb..ca58fd43 100644 --- a/tests/unit/daq/operations/face_detection/test_face_detection_logic.py +++ b/tests/unit/daq/operations/face_detection/test_face_detection_logic.py @@ -1,17 +1,18 @@ import numpy as np import pytest + from aare.daq.operations.face_detection.utils import ( - box_height_from_tuple, box_area_from_tuple, - prepare_samples, - cos_model, - mad_filter, - fit_metrics, - fit_cosine, - get_samples_out, + box_height_from_tuple, choose_best_fit, - get_flat_face, chose_best_angle, + cos_model, + fit_cosine, + fit_metrics, + get_flat_face, + get_samples_out, + mad_filter, + prepare_samples, ) @@ -101,7 +102,7 @@ def test_choose_best_fit(): "Height": {"angle": 45, "params": {"rmse": 0.1, "mae": 0.1, "r2": 0.95}}, "Area": {"angle": 50, "params": {"rmse": 0.05, "mae": 0.05, "r2": 0.98}}, } - angle, fit, name = choose_best_fit(fits) + angle, _fit, name = choose_best_fit(fits) assert name == "Area" assert angle == 50.0 diff --git a/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py b/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py index 2dc7e809..8059a644 100644 --- a/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py +++ b/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py @@ -88,7 +88,7 @@ def test_service_succeeds_when_correction_pass_has_valid_target(monkeypatch, con return AngleAnalysis( angle_deg=kwargs["angle"], classes=[MLBoxType.CRYSTAL.value] if kwargs["angle"] == 0 else [], - has_valid_target=True if kwargs["angle"] == 0 else False, + has_valid_target=kwargs["angle"] == 0, ignore_only=False, ) return AngleAnalysis( diff --git a/tests/unit/daq/test_autofocus.py b/tests/unit/daq/test_autofocus.py index 77e398c0..ce108fbd 100644 --- a/tests/unit/daq/test_autofocus.py +++ b/tests/unit/daq/test_autofocus.py @@ -1,5 +1,6 @@ -import numpy as np import cv2 +import numpy as np + from aare.daq.autofocus import calculate_focus_measure diff --git a/tests/unit/daq/test_beamcenterfit.py b/tests/unit/daq/test_beamcenterfit.py index c290f09e..8155dcb9 100644 --- a/tests/unit/daq/test_beamcenterfit.py +++ b/tests/unit/daq/test_beamcenterfit.py @@ -1,6 +1,7 @@ -import pytest import numpy as np -from aare.daq.beamcenterfit import beamcenter_fit, Gaussian2Dfit +import pytest + +from aare.daq.beamcenterfit import Gaussian2Dfit, beamcenter_fit def create_synthetic_beam_image( @@ -50,7 +51,6 @@ def test_beamcenter_fit_no_converge(): # but we want to test the failure path. # To truly force non-convergence we might need a more extreme case, # but return None is better than exit() anyway. - pass def test_beamcenter_fit_no_contours(): diff --git a/tests/unit/daq/test_face_detection.py b/tests/unit/daq/test_face_detection.py index e1075e09..887beb29 100644 --- a/tests/unit/daq/test_face_detection.py +++ b/tests/unit/daq/test_face_detection.py @@ -77,21 +77,9 @@ def test_public_face_detection_uses_execute_face_detection(monkeypatch): daq = object.__new__(AareDAQ) cfg = types.SimpleNamespace(try_set_busy=lambda timeout=360: None, state_busy=False) - setattr(daq, "_cfg", cfg) + daq._cfg = cfg - setattr( - daq, - "_execute_face_detection", - lambda **kwargs: FaceDetectionResult( - success=True, - payload={ - "running": False, - "samples": [{"angle": 45}], - "height_fit": {}, - "area_fit": {}, - }, - ), - ) + daq._execute_face_detection = lambda **kwargs: FaceDetectionResult(success=True, payload={"running": False, "samples": [{"angle": 45}], "height_fit": {}, "area_fit": {}}) result = daq.face_detection(steps=7, step_size=30) diff --git a/tests/unit/daq/test_mlbox.py b/tests/unit/daq/test_mlbox.py index 98ef9a9e..650e8c09 100644 --- a/tests/unit/daq/test_mlbox.py +++ b/tests/unit/daq/test_mlbox.py @@ -70,7 +70,7 @@ def test_filter_predictions(mlbox): mlbox._filter_predictions(out, confidence_min=0.5) assert len(out.boxes) == 1 - assert list(out.boxes.values())[0].conf == 0.8 + assert next(iter(out.boxes.values())).conf == 0.8 @patch("cv2.imdecode") diff --git a/tests/unit/daq/test_mount.py b/tests/unit/daq/test_mount.py index 808272fd..2abdc2f1 100644 --- a/tests/unit/daq/test_mount.py +++ b/tests/unit/daq/test_mount.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from types import SimpleNamespace from unittest.mock import MagicMock @@ -61,13 +61,13 @@ def test_was_previous_sample_unmounted_since_prefers_tell_phase_confirmation(): previous_sample = _make_sample(1, "old") daq = _make_daq(previous_sample) - started_at = datetime.now(timezone.utc) - timedelta(seconds=5) + started_at = datetime.now(UTC) - timedelta(seconds=5) daq._safe_tell_state = MagicMock( return_value=TellStateModel( activity=TellActivityEnum.MOUNTING, operation="mount", phase=TellPhaseEnum.PICKING_NEW_SAMPLE, - last_update_ts=datetime.now(timezone.utc).isoformat(), + last_update_ts=datetime.now(UTC).isoformat(), last_event_class="Motion Sync", last_event_value="Sample get on Puck", ) diff --git a/tests/unit/daq/test_server.py b/tests/unit/daq/test_server.py index 7a535775..c1343cad 100644 --- a/tests/unit/daq/test_server.py +++ b/tests/unit/daq/test_server.py @@ -2,7 +2,6 @@ import os from types import SimpleNamespace from unittest.mock import patch - os.environ["JWT_AAREDAQ_KEY"] = "test_key_for_unit_testing" diff --git a/tests/unit/daq/test_spreadsheetupdater.py b/tests/unit/daq/test_spreadsheetupdater.py index 2255bb87..56ce46d9 100644 --- a/tests/unit/daq/test_spreadsheetupdater.py +++ b/tests/unit/daq/test_spreadsheetupdater.py @@ -25,9 +25,8 @@ def test_get_ws_headers_success(): def test_get_ws_headers_fail(): - with patch("os.getenv", return_value=None): - with pytest.raises(ValueError): - get_ws_headers() + with patch("os.getenv", return_value=None), pytest.raises(ValueError): + get_ws_headers() def test_set_spreadsheet_in_redis(mock_config): diff --git a/tests/unit/devices/test_enum_pv.py b/tests/unit/devices/test_enum_pv.py index 6ae6cd4d..89658a5e 100644 --- a/tests/unit/devices/test_enum_pv.py +++ b/tests/unit/devices/test_enum_pv.py @@ -1,6 +1,8 @@ -import pytest -from unittest.mock import patch from enum import Enum +from unittest.mock import patch + +import pytest + from aare.devices.enum_pv import EnumPV @@ -29,14 +31,14 @@ def mock_pvs(): def test_enum_pv_init_fail(mock_pvs): - set_pv, get_pv = mock_pvs + set_pv, _get_pv = mock_pvs set_pv.enum_strs = None with pytest.raises(RuntimeError): EnumPV("test", "SET", "GET") def test_enum_pv_init_success(mock_pvs): - set_pv, get_pv = mock_pvs + set_pv, _get_pv = mock_pvs set_pv.enum_strs = ("State1", "State2") epv = EnumPV("test", "SET", "GET") assert epv.name == "test" @@ -51,7 +53,7 @@ def test_enum_pv_position(mock_pvs): def test_enum_pv_resolve_enum(mock_pvs): - set_pv, get_pv = mock_pvs + set_pv, _get_pv = mock_pvs set_pv.enum_strs = ("State1", "State2") epv = EnumPV("test", "SET", "GET") @@ -63,7 +65,7 @@ def test_enum_pv_resolve_enum(mock_pvs): def test_enum_pv_resolve_int(mock_pvs): - set_pv, get_pv = mock_pvs + set_pv, _get_pv = mock_pvs set_pv.enum_strs = ("State1", "State2") epv = EnumPV("test", "SET", "GET") @@ -75,7 +77,7 @@ def test_enum_pv_resolve_int(mock_pvs): def test_enum_pv_resolve_str(mock_pvs): - set_pv, get_pv = mock_pvs + set_pv, _get_pv = mock_pvs set_pv.enum_strs = (" State1 ", "State2") epv = EnumPV("test", "SET", "GET") @@ -87,7 +89,7 @@ def test_enum_pv_resolve_str(mock_pvs): def test_enum_pv_resolve_invalid_type(mock_pvs): - set_pv, get_pv = mock_pvs + set_pv, _get_pv = mock_pvs set_pv.enum_strs = ("State1", "State2") epv = EnumPV("test", "SET", "GET") with pytest.raises(TypeError): diff --git a/tests/unit/devices/test_mx_lib.py b/tests/unit/devices/test_mx_lib.py index e7158f1c..4baeb224 100644 --- a/tests/unit/devices/test_mx_lib.py +++ b/tests/unit/devices/test_mx_lib.py @@ -1,7 +1,9 @@ -import pytest from unittest.mock import MagicMock, patch -from aare.devices.mx_lib import wait_for_movement_to_finish, pv_wait, is_epics_type, clean_filename -from epics import Motor, PV + +import pytest +from epics import PV, Motor + +from aare.devices.mx_lib import clean_filename, is_epics_type, pv_wait, wait_for_movement_to_finish def test_clean_filename(): @@ -47,9 +49,8 @@ def test_wait_for_movement_to_finish_timeout(mock_poll): mock_motor.units = "mm" mock_motor._prefix = "MOT1:" - with patch("time.time", side_effect=[0, 0, 100, 101]): - with pytest.raises(TimeoutError): - wait_for_movement_to_finish(mock_motor) + with patch("time.time", side_effect=[0, 0, 100, 101]), pytest.raises(TimeoutError): + wait_for_movement_to_finish(mock_motor) @patch("aare.devices.mx_lib.wait_motor_position") diff --git a/tests/unit/devices/test_my_motor.py b/tests/unit/devices/test_my_motor.py index eefd4cad..c6cb1258 100644 --- a/tests/unit/devices/test_my_motor.py +++ b/tests/unit/devices/test_my_motor.py @@ -34,7 +34,7 @@ def test_my_motor_init(mock_motor_base): def test_my_motor_speed(mock_motor_base): - _, mock_get, mock_put, _ = mock_motor_base + _, mock_get, _mock_put, _ = mock_motor_base m = MyMotor("MTR") with patch.object(MyMotor, "name", create=True, new_callable=PropertyMock) as mock_name: mock_name.return_value = "MTR" @@ -94,7 +94,7 @@ def test_my_motor_units(mock_motor_base): def test_my_motor_limits(mock_motor_base): - _, mock_get, mock_put, _ = mock_motor_base + _, mock_get, _mock_put, _ = mock_motor_base m = MyMotor("MTR") with patch.object(MyMotor, "name", create=True, new_callable=PropertyMock) as mock_name: mock_name.return_value = "MTR" diff --git a/tests/unit/devices/test_workflow_tools.py b/tests/unit/devices/test_workflow_tools.py index d25419a0..c8136951 100644 --- a/tests/unit/devices/test_workflow_tools.py +++ b/tests/unit/devices/test_workflow_tools.py @@ -1,4 +1,5 @@ import pytest + from aare.devices.workflow_tools import wait_position diff --git a/tests/unit/gui/test_auth_mock.py b/tests/unit/gui/test_auth_mock.py index 8685fa43..836ef743 100644 --- a/tests/unit/gui/test_auth_mock.py +++ b/tests/unit/gui/test_auth_mock.py @@ -15,7 +15,7 @@ def test_auth_success(mocker): assert token == "fake_token_abc.123.xyz" mock_run.assert_called_once() - args, kwargs = mock_run.call_args + args, _kwargs = mock_run.call_args assert "curl" in args[0] assert "--cacert" in args[0] assert "/tmp/test-cert.pem" in args[0] diff --git a/tests/unit/gui/test_axis_video_thread.py b/tests/unit/gui/test_axis_video_thread.py index 828fb594..44ec9270 100644 --- a/tests/unit/gui/test_axis_video_thread.py +++ b/tests/unit/gui/test_axis_video_thread.py @@ -1,8 +1,10 @@ -import pytest -import numpy as np -import cv2 from unittest.mock import MagicMock, patch + +import cv2 +import numpy as np +import pytest from PySide6.QtGui import QImage + from aare.gui.threads.axis_video_thread import VideoThread diff --git a/tests/unit/gui/test_login.py b/tests/unit/gui/test_login.py index 0b761a1f..12f2d50c 100644 --- a/tests/unit/gui/test_login.py +++ b/tests/unit/gui/test_login.py @@ -1,8 +1,10 @@ +from unittest.mock import MagicMock, patch + +import jwt import pytest from PySide6.QtCore import Qt + from aare.gui.widgets.login import LoginDialog -from unittest.mock import MagicMock, patch -import jwt @pytest.fixture diff --git a/tests/unit/gui/test_message_box.py b/tests/unit/gui/test_message_box.py index 8f2464fd..5ed7923f 100644 --- a/tests/unit/gui/test_message_box.py +++ b/tests/unit/gui/test_message_box.py @@ -1,13 +1,15 @@ from unittest.mock import MagicMock, patch -from PySide6.QtWidgets import QMessageBox -from aare.gui.widgets.message_box import ( - reply_box, - timer_box, - ring_current_low_check, - experiment_hutch_shutter_check, - ring_current_auto_check, -) + import pytest +from PySide6.QtWidgets import QMessageBox + +from aare.gui.widgets.message_box import ( + experiment_hutch_shutter_check, + reply_box, + ring_current_auto_check, + ring_current_low_check, + timer_box, +) def test_reply_box(qtbot): diff --git a/tests/unit/gui/test_sse_client.py b/tests/unit/gui/test_sse_client.py index b1e20f1d..866c208d 100644 --- a/tests/unit/gui/test_sse_client.py +++ b/tests/unit/gui/test_sse_client.py @@ -1,8 +1,10 @@ +from unittest.mock import MagicMock + import pytest from PySide6.QtCore import QByteArray from PySide6.QtNetwork import QNetworkReply + from aare.gui.threads.sse_client import SSEClient -from unittest.mock import MagicMock @pytest.fixture diff --git a/tests/unit/gui/test_tutorials.py b/tests/unit/gui/test_tutorials.py index 2955477e..3f049cfb 100644 --- a/tests/unit/gui/test_tutorials.py +++ b/tests/unit/gui/test_tutorials.py @@ -1,19 +1,20 @@ import pytest from PySide6.QtWidgets import QWidget + from aare.gui.tutorials.tutorial_manager import TutorialManager from aare.gui.tutorials.tutorial_models import ( + StepKind, + TargetKind, + TutorialMode, TutorialScenario, TutorialStepDefinition, - TutorialMode, - StepKind, - TutorialTextRef, TutorialTarget, - TargetKind, + TutorialTextRef, ) from aare.gui.tutorials.tutorial_runtime import ( DictionaryTextResolver, - NoOpActionExecutor, DictTargetResolver, + NoOpActionExecutor, TutorialEventBus, ) diff --git a/tests/unit/utils/test_beamline_dispatch.py b/tests/unit/utils/test_beamline_dispatch.py index ff29f700..8feadae6 100644 --- a/tests/unit/utils/test_beamline_dispatch.py +++ b/tests/unit/utils/test_beamline_dispatch.py @@ -9,9 +9,8 @@ from aare.beamline_dispatch.x10sa import X10saDispatch def test_dispatch_cannot_be_instantiated_without_env(): - with patch.dict("os.environ", {}, clear=True): - with pytest.raises(ValueError) as e: - _ = get_beamline_dispatch() + with patch.dict("os.environ", {}, clear=True), pytest.raises(ValueError) as e: + _ = get_beamline_dispatch() assert e.match("Please set the BEAMLINE environment variable") @@ -23,9 +22,8 @@ def test_simulated_dispatch_can_be_instantiated(): @pytest.mark.skipif(importlib.util.find_spec("pxi_bec") is None, reason="run only for pxi flavour") def test_pxi_dispatch_can_be_instantiated(): - with patch.dict("os.environ", {"BEAMLINE": "X06SA"}): - with pytest.raises(TypeError) as e: - _ = get_beamline_dispatch() + with patch.dict("os.environ", {"BEAMLINE": "X06SA"}), pytest.raises(TypeError) as e: + _ = get_beamline_dispatch() assert e.match("Can't instantiate abstract class") @@ -42,7 +40,6 @@ def test_pxii_dispatch_can_be_instantiated(): importlib.util.find_spec("pxiii_bec") is None, reason="run only for pxiii flavour" ) def test_pxiii_dispatch_can_be_instantiated(): - with patch.dict("os.environ", {"BEAMLINE": "X06DA"}): - with pytest.raises(TypeError) as e: - _ = get_beamline_dispatch() + with patch.dict("os.environ", {"BEAMLINE": "X06DA"}), pytest.raises(TypeError) as e: + _ = get_beamline_dispatch() assert e.match("Can't instantiate abstract class")