style: ruff check --fix --unsafe-fixes

This commit is contained in:
2026-07-31 15:57:17 +02:00
parent 70c81abdd8
commit f63107a1e0
87 changed files with 413 additions and 481 deletions
+10 -13
View File
@@ -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
+8 -8
View File
@@ -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
+6 -5
View File
@@ -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:
+14 -15
View File
@@ -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
+1 -2
View File
@@ -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
+1 -1
View File
@@ -1,5 +1,5 @@
import numpy as np
import cv2
import numpy as np
from scipy.optimize import curve_fit
+5 -8
View File
@@ -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
+23 -33
View File
@@ -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,
+15 -17
View File
@@ -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.
"""
@@ -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
@@ -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")
@@ -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",
]
+1 -1
View File
@@ -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"]
+1 -1
View File
@@ -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(
+2 -2
View File
@@ -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)
+4 -5
View File
@@ -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"],
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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
+12 -13
View File
@@ -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,
):
+1 -2
View File
@@ -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)
+7 -8
View File
@@ -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")
+1 -1
View File
@@ -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):
+1 -1
View File
@@ -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()
+18 -19
View File
@@ -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:
+10 -9
View File
@@ -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")
+4 -2
View File
@@ -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:
+6 -10
View File
@@ -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}")
+2 -1
View File
@@ -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
+10 -11
View File
@@ -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()&"))
+4 -10
View File
@@ -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)
+4 -5
View File
@@ -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.
+1 -1
View File
@@ -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:
+2 -2
View File
@@ -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)
+2 -4
View File
@@ -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(
+1
View File
@@ -1,4 +1,5 @@
import json
from PySide6.QtCore import QSettings
+1 -3
View File
@@ -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):
+1 -1
View File
@@ -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())
+2 -2
View File
@@ -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)
+1 -1
View File
@@ -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
+1 -2
View File
@@ -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
+1 -1
View File
@@ -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
@@ -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
+3 -4
View File
@@ -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
@@ -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
+1 -1
View File
@@ -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):
+1 -1
View File
@@ -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,
+3 -4
View File
@@ -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,
@@ -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:
+21 -32
View File
@@ -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
+4 -4
View File
@@ -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
+45 -52
View File
@@ -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())
+7 -8
View File
@@ -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:
@@ -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):
+5 -7
View File
@@ -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)
+3 -3
View File
@@ -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)
+1 -1
View File
@@ -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"
+6 -7
View File
@@ -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")
+6 -8
View File
@@ -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)
+1 -1
View File
@@ -1,4 +1,4 @@
from PySide6.QtCore import Signal, Qt
from PySide6.QtCore import Qt, Signal
from PySide6.QtWidgets import QLabel
@@ -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
+4 -4
View File
@@ -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:
+4 -4
View File
@@ -1,12 +1,12 @@
from PySide6.QtCore import Qt
from PySide6.QtWidgets import (
QDialog,
QVBoxLayout,
QPushButton,
QLabel,
QComboBox,
QCompleter,
QDialog,
QLabel,
QMessageBox,
QPushButton,
QVBoxLayout,
)
+1 -1
View File
@@ -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):
-2
View File
@@ -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)
+3 -3
View File
@@ -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
+1 -1
View File
@@ -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
+2 -1
View File
@@ -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
@@ -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
@@ -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(
+2 -1
View File
@@ -1,5 +1,6 @@
import numpy as np
import cv2
import numpy as np
from aare.daq.autofocus import calculate_focus_measure
+3 -3
View File
@@ -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():
+2 -14
View File
@@ -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)
+1 -1
View File
@@ -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")
+3 -3
View File
@@ -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",
)
-1
View File
@@ -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"
+2 -3
View File
@@ -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):
+10 -8
View File
@@ -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):
+7 -6
View File
@@ -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")
+2 -2
View File
@@ -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"
@@ -1,4 +1,5 @@
import pytest
from aare.devices.workflow_tools import wait_position
+1 -1
View File
@@ -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]
+5 -3
View File
@@ -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
+4 -2
View File
@@ -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
+10 -8
View File
@@ -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):
+3 -1
View File
@@ -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
+6 -5
View File
@@ -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,
)
+6 -9
View File
@@ -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")