operations refactor: standardised operations, added rotation operation, updated tests.

This commit is contained in:
2026-06-17 16:25:02 +02:00
parent 9be8c1a3c8
commit 087ce0482f
22 changed files with 945 additions and 439 deletions
+215 -93
View File
@@ -12,7 +12,6 @@ import numpy as np
from aareDB import SampleEventType
from jfjoch_client import ScanResult, ScanResultImagesInner
from aare.common.tell_models import TellStateModel, TellPhaseEnum
from aare.daq import workflows
from aare.daq.aaredb import AareWrapper
@@ -20,6 +19,8 @@ from aare.daq.config import BeamlineConfig, ABR_POS_MOUNT
from aare.daq.config import BeamlineStateEnum
from aare.daq.devices import BeamlineDevices
from aare.daq.mlbox import MlBox
from aare.common.tell_models import TellStateModel, TellPhaseEnum
from aare.common.beamline import MXBeamline, cfg_get
from aare.common.coordinate import Coordinate, SmargonCoordinate, AerotechCoordinate
from aare.common.diffraction_geometry import DiffractionGeometry
@@ -47,30 +48,6 @@ from aare.common.automation_models import (
from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid
from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan
from aare.common.sample_geometry import SampleGeometryModel
from aare.daq.operations.face_detection import FaceDetectionContext, FaceDetectionService, FaceDetectionResult
from aare.daq.operations.loop_centering import LoopCenteringService, LoopCenteringContext
from aare.daq.operations.loop_centering.models import LoopCenteringSettings
from aare.daq.operations.mounting.service import MountingService
from aare.daq.operations.mounting.models import MountingResult, MountingContext
from aare.daq.operations.raster.models import RasterContext
from aare.daq.operations.raster.service import RasterService
from aare.daq.operations.common.ml_bounding_box import get_ml_bounding_box
from aare.daq.operations.screenshot.service import ScreenshotService
from aare.daq.operations.common.simulate_scan_result import (
build_fake_rotation_result,
)
from aare.daq.operations.common.runtime import (
DAQRuntimeState,
FaceDetectionProgressEmitter,
OperationServices,
PredictionProvider,
StateController,
TraceWriter,
)
from aare.devices.area_detector import AutoEnum
from aare.devices.jfjoch import JFJochWrapper
from aare.common.exception_handler import (
TransformationInvalidException,
StateTransitionFailed,
@@ -93,6 +70,55 @@ from aare.common.exception_handler import (
AutoRasterSampleSkipped
)
from aare.daq.operations.face_detection import FaceDetectionContext, FaceDetectionService, FaceDetectionResult
from aare.daq.operations.face_detection.models import (
FaceDetectionDependencies,
FaceDetectionSettings,
)
from aare.daq.operations.loop_centering import LoopCenteringService, LoopCenteringContext
from aare.daq.operations.loop_centering.models import (
LoopCenteringSettings,
LoopCenteringDependencies,
)
from aare.daq.operations.mounting.service import MountingService
from aare.daq.operations.mounting.models import (
MountingResult,
MountingContext,
MountingDependencies,
MountingSettings,
)
from aare.daq.operations.raster.models import (
RasterContext,
RasterSettings,
RasterDependencies,
)
from aare.daq.operations.raster.service import RasterService
from aare.daq.operations.common.ml_bounding_box import get_ml_bounding_box
from aare.daq.operations.rotation.models import (
RotationSettings,
RotationDependencies,
RotationContext,
)
from aare.daq.operations.rotation.service import RotationService
from aare.daq.operations.screenshot.service import ScreenshotService
from aare.daq.operations.common.simulate_scan_result import (
build_fake_rotation_result,
)
from aare.daq.operations.common.runtime import DAQRuntimeState
from aare.daq.operations.common.services import (
FaceDetectionProgressEmitter,
OperationServices,
PredictionProvider,
StateController,
TraceWriter,
SampleEventPublisher,
ScanIngestionService,
DataCollectionPreparer,
)
from aare.devices.area_detector import AutoEnum
from aare.devices.jfjoch import JFJochWrapper
logger = setup_logger("aareDAQ")
@@ -171,6 +197,45 @@ class _DAQTraceAppender:
self._daq._append_smargon_trace(sample_id=sample_id, event=event)
class _DAQSampleEventSender:
def __init__(self, daq: "AareDAQ"):
self._daq = daq
def send_sample_event(self, sample_id: int, event_type, comment: str | None = None) -> None:
self._daq._AareDAQ__aare.send_sample_event(sample_id, event_type, comment)
class _DAQScanIngestor:
def __init__(self, daq: "AareDAQ"):
self._daq = daq
def ingest_scan(self, *, sample, result, geom, beam_mark_pxl) -> None:
self._daq._AareDAQ__aare.ingest_scan(
sample=sample,
result=result,
geom=geom,
beam_mark_pxl=beam_mark_pxl,
)
def ingest_gridscan(self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl) -> None:
self._daq._AareDAQ__aare.ingest_gridscan(
sample=sample,
raster_result=raster_result,
raster_request=raster_request,
geom=geom,
com=com,
beam_mark_pxl=beam_mark_pxl,
)
class _DAQDatacollectionSetupRunner:
def __init__(self, daq: "AareDAQ"):
self._daq = daq
def prepare(self, request, screening: bool = False) -> None:
self._daq._AareDAQ__setup_datacollection(request=request, screening=screening)
class _LoopCenteringPredictionGetter:
def __init__(self, daq: "AareDAQ", settings: LoopCenteringSettings):
self._daq = daq
@@ -835,6 +900,21 @@ class AareDAQ:
def _build_operation_services(self) -> OperationServices:
return OperationServices(
screenshots=self._screenshot_service,
state=StateController(
setter=_DAQStateSetter(self),
),
traces=TraceWriter(
appender=_DAQTraceAppender(self),
),
events=SampleEventPublisher(
sender=_DAQSampleEventSender(self),
),
ingestion=ScanIngestionService(
ingestor=_DAQScanIngestor(self),
),
datacollection=DataCollectionPreparer(
runner=_DAQDatacollectionSetupRunner(self),
),
)
def _create_loop_centering_settings(self) -> LoopCenteringSettings:
@@ -842,85 +922,95 @@ class AareDAQ:
def _create_loop_centering_service(self) -> LoopCenteringService:
settings = self._create_loop_centering_settings()
services = self._build_operation_services()
services.predictions = PredictionProvider(
getter=_LoopCenteringPredictionGetter(self, settings),
)
return LoopCenteringService(
context=LoopCenteringContext(
cfg=self.__cfg,
devs=self.__devs,
mlbox=self.__mlbox,
settings=settings,
deps=LoopCenteringDependencies(
cfg=self.__cfg,
devs=self.__devs,
mlbox=self.__mlbox,
),
runtime=self._build_runtime_state(),
services=self._build_operation_services(),
trace_writer=TraceWriter(
appender=_DAQTraceAppender(self),
),
prediction_provider=PredictionProvider(
getter=_LoopCenteringPredictionGetter(self, settings),
),
services=services,
settings=settings,
),
logger=logger,
)
def _create_face_detection_service(self) -> FaceDetectionService:
services = self._build_operation_services()
services.face_detection_progress = FaceDetectionProgressEmitter(
reporter=_FaceDetectionProgressReporter(self),
)
return FaceDetectionService(
context=FaceDetectionContext(
cfg=self.__cfg,
devs=self.__devs,
mlbox=self.__mlbox,
runtime=self._build_runtime_state(),
progress_emitter=FaceDetectionProgressEmitter(
reporter=_FaceDetectionProgressReporter(self),
deps=FaceDetectionDependencies(
cfg=self.__cfg,
devs=self.__devs,
mlbox=self.__mlbox,
),
runtime=self._build_runtime_state(),
services=services,
settings=FaceDetectionSettings(),
),
logger=logger,
)
def _create_mounting_service(self):
def _create_mounting_service(self) -> MountingService:
return MountingService(
context=MountingContext(
cfg=self.__cfg,
devs=self.__devs,
mount_position=ABR_POS_MOUNT,
deps=MountingDependencies(
cfg=self.__cfg,
devs=self.__devs,
),
settings=MountingSettings(
mount_position=ABR_POS_MOUNT,
),
),
logger=logger,
)
def _create_raster_service(self):
def _create_raster_service(self) -> RasterService:
return RasterService(
context=RasterContext(
cfg=self.__cfg,
devs=self.__devs,
mlbox=self.__mlbox,
jfjoch=self.__jfjoch,
aare=self.__aare,
runtime=self._build_runtime_state(),
state_controller=StateController(
setter=_DAQStateSetter(self),
deps=RasterDependencies(
cfg=self.__cfg,
devs=self.__devs,
mlbox=self.__mlbox,
jfjoch=self.__jfjoch,
aare=self.__aare,
),
runtime=self._build_runtime_state(),
services=self._build_operation_services(),
auto_raster_max_images=self.AUTO_RASTER_MAX_IMAGES,
auto_raster_min_cell_size_mm=self.AUTO_RASTER_MIN_CELL_SIZE_MM,
auto_raster_skip_if_exceed_max_image_threshold=self.AUTO_RASTER_SKIP_IF_EXCEED_MAX_IMAGE_THRESHOLD,
settings=RasterSettings(
auto_raster_max_images=self.AUTO_RASTER_MAX_IMAGES,
auto_raster_min_cell_size_mm=self.AUTO_RASTER_MIN_CELL_SIZE_MM,
auto_raster_skip_if_exceed_max_image_threshold=self.AUTO_RASTER_SKIP_IF_EXCEED_MAX_IMAGE_THRESHOLD,
),
),
logger=logger,
)
#
# def _create_rotation_service(self):
# return RotationService(
# context=RotationContext(
# cfg=self.__cfg,
# devs=self.__devs,
# jfjoch=self.__jfjoch,
# aare=self.__aare,
# sample_provider=lambda: self.sample,
# sample_geometry_provider=lambda: self.sample_geometry,
# status_provider=lambda: self.status,
# set_state=self.__set_state,
# save_screenshot_db=self.save_screenshot_db,
# ),
# logger=logger,
# )
def _create_rotation_service(self) -> RotationService:
return RotationService(
context=RotationContext(
deps=RotationDependencies(
cfg=self.__cfg,
devs=self.__devs,
jfjoch=self.__jfjoch,
aare=self.__aare,
),
runtime=self._build_runtime_state(),
services=self._build_operation_services(),
settings=RotationSettings(),
),
logger=logger,
)
#--------------------------------------------
# Operation Handlers
@@ -1362,6 +1452,56 @@ class AareDAQ:
)
return None
# def _execute_rotation_sequence(self, rotation_request: RotationScanRequest) -> CompletedRotationScan | None:
# """
# Execute rotation scan.
#
# Args:
# rotation_request: Rotation scan parameters
#
# Returns:
# CompletedRotationScan result or None if failed
# """
# try:
# status= self.status
# if self.__cfg.simulated_detector:
# logger.info("Simulated detector mode enabled; skipping JFJoch start.")
# else:
# self.__jfjoch.measure_rotation(rotation_request, status, self.__cfg.xrf)
#
# self.__setup_datacollection(request=rotation_request)
# if self.sample is not None and self.sample.db_id is not None:
# self.__aare.send_sample_event(self.sample.db_id, SampleEventType.COLLECTING)
#
# self.__set_state(BeamlineStateEnum.DataCollection)
# result = self.__rotation(rotation_request)
# self.__set_state(BeamlineStateEnum.SampleAlignment)
# if self.sample is not None and self.sample.db_id is not None:
# self.save_screenshot_db(self.sample.db_id, "scan_preview")
# self.__aare.send_sample_event(self.sample.db_id, SampleEventType.COLLECTED)
# self.__aare.ingest_scan(sample=self.sample, result=result.result,
# geom=self.sample_geometry, beam_mark_pxl=self.__cfg.get_beam_mark(self.zoom))
# return result
# except JFJochCommunicationError as e:
# logger.error(f"Rotation sequence failed due to JFJoch Communication error: {e}")
# self._handle_operation_error(
# operation=DAQOperation.ROTATION,
# sample=self.sample,
# error=e,
# event_type=SampleEventType.COLLECTIONFAILED,
# additional_comment=f"JFJoch communication error: {e}"
# )
# raise
# except Exception as e:
# logger.error(f"Rotation sequence failed: {e}")
# self._handle_operation_error(
# operation=DAQOperation.ROTATION,
# sample=self.sample,
# error=e,
# event_type=SampleEventType.COLLECTIONFAILED
# )
# raise
def _execute_rotation_sequence(self, rotation_request: RotationScanRequest) -> CompletedRotationScan | None:
"""
Execute rotation scan.
@@ -1373,25 +1513,7 @@ class AareDAQ:
CompletedRotationScan result or None if failed
"""
try:
status= self.status
if self.__cfg.simulated_detector:
logger.info("Simulated detector mode enabled; skipping JFJoch start.")
else:
self.__jfjoch.measure_rotation(rotation_request, status, self.__cfg.xrf)
self.__setup_datacollection(request=rotation_request)
if self.sample is not None and self.sample.db_id is not None:
self.__aare.send_sample_event(self.sample.db_id, SampleEventType.COLLECTING)
self.__set_state(BeamlineStateEnum.DataCollection)
result = self.__rotation(rotation_request)
self.__set_state(BeamlineStateEnum.SampleAlignment)
if self.sample is not None and self.sample.db_id is not None:
self.save_screenshot_db(self.sample.db_id, "scan_preview")
self.__aare.send_sample_event(self.sample.db_id, SampleEventType.COLLECTED)
self.__aare.ingest_scan(sample=self.sample, result=result.result,
geom=self.sample_geometry, beam_mark_pxl=self.__cfg.get_beam_mark(self.zoom))
return result
return self._create_rotation_service().run(rotation_request)
except JFJochCommunicationError as e:
logger.error(f"Rotation sequence failed due to JFJoch Communication error: {e}")
self._handle_operation_error(
+55
View File
@@ -0,0 +1,55 @@
from dataclasses import dataclass
from typing import Generic, Protocol, TypeVar
from aare.daq.config import BeamlineConfig
from aare.daq.devices import BeamlineDevices
from aare.daq.operations.common.runtime import DAQRuntimeState
from aare.daq.operations.common.services import OperationServices
class OperationDependencies(Protocol):
"""Marker protocol for operation dependency bundles."""
class OperationSettings(Protocol):
"""Marker protocol for operation settings bundles."""
DepsT = TypeVar("DepsT", bound=OperationDependencies)
SettingsT = TypeVar("SettingsT", bound=OperationSettings)
@dataclass
class BeamlineDependencies:
cfg: BeamlineConfig
devs: BeamlineDevices
@dataclass
class OperationResult:
success: bool
error: Exception | None = None
comment: str | None = None
@property
def failed(self) -> bool:
return not self.success
@dataclass
class BaseOperationContext(Generic[DepsT, SettingsT]):
deps: DepsT
runtime: DAQRuntimeState
services: OperationServices
settings: SettingsT
@property
def sample(self):
return self.runtime.sample
@property
def sample_geometry(self):
return self.runtime.sample_geometry
@property
def status(self):
return self.runtime.status
+3 -61
View File
@@ -1,13 +1,8 @@
from dataclasses import dataclass
from typing import Callable, Protocol, TypeVar
from aare.common.exception_handler import AareException
from typing import Protocol
from aare.common.models import DAQStatusModel, SampleShortInfo
from aare.common.sample_geometry import SampleGeometryModel
from aare.daq.config import BeamlineStateEnum
from aare.daq.operations.screenshot.service import ScreenshotService
class SampleProvider(Protocol):
@@ -25,23 +20,7 @@ class StatusProvider(Protocol):
def status(self) -> DAQStatusModel: ...
class StateSetter(Protocol):
def set_state(self, target: BeamlineStateEnum) -> None: ...
class SmargonTraceAppender(Protocol):
def append_smargon_trace(self, *, sample_id: int | None, event: str) -> None: ...
class PredictionGetter(Protocol):
def get_predictions(self): ...
class FaceDetectionProgressReporter(Protocol):
def emit_progress(self, payload: dict) -> None: ...
@dataclass
@dataclass(frozen=True)
class DAQRuntimeState:
sample_provider: SampleProvider
sample_geometry_provider: SampleGeometryProvider
@@ -57,41 +36,4 @@ class DAQRuntimeState:
@property
def status(self) -> DAQStatusModel:
return self.status_provider.status
@dataclass
class StateController:
setter: StateSetter
def set_state(self, target: BeamlineStateEnum) -> None:
self.setter.set_state(target)
@dataclass
class TraceWriter:
appender: SmargonTraceAppender
def append_smargon_trace(self, *, sample_id: int | None, event: str) -> None:
self.appender.append_smargon_trace(sample_id=sample_id, event=event)
@dataclass
class PredictionProvider:
getter: PredictionGetter
def get_predictions(self):
return self.getter.get_predictions()
@dataclass
class FaceDetectionProgressEmitter:
reporter: FaceDetectionProgressReporter
def emit_progress(self, payload: dict) -> None:
self.reporter.emit_progress(payload)
@dataclass
class OperationServices:
screenshots: ScreenshotService
return self.status_provider.status
+118
View File
@@ -0,0 +1,118 @@
from dataclasses import dataclass
from typing import Protocol
from aare.daq.config import BeamlineStateEnum
from aare.daq.operations.screenshot.service import ScreenshotService
class StateSetter(Protocol):
def set_state(self, target: BeamlineStateEnum) -> None: ...
class SmargonTraceAppender(Protocol):
def append_smargon_trace(self, *, sample_id: int | None, event: str) -> None: ...
class PredictionGetter(Protocol):
def get_predictions(self): ...
class FaceDetectionProgressReporter(Protocol):
def emit_progress(self, payload: dict) -> None: ...
class SampleEventSender(Protocol):
def send_sample_event(self, sample_id: int, event_type, comment: str | None = None) -> None: ...
class ScanIngestor(Protocol):
def ingest_scan(self, *, sample, result, geom, beam_mark_pxl) -> None: ...
def ingest_gridscan(self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl) -> None: ...
class DatacollectionSetupRunner(Protocol):
def prepare(self, request, screening: bool = False) -> None: ...
@dataclass(frozen=True)
class StateController:
setter: StateSetter
def set_state(self, target: BeamlineStateEnum) -> None:
self.setter.set_state(target)
@dataclass(frozen=True)
class TraceWriter:
appender: SmargonTraceAppender
def append_smargon_trace(self, *, sample_id: int | None, event: str) -> None:
self.appender.append_smargon_trace(sample_id=sample_id, event=event)
@dataclass(frozen=True)
class PredictionProvider:
getter: PredictionGetter
def get_predictions(self):
return self.getter.get_predictions()
@dataclass(frozen=True)
class FaceDetectionProgressEmitter:
reporter: FaceDetectionProgressReporter
def emit_progress(self, payload: dict) -> None:
self.reporter.emit_progress(payload)
@dataclass(frozen=True)
class SampleEventPublisher:
sender: SampleEventSender
def send(self, sample_id: int, event_type, comment: str | None = None) -> None:
self.sender.send_sample_event(sample_id, event_type, comment)
@dataclass(frozen=True)
class ScanIngestionService:
ingestor: ScanIngestor
def ingest_scan(self, *, sample, result, geom, beam_mark_pxl) -> None:
self.ingestor.ingest_scan(
sample=sample,
result=result,
geom=geom,
beam_mark_pxl=beam_mark_pxl,
)
def ingest_gridscan(self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl) -> None:
self.ingestor.ingest_gridscan(
sample=sample,
raster_result=raster_result,
raster_request=raster_request,
geom=geom,
com=com,
beam_mark_pxl=beam_mark_pxl,
)
@dataclass(frozen=True)
class DataCollectionPreparer:
runner: DatacollectionSetupRunner
def prepare(self, request, screening: bool = False) -> None:
self.runner.prepare(request, screening=screening)
@dataclass
class OperationServices:
screenshots: ScreenshotService
state: StateController | None = None
traces: TraceWriter | None = None
predictions: PredictionProvider | None = None
face_detection_progress: FaceDetectionProgressEmitter | None = None
events: SampleEventPublisher | None = None
ingestion: ScanIngestionService | None = None
datacollection: DataCollectionPreparer | None = None
@@ -1,18 +1,29 @@
from dataclasses import dataclass
from aare.daq.config import BeamlineConfig
from aare.daq.devices import BeamlineDevices
from aare.daq.mlbox import MlBox
from aare.daq.operations.common.runtime import DAQRuntimeState, FaceDetectionProgressEmitter
from aare.daq.operations.common.models import (
BaseOperationContext,
BeamlineDependencies,
)
@dataclass
class FaceDetectionContext:
cfg: BeamlineConfig
devs: BeamlineDevices
class FaceDetectionDependencies(BeamlineDependencies):
mlbox: MlBox
runtime: DAQRuntimeState
progress_emitter: FaceDetectionProgressEmitter
@dataclass
class FaceDetectionSettings:
steps: int = 14
step_size: int = 15
face_min_ratio: float = 0.3
@dataclass
class FaceDetectionContext(
BaseOperationContext[FaceDetectionDependencies, FaceDetectionSettings]
):
pass
@dataclass
@@ -1,9 +1,10 @@
import time
import aare.daq.operations.face_detection.utils as fd
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.logger_events import log_duration, log_ml_bundle_meta
from aare.common.models import MLBoxModel, ZoomModeEnum
import aare.daq.operations.face_detection.utils as fd
from aare.daq.mlbox import MLBoxPredictionResult
from aare.daq.operations.face_detection.models import (
FaceDetectionContext,
@@ -16,6 +17,12 @@ class FaceDetectionService:
self.ctx = context
self.logger = logger
def _progress_emitter(self):
emitter = self.ctx.services.face_detection_progress
if emitter is None:
raise RuntimeError("FaceDetectionService requires services.face_detection_progress")
return emitter
def _log_warning(self, message: str) -> None:
warning = getattr(self.logger, "warning", None)
if callable(warning):
@@ -32,7 +39,7 @@ class FaceDetectionService:
angle: int,
boxes_face: dict[int, tuple[float, float, float, float]],
) -> None:
self.ctx.progress_emitter.emit_progress(
self._progress_emitter().emit_progress(
{
"running": True,
"current_angle_deg": angle,
@@ -49,7 +56,7 @@ class FaceDetectionService:
"height_fit": {},
"area_fit": {},
}
self.ctx.progress_emitter.emit_progress(payload)
self._progress_emitter().emit_progress(payload)
return payload
def _centre_correction(self, model: MLBoxModel, tolerance: float = 0.2) -> None:
@@ -65,28 +72,31 @@ class FaceDetectionService:
if beam_y != 0 and abs(centre_y - beam_y) / abs(beam_y) > tolerance:
coord = geom.picture_to_smargon(Coordinate(x=beam_x, y=centre_y))
self.ctx.devs.smargon_pos = SmargonCoordinate(sh_mm=coord)
self.ctx.devs.smargon_wait(60)
self.ctx.deps.devs.smargon_pos = SmargonCoordinate(sh_mm=coord)
self.ctx.deps.devs.smargon_wait(60)
def run(
self,
*,
steps: int = 14,
step_size: int = 15,
face_min_ratio: float = 0.3,
self,
*,
steps: int | None = None,
step_size: int | None = None,
face_min_ratio: float | None = None,
) -> FaceDetectionResult:
steps = self.ctx.settings.steps if steps is None else steps
step_size = self.ctx.settings.step_size if step_size is None else step_size
face_min_ratio = self.ctx.settings.face_min_ratio if face_min_ratio is None else face_min_ratio
try:
self.ctx.cfg.zoom_mode = ZoomModeEnum.LoopCenter
self.ctx.devs.lamp_light = 2.5
self.ctx.deps.cfg.zoom_mode = ZoomModeEnum.LoopCenter
self.ctx.deps.devs.lamp_light = 2.5
zoom_value = self.ctx.devs.zoom
zoom_value = self.ctx.deps.devs.zoom
self.logger.info("Starting face detection sequence")
self.ctx.devs.set_zoom(zoom_value, wait=True)
self.ctx.deps.devs.set_zoom(zoom_value, wait=True)
boxes_face: dict[int, tuple[float, float, float, float]] = {}
boxes_loop: dict[int, tuple[float, float, float, float]] = {}
curr_angle = int(self.ctx.devs.aerotech_omega)
curr_angle = int(self.ctx.deps.devs.aerotech_omega)
total_range = steps * step_size + 1
start_angle = curr_angle if curr_angle + total_range < 720 else 0
end_angle = curr_angle + total_range
@@ -94,7 +104,7 @@ class FaceDetectionService:
for angle in range(start_angle, end_angle, step_size):
self.logger.debug(f"moving to angle: {angle}")
rotate_time = time.perf_counter()
self.ctx.devs.aerotech_omega = angle
self.ctx.deps.devs.aerotech_omega = angle
log_duration(
self.logger,
"Completed Aerotech move during face detection",
@@ -102,7 +112,7 @@ class FaceDetectionService:
extra={"angle_deg": angle},
)
prediction_result: MLBoxPredictionResult = self.ctx.mlbox.predict(
prediction_result: MLBoxPredictionResult = self.ctx.deps.mlbox.predict(
preferred_class=(3, 0),
return_image=True,
return_bundle_meta=True,
@@ -183,7 +193,7 @@ class FaceDetectionService:
self.logger.debug(f"chosen fit: {best_name}")
self.logger.info(f"best angle: {flat_face_angle}")
self.ctx.devs.aerotech_omega = flat_face_angle
self.ctx.deps.devs.aerotech_omega = flat_face_angle
samples_out = fd.get_samples_out(boxes)
self.logger.info(f"Face detection sequence complete")
@@ -206,7 +216,7 @@ class FaceDetectionService:
"best_angle_deg": best_fit_angle_area,
},
}
self.ctx.progress_emitter.emit_progress(payload)
self._progress_emitter().emit_progress(payload)
return FaceDetectionResult(success=True, payload=payload)
except Exception as e:
@@ -219,4 +229,4 @@ class FaceDetectionService:
comment="Face detection sequence failed",
)
finally:
self.ctx.cfg.zoom_mode = ZoomModeEnum.User
self.ctx.deps.cfg.zoom_mode = ZoomModeEnum.User
@@ -104,7 +104,7 @@ class LoopCenteringAnalyzer:
if box.cls == MLBoxType.PIN:
pin = box
best_box = self.ctx.mlbox.get_preferred_class_box(
best_box = self.ctx.deps.mlbox.get_preferred_class_box(
boxes,
(MLBoxType.CRYSTAL, MLBoxType.LOOP_FACE, MLBoxType.LOOP_ALL, MLBoxType.PIN),
)
@@ -124,7 +124,7 @@ class LoopCenteringAnalyzer:
centre_y = y1 + (y2 - y1) / 2
centre_x = x1
elif pin:
position_dict = self.ctx.mlbox.check_box_relation(pin, best_box)
position_dict = self.ctx.deps.mlbox.check_box_relation(pin, best_box)
if position_dict["overlap_y"] and position_dict["overlap_x"]:
centre_y = y1 + (y2 - y1) / 2
centre_x = x1
@@ -157,7 +157,10 @@ class LoopCenteringAnalyzer:
zoom_value: float,
sample_id: int | None,
) -> AngleAnalysis:
prediction_result = self.ctx.prediction_provider.get_predictions()
if self.ctx.services.predictions is None:
raise RuntimeError("LoopCenteringAnalyzer requires services.predictions")
prediction_result = self.ctx.services.predictions.get_predictions()
log_ml_bundle_meta(
self.logger,
f"loop_center_angle_{angle_deg}_zoom_{zoom_value:.0f}",
@@ -1,17 +1,18 @@
from dataclasses import dataclass, field
from aare.common.coordinate import SmargonCoordinate
from aare.daq.config import BeamlineConfig
from aare.daq.devices import BeamlineDevices
from aare.daq.mlbox import MlBox
from aare.daq.operations.common.runtime import (
DAQRuntimeState,
OperationServices,
PredictionProvider,
TraceWriter,
from aare.daq.operations.common.models import (
BaseOperationContext,
BeamlineDependencies,
)
@dataclass
class LoopCenteringDependencies(BeamlineDependencies):
mlbox: MlBox
@dataclass
class AngleAnalysis:
angle_deg: int
@@ -43,12 +44,7 @@ class LoopCenteringSettings:
@dataclass
class LoopCenteringContext:
cfg: BeamlineConfig
devs: BeamlineDevices
mlbox: MlBox
settings: LoopCenteringSettings
runtime: DAQRuntimeState
services: OperationServices
trace_writer: TraceWriter
prediction_provider: PredictionProvider
class LoopCenteringContext(
BaseOperationContext[LoopCenteringDependencies, LoopCenteringSettings]
):
pass
@@ -31,7 +31,7 @@ class LoopCenteringService:
) -> AngleAnalysis:
self.logger.debug(f"Moving to new omega: {angle}")
time_to_move_aerotech = time.perf_counter()
self.ctx.devs.aerotech_omega = angle
self.ctx.deps.devs.aerotech_omega = angle
log_duration(
self.logger,
"Completed Aerotech move during loop centering",
@@ -47,8 +47,8 @@ class LoopCenteringService:
if analysis.has_valid_target and not analysis.ignore_only and analysis.final_target is not None:
time_to_move_smargon = time.perf_counter()
self.ctx.devs.smargon_pos = analysis.final_target
self.ctx.devs.smargon_wait(60)
self.ctx.deps.devs.smargon_pos = analysis.final_target
self.ctx.deps.devs.smargon_wait(60)
log_duration(
self.logger,
"Completed Smargon move during loop centering",
@@ -58,7 +58,7 @@ class LoopCenteringService:
analysis.moved = True
if sample_id is not None and trace_all_alc_moves:
self.ctx.trace_writer.append_smargon_trace(
self.ctx.services.traces.append_smargon_trace(
sample_id=sample_id,
event=f"alc_move_zoom_{zoom_value:.0f}_angle_{angle}",
)
@@ -96,11 +96,11 @@ class LoopCenteringService:
try:
zoom_value = settings.zoom_value
if abs(self.ctx.devs.zoom - zoom_value) > 1e-6:
self.ctx.devs.samcam_auto(AutoEnum.AUTO)
self.ctx.devs.zoom = zoom_value
if abs(self.ctx.deps.devs.zoom - zoom_value) > 1e-6:
self.ctx.deps.devs.samcam_auto(AutoEnum.AUTO)
self.ctx.deps.devs.zoom = zoom_value
time.sleep(0.2)
self.ctx.devs.samcam_auto(AutoEnum.ONCE)
self.ctx.deps.devs.samcam_auto(AutoEnum.ONCE)
if sample_id is not None:
self.logger.info(
@@ -158,7 +158,7 @@ class LoopCenteringService:
if valid_seen_correction:
if sample_id is not None:
self.logger.info(f"sample {sample_id} centered")
self.ctx.trace_writer.append_smargon_trace(sample_id=sample_id, event="alc_success")
self.ctx.services.traces.append_smargon_trace(sample_id=sample_id, event="alc_success")
return LoopCenteringResult(success=True)
failure_reason = f"ALC exceeded max attempts ({max_attempts})"
+13 -5
View File
@@ -2,15 +2,23 @@ from dataclasses import dataclass
from aare.common.coordinate import AerotechCoordinate
from aare.common.models import SampleShortInfo
from aare.daq.config import BeamlineConfig
from aare.daq.devices import BeamlineDevices
from aare.daq.operations.common.models import BeamlineDependencies
@dataclass
class MountingDependencies(BeamlineDependencies):
pass
@dataclass
class MountingSettings:
mount_position: AerotechCoordinate
@dataclass
class MountingContext:
cfg: BeamlineConfig
devs: BeamlineDevices
mount_position: AerotechCoordinate
deps: MountingDependencies
settings: MountingSettings
@dataclass
+25 -25
View File
@@ -20,22 +20,22 @@ class MountingService:
self.logger = logger
def reset_mount_failure_counter(self, reason: str) -> None:
previous_count = self.ctx.cfg.get_mount_failure_streak()
previous_count = self.ctx.deps.cfg.get_mount_failure_streak()
if previous_count > 0:
self.logger.warning(
f"Resetting mount failure counter from {previous_count} due to: {reason}"
)
else:
self.logger.debug(f"Mount failure counter already clear: {reason}")
self.ctx.cfg.reset_mount_failure_streak()
self.ctx.deps.cfg.reset_mount_failure_streak()
def _magnet_position_sensor_check(self, timeout: float = 1.0) -> None:
if self.ctx.devs.magnet_position_sensor.value != 0:
if self.ctx.deps.devs.magnet_position_sensor.value != 0:
self.logger.warning(
"Goniometer is not in position based on magnet position sensor readout"
)
for _ in range(round(timeout * 10.0)):
if self.ctx.devs.magnet_position_sensor.value == 0:
if self.ctx.deps.devs.magnet_position_sensor.value == 0:
return
time.sleep(0.1)
@@ -45,7 +45,7 @@ class MountingService:
raise Exception("Goniometer is not in position based on magnet position sensor readout")
def _handle_consecutive_mount_failure(self) -> None:
count = self.ctx.cfg.increment_mount_failure_streak()
count = self.ctx.deps.cfg.increment_mount_failure_streak()
self.logger.warning(f"Consecutive mount failure count is now {count}")
if count == DRY_AFTER_FAIL_COUNT:
@@ -72,7 +72,7 @@ class MountingService:
def _mount_handler(self, target) -> None:
self._prepare_mount_hardware()
value = self.ctx.devs.tell.mount(
value = self.ctx.deps.devs.tell.mount(
address=target.tell_address(),
force=True,
auto_unmount=True,
@@ -82,8 +82,8 @@ class MountingService:
)
if isinstance(value, str):
self.ctx.devs.tell.check_command_ok()
self.logger.error(f"{self.ctx.devs.tell.get_result(self.ctx.devs.tell._last_cmd_id)}")
self.ctx.deps.devs.tell.check_command_ok()
self.logger.error(f"{self.ctx.deps.devs.tell.get_result(self.ctx.deps.devs.tell._last_cmd_id)}")
self.logger.error(f"Unexpected string response from Tell mount: {value}")
raise CriticalTellException(f"Critical error in TELL mount: unexpected response '{value}'")
@@ -115,38 +115,38 @@ class MountingService:
raise CriticalTellException(f"Critical error in TELL mount: {message}")
def _prepare_mount_hardware(self) -> None:
self.ctx.devs.smargon_move_home()
self.ctx.devs.aerotech_pos = self.ctx.mount_position
self.ctx.deps.devs.smargon_move_home()
self.ctx.deps.devs.aerotech_pos = self.ctx.settings.mount_position
self._magnet_position_sensor_check(timeout=360.0)
self.ctx.devs.tell.check_enable_motion()
self.ctx.devs.tell.wait_not_busy()
self.ctx.devs.tell.set_in_mount_position(True)
self.ctx.deps.devs.tell.check_enable_motion()
self.ctx.deps.devs.tell.wait_not_busy()
self.ctx.deps.devs.tell.set_in_mount_position(True)
def _unmount_current_sample(self, timeout: float = 60.0):
self._prepare_mount_hardware()
previous_sample = self.ctx.cfg.current_sample
previous_sample = self.ctx.deps.cfg.current_sample
if previous_sample is not None:
self.logger.debug(f"Unmounting sample: {previous_sample}")
self.ctx.devs.tell.unmount(wait=True, timeout=timeout)
self.ctx.cfg.current_sample = None
self.ctx.deps.devs.tell.unmount(wait=True, timeout=timeout)
self.ctx.deps.cfg.current_sample = None
return previous_sample
def dry(self, *, park: bool, unmount: bool = False) -> None:
self.ctx.devs.tell.check_enable_motion()
self.ctx.devs.tell.wait_not_busy()
self.ctx.devs.tell.set_in_mount_position(True)
self.ctx.deps.devs.tell.check_enable_motion()
self.ctx.deps.devs.tell.wait_not_busy()
self.ctx.deps.devs.tell.set_in_mount_position(True)
if self.ctx.cfg.current_sample and unmount:
if self.ctx.deps.cfg.current_sample and unmount:
self._prepare_mount_hardware()
self._unmount_current_sample(timeout=60.0)
if park:
self.ctx.devs.tell.dry(wait_cold=-1, wait=True)
self.ctx.deps.devs.tell.dry(wait_cold=-1, wait=True)
else:
self.ctx.devs.tell.dry(wait=True)
self.ctx.deps.devs.tell.dry(wait=True)
def execute(self, *, target) -> MountingResult:
previous_sample = self.ctx.cfg.current_sample
previous_sample = self.ctx.deps.cfg.current_sample
try:
if target is None:
@@ -173,8 +173,8 @@ class MountingService:
raise
self.logger.debug("Mounting succeeded, resetting mount failure counter")
self.ctx.cfg.reset_mount_failure_streak()
self.ctx.cfg.current_sample = target
self.ctx.deps.cfg.reset_mount_failure_streak()
self.ctx.deps.cfg.current_sample = target
return MountingResult(
success=True,
+13 -24
View File
@@ -2,46 +2,35 @@ from dataclasses import dataclass
from typing import TYPE_CHECKING
from aare.common.raster_grid import RasterGridRequest
from aare.daq.operations.common.runtime import (
DAQRuntimeState,
OperationServices,
StateController,
from aare.daq.mlbox import MlBox
from aare.daq.operations.common.models import (
BaseOperationContext,
BeamlineDependencies,
)
if TYPE_CHECKING:
from aare.daq.aaredb import AareWrapper
from aare.daq.config import BeamlineConfig
from aare.daq.devices import BeamlineDevices
from aare.daq.mlbox import MlBox
from aare.devices.jfjoch import JFJochWrapper
@dataclass
class RasterContext:
cfg: "BeamlineConfig"
devs: "BeamlineDevices"
mlbox: "MlBox"
class RasterDependencies(BeamlineDependencies):
mlbox: MlBox
jfjoch: "JFJochWrapper"
aare: "AareWrapper"
runtime: DAQRuntimeState
state_controller: StateController
services: OperationServices
@dataclass
class RasterSettings:
auto_raster_max_images: int
auto_raster_min_cell_size_mm: float
auto_raster_skip_if_exceed_max_image_threshold: bool
wait_for_screenshot_s: float = 0.3
@property
def sample(self):
return self.runtime.sample
@property
def sample_geometry(self):
return self.runtime.sample_geometry
@property
def status(self):
return self.runtime.status
@dataclass
class RasterContext(BaseOperationContext[RasterDependencies, RasterSettings]):
pass
@dataclass
+99 -57
View File
@@ -85,7 +85,7 @@ class RasterService:
return
try:
diffraction_image = self.ctx.jfjoch.get_diffraction_image(image_id)
diffraction_image = self.ctx.deps.jfjoch.get_diffraction_image(image_id)
except NotFoundException:
self.logger.warning(
"JFJoch diffraction preview image was not found after raster; continuing without upload",
@@ -101,7 +101,7 @@ class RasterService:
)
return
self.ctx.aare.upload_jpg(sample_id, filename, diffraction_image)
self.ctx.deps.aare.upload_jpg(sample_id, filename, diffraction_image)
def ml_bounding_box(
self,
@@ -109,15 +109,15 @@ class RasterService:
filename: str | None = None,
) -> RasterGridRequest | None:
return get_ml_bounding_box(
mlbox=self.ctx.mlbox,
mlbox=self.ctx.deps.mlbox,
sample=self.ctx.sample,
sample_geometry=self.ctx.sample_geometry,
upload_image=self.ctx.aare.upload_image,
upload_image=self.ctx.deps.aare.upload_image,
logger=self.logger,
filename=filename,
max_images=self.ctx.auto_raster_max_images,
min_cell_size_mm=self.ctx.auto_raster_min_cell_size_mm,
skip_if_exceed_max_image_threshold=self.ctx.auto_raster_skip_if_exceed_max_image_threshold,
max_images=self.ctx.settings.auto_raster_max_images,
min_cell_size_mm=self.ctx.settings.auto_raster_min_cell_size_mm,
skip_if_exceed_max_image_threshold=self.ctx.settings.auto_raster_skip_if_exceed_max_image_threshold,
)
def auto_center_line_scan_top_left(
@@ -137,7 +137,7 @@ class RasterService:
line_scan_centre = geom.picture_to_smargon(Coordinate(x=beam_x_pxl, y=beam_y_pxl))
n_y = default_n_y
prediction_result: MLBoxPredictionResult = self.ctx.mlbox.predict(
prediction_result: MLBoxPredictionResult = self.ctx.deps.mlbox.predict(
preferred_class=(3, 0),
return_image=False,
return_bundle_meta=True,
@@ -275,17 +275,20 @@ class RasterService:
def execute(
self,
request: RasterGridRequest,
wait_for_screenshot: float | None = 0.3,
wait_for_screenshot: float | None = None,
) -> CompletedRasterGridElem:
total_time = request.exp_time_s * request.n_x * request.n_y + request.n_y * 0.3
if wait_for_screenshot is None:
wait_for_screenshot = self.ctx.settings.wait_for_screenshot_s
if total_time > 1200:
raise RasterScanException(f"Raster scan is too long {total_time}s > 20min")
status = self.ctx.status
smargon_top_left = request.smargon_top_left
self.ctx.devs.set_smargon_pos(
self.ctx.deps.devs.set_smargon_pos(
SmargonCoordinate(
sh_mm=smargon_top_left.sh_mm,
phi_deg=smargon_top_left.phi_deg,
@@ -304,17 +307,17 @@ class RasterService:
try:
if self.ctx.sample is not None and self.ctx.sample.db_id is not None:
self.ctx.aare.create_gridscan_run(self.ctx.sample, request, status)
self.ctx.deps.aare.create_gridscan_run(self.ctx.sample, request, status)
if not self.ctx.cfg.simulated_detector:
self.ctx.jfjoch.wait_till_running(timeout=60.0)
if not self.ctx.deps.cfg.simulated_detector:
self.ctx.deps.jfjoch.wait_till_running(timeout=60.0)
else:
self.logger.info("Simulated detector mode enabled; faking jfjoch intilalisation.")
self.logger.debug(
f"Starting grid scan with {request.n_x}x{request.n_y} points, exp time {request.exp_time_s}s"
)
self.ctx.devs.aerotech.grid_scan(
self.ctx.deps.devs.aerotech.grid_scan(
grid_elem_size_y_um=request.grid_size_mm.y * 1000,
grid_elem_size_x_um=request.grid_size_mm.x * 1000,
grid_elem_count_x=request.n_x,
@@ -323,24 +326,24 @@ class RasterService:
run_async=True,
)
self.ctx.devs.aerotech.wait_till_done(
self.ctx.deps.devs.aerotech.wait_till_done(
timeout=int(round(total_time + total_time * 0.1 + 60, 0))
)
if isinstance(self.ctx.cfg.abr_meas_pos, Coordinate):
coord = self.ctx.cfg.abr_meas_pos
if isinstance(self.ctx.deps.cfg.abr_meas_pos, Coordinate):
coord = self.ctx.deps.cfg.abr_meas_pos
else:
coord = self.ctx.cfg.abr_meas_pos.at_mm
self.ctx.devs.aerotech_pos = AerotechCoordinate(
coord = self.ctx.deps.cfg.abr_meas_pos.at_mm
self.ctx.deps.devs.aerotech_pos = AerotechCoordinate(
at_mm=coord,
omega_deg=self.ctx.devs.aerotech_omega,
omega_deg=self.ctx.deps.devs.aerotech_omega,
)
self.ctx.devs.aerotech.wait_till_done(timeout=60)
self.ctx.deps.devs.aerotech.wait_till_done(timeout=60)
x = None
y = None
if self.ctx.cfg.simulated_detector:
if self.ctx.deps.cfg.simulated_detector:
self.logger.info("Simulated detector mode enabled; using fake raster result.")
scan_result = generate_no_beam_scan_result(request)
result_array = create_quality_filtered_array(
@@ -361,7 +364,7 @@ class RasterService:
chi_deg=request.smargon_top_left.chi_deg,
)
else:
scan_result = self.ctx.jfjoch.wait_till_done(60)
scan_result = self.ctx.deps.jfjoch.wait_till_done(60)
if scan_result is None:
self.logger.error(
"JFJoch returned no ScanResult for raster",
@@ -439,8 +442,8 @@ class RasterService:
chi_deg=request.smargon_top_left.chi_deg,
)
self.ctx.devs.smargon_pos = target_smargon
self.ctx.devs.smargon_wait(timeout=180)
self.ctx.deps.devs.smargon_pos = target_smargon
self.ctx.deps.devs.smargon_wait(timeout=180)
self.logger.info(
"Moved Smargon to raster centre",
@@ -460,7 +463,10 @@ class RasterService:
sample_id = self.ctx.sample.db_id if self.ctx.sample and self.ctx.sample.db_id is not None else None
if sample_id:
self.logger.debug("moving to XtalSnapshot to take a screenshot of the sample")
self.ctx.state_controller.set_state(BeamlineStateEnum.XtalSnapshot)
if self.ctx.services.state is None:
raise RuntimeError("RasterService requires services.state")
self.ctx.services.state.set_state(BeamlineStateEnum.XtalSnapshot)
if wait_for_screenshot and wait_for_screenshot > 0:
time.sleep(wait_for_screenshot)
@@ -469,14 +475,24 @@ class RasterService:
f"{sample_id}_post_raster_{int(request.omega_deg)}deg",
)
self.ctx.aare.ingest_gridscan(
sample=self.ctx.sample,
raster_result=scan_result,
raster_request=request,
geom=self.ctx.sample_geometry,
com=None,
beam_mark_pxl=self.ctx.cfg.get_beam_mark(self.ctx.devs.zoom),
)
if self.ctx.services.ingestion is not None:
self.ctx.services.ingestion.ingest_gridscan(
sample=self.ctx.sample,
raster_result=scan_result,
raster_request=request,
geom=self.ctx.sample_geometry,
com=None,
beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom),
)
else:
self.ctx.deps.aare.ingest_gridscan(
sample=self.ctx.sample,
raster_result=scan_result,
raster_request=request,
geom=self.ctx.sample_geometry,
com=None,
beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom),
)
if com is not None and com.max_image is not None:
diffraction_image_filename = (
@@ -498,11 +514,19 @@ class RasterService:
scan_result=scan_result,
request=request,
)
self.ctx.aare.send_sample_event(
self.ctx.sample.db_id,
event_type=SampleEventType.RASTERED,
comment=f"Raster completed at {request.omega_deg:.1f} deg",
)
if self.ctx.services.events is not None:
self.ctx.services.events.send(
self.ctx.sample.db_id,
SampleEventType.RASTERED,
comment=f"Raster completed at {request.omega_deg:.1f} deg",
)
else:
self.ctx.deps.aare.send_sample_event(
self.ctx.sample.db_id,
event_type=SampleEventType.RASTERED,
comment=f"Raster completed at {request.omega_deg:.1f} deg",
)
self.logger.info(
"Raster finished",
@@ -556,11 +580,18 @@ class RasterService:
),
)
self.ctx.aare.send_sample_event(
sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
if self.ctx.services.events is not None:
self.ctx.services.events.send(
sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
else:
self.ctx.deps.aare.send_sample_event(
sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
r = self.ml_bounding_box(sample.db_id, f"ml_{geom.omega_deg:.2f}deg")
@@ -575,7 +606,7 @@ class RasterService:
},
),
)
self.ctx.devs.aerotech_omega = geom.omega_deg + 90.0
self.ctx.deps.devs.aerotech_omega = geom.omega_deg + 90.0
time.sleep(0.2)
r = self.ml_bounding_box(sample.db_id, f"ml_{geom.omega_deg + 90.0:.2f}deg")
@@ -620,14 +651,16 @@ class RasterService:
grid.omega_deg = geom.omega_deg
status = self.ctx.status
if not self.ctx.cfg.simulated_detector:
if not self.ctx.deps.cfg.simulated_detector:
self.logger.info("initialise detector")
self.ctx.jfjoch.measure_raster(grid, status)
self.ctx.deps.jfjoch.measure_raster(grid, status)
self.logger.info("detector initialised")
else:
self.logger.info("Simulated detector mode enabled; using fake raster result.")
self.ctx.state_controller.set_state(BeamlineStateEnum.DataCollection)
if self.ctx.services.state is None:
raise RuntimeError("RasterService requires services.state")
self.ctx.services.state.set_state(BeamlineStateEnum.DataCollection)
self.logger.info(
"Running first auto-center raster",
@@ -644,12 +677,19 @@ class RasterService:
res1 = self.execute(grid)
grid.omega_deg += 90
self.ctx.aare.send_sample_event(
self.ctx.sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
self.ctx.devs.aerotech_omega = grid.omega_deg
if self.ctx.services.events is not None:
self.ctx.services.events.send(
self.ctx.sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
else:
self.ctx.deps.aare.send_sample_event(
self.ctx.sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
self.ctx.deps.devs.aerotech_omega = grid.omega_deg
grid.n_x = 1
grid.file_prefix = f"{old_prefix}_{grid.omega_deg}deg"
@@ -681,14 +721,16 @@ class RasterService:
)
status = self.ctx.status
if not self.ctx.cfg.simulated_detector:
if not self.ctx.deps.cfg.simulated_detector:
self.logger.info(f"initialise detector for raster at {grid.omega_deg}")
self.ctx.jfjoch.measure_raster(grid, status)
self.ctx.deps.jfjoch.measure_raster(grid, status)
self.logger.info("detector initialised")
else:
self.logger.info("Simulated detector mode enabled; using fake raster result.")
self.ctx.state_controller.set_state(BeamlineStateEnum.DataCollection)
if self.ctx.services.state is None:
raise RuntimeError("RasterService requires services.state")
self.ctx.services.state.set_state(BeamlineStateEnum.DataCollection)
self.logger.info(
"Running second auto-center raster",
@@ -0,0 +1,27 @@
from dataclasses import dataclass
from typing import TYPE_CHECKING
from aare.daq.operations.common.models import (
BaseOperationContext,
BeamlineDependencies,
)
if TYPE_CHECKING:
from aare.daq.aaredb import AareWrapper
from aare.devices.jfjoch import JFJochWrapper
@dataclass
class RotationDependencies(BeamlineDependencies):
jfjoch: "JFJochWrapper"
aare: "AareWrapper"
@dataclass
class RotationSettings:
preview_filename: str = "scan_preview"
@dataclass
class RotationContext(BaseOperationContext[RotationDependencies, RotationSettings]):
pass
+119
View File
@@ -0,0 +1,119 @@
import copy
import time
from aareDB import SampleEventType
from aare.common.coordinate import SmargonCoordinate
from aare.common.models import BeamlineStateEnum
from aare.common.rotation_scan import CompletedRotationScan, RotationScanRequest
from aare.daq.operations.common.simulate_scan_result import build_fake_rotation_result
from aare.daq.operations.rotation.models import RotationContext
class RotationService:
def __init__(self, *, context: RotationContext, logger):
self.ctx = context
self.logger = logger
def _prepare_detector(self, request: RotationScanRequest) -> None:
if self.ctx.deps.cfg.simulated_detector:
self.logger.info("Simulated detector mode enabled; skipping JFJoch start.")
return
self.ctx.deps.jfjoch.measure_rotation(
request,
self.ctx.status,
self.ctx.deps.cfg.xrf,
)
def _execute_scan(self, request: RotationScanRequest) -> CompletedRotationScan:
omega_start = self.ctx.deps.devs.aerotech_omega
if request.exp_time_s < 0.004:
self.logger.warning("Exposure time too short for PXII rotation scan")
request.exp_time_s = 0.004
total_time = request.exp_time_s * request.steps
sample = self.ctx.sample
if sample is not None and sample.db_id is not None:
self.ctx.deps.aare.create_rotation_run(sample, request, self.ctx.status)
if not self.ctx.deps.cfg.simulated_detector:
self.ctx.deps.jfjoch.wait_till_running(timeout=60.0)
if request.screening:
self.ctx.deps.devs.aerotech.screening_scan(
rotation_deg=request.steps * request.incr_omega_deg,
wedge_deg=request.wedge_omega_deg,
time_sec=total_time,
steps=request.steps,
run_async=True,
)
else:
self.ctx.deps.devs.aerotech.rotation_scan(
rotation_deg=request.steps * request.incr_omega_deg,
time_sec=total_time,
start_pos_deg=request.start_omega_deg,
run_async=True,
)
if request.start is not None and request.end is not None:
smargon_time_step = request.exp_time_s / float(request.steps)
pos_step = (request.end.sh_mm - request.start.sh_mm) * (1.0 / float(request.steps))
for i in range(request.steps):
self.ctx.deps.devs.smargon.target = SmargonCoordinate(
sh_mm=request.start.sh_mm + pos_step * i
)
time.sleep(smargon_time_step)
self.ctx.deps.devs.aerotech.wait_till_done(timeout=int(round(total_time + 60, 0)))
self.ctx.deps.devs.aerotech_omega = omega_start
if self.ctx.deps.cfg.simulated_detector:
self.logger.warning("Detector in simulation mode, returning fake zero rotation result.")
return build_fake_rotation_result(request, start_angle=float(omega_start))
scan_result = self.ctx.deps.jfjoch.wait_till_done(60)
return CompletedRotationScan(
request=copy.deepcopy(request),
result=scan_result,
)
def run(self, request: RotationScanRequest) -> CompletedRotationScan:
sample = self.ctx.sample
self._prepare_detector(request)
self.ctx.services.datacollection.prepare(request)
if sample is not None and sample.db_id is not None and self.ctx.services.events is not None:
self.ctx.services.events.send(sample.db_id, SampleEventType.COLLECTING)
if self.ctx.services.state is not None:
self.ctx.services.state.set_state(BeamlineStateEnum.DataCollection)
result = self._execute_scan(request)
if self.ctx.services.state is not None:
self.ctx.services.state.set_state(BeamlineStateEnum.SampleAlignment)
if sample is not None and sample.db_id is not None:
if self.ctx.services.screenshots is not None:
self.ctx.services.screenshots.save_to_db(
sample.db_id,
self.ctx.settings.preview_filename,
)
if self.ctx.services.events is not None:
self.ctx.services.events.send(sample.db_id, SampleEventType.COLLECTED)
if self.ctx.services.ingestion is not None:
self.ctx.services.ingestion.ingest_scan(
sample=sample,
result=result.result,
geom=self.ctx.sample_geometry,
beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom),
)
return result
@@ -4,17 +4,18 @@ import pytest
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.models import BoundingBoxModel, MLBoxModel, MLBoxType, ZoomModeEnum
from aare.daq.operations.common import runtime
from aare.daq.operations.face_detection.models import (
FaceDetectionContext,
FaceDetectionDependencies,
FaceDetectionResult,
FaceDetectionSettings,
)
from aare.daq.operations.face_detection.service import FaceDetectionService
from aare.daq.operations.common.runtime import (
DAQRuntimeState,
SampleProvider,
StatusProvider,
FaceDetectionProgressEmitter)
from aare.daq.operations.common.runtime import DAQRuntimeState
from aare.daq.operations.common.services import (
FaceDetectionProgressEmitter,
OperationServices,
)
class DummyGeometry:
def __init__(self):
@@ -61,22 +62,30 @@ def context():
status_provider=types.SimpleNamespace(status=None),
)
progress_emitter = FaceDetectionProgressEmitter(
reporter=types.SimpleNamespace(emit_progress=emit_progress)
services = OperationServices(
screenshots=types.SimpleNamespace(
save_to_db=lambda *args, **kwargs: None
),
face_detection_progress=FaceDetectionProgressEmitter(
reporter=types.SimpleNamespace(emit_progress=emit_progress)
),
)
ctx = FaceDetectionContext(
cfg=cfg,
devs=devs,
mlbox=types.SimpleNamespace(
predict=lambda **kwargs: types.SimpleNamespace(
box=None,
target_point=None,
focus=None,
deps=FaceDetectionDependencies(
cfg=cfg,
devs=devs,
mlbox=types.SimpleNamespace(
predict=lambda **kwargs: types.SimpleNamespace(
box=None,
target_point=None,
focus=None,
),
),
),
runtime=runtime_state,
progress_emitter=progress_emitter
services=services,
settings=FaceDetectionSettings(),
)
ctx._progress_events = progress_events
return ctx
@@ -86,7 +95,7 @@ def test_service_returns_empty_payload_when_no_boxes(monkeypatch, context, mock_
service = FaceDetectionService(context=context, logger=mock_logger)
monkeypatch.setattr(
context.mlbox,
context.deps.mlbox,
"predict",
lambda **kwargs: types.SimpleNamespace(
box=None,
@@ -103,7 +112,7 @@ def test_service_returns_empty_payload_when_no_boxes(monkeypatch, context, mock_
assert result.payload["samples"] == []
assert result.payload["height_fit"] == {}
assert result.payload["area_fit"] == {}
assert context.cfg.zoom_mode == ZoomModeEnum.User
assert context.deps.cfg.zoom_mode == ZoomModeEnum.User
assert len(context._progress_events) >= 1
assert context._progress_events[-1]["running"] is False
@@ -120,7 +129,7 @@ def test_service_prefers_face_boxes_when_ratio_is_high(monkeypatch, context, moc
]
)
monkeypatch.setattr(context.mlbox, "predict", lambda **kwargs: next(predictions))
monkeypatch.setattr(context.deps.mlbox, "predict", lambda **kwargs: next(predictions))
monkeypatch.setattr(
"aare.daq.operations.face_detection.service.fd.get_flat_face",
@@ -144,7 +153,7 @@ def test_service_prefers_face_boxes_when_ratio_is_high(monkeypatch, context, moc
assert result.payload["running"] is False
assert result.payload["height_fit"]["best_angle_deg"] == 50
assert result.payload["area_fit"]["best_angle_deg"] == 45
assert context.devs.aerotech_omega == 47
assert context.deps.devs.aerotech_omega == 47
assert [sample["angle"] for sample in result.payload["samples"]] == [0, 30, 90]
assert context._progress_events[-1]["running"] is False
@@ -162,7 +171,7 @@ def test_service_falls_back_to_loop_all_when_face_ratio_is_low(monkeypatch, cont
]
)
monkeypatch.setattr(context.mlbox, "predict", lambda **kwargs: next(predictions))
monkeypatch.setattr(context.deps.mlbox, "predict", lambda **kwargs: next(predictions))
captured = {}
@@ -185,7 +194,7 @@ def test_service_falls_back_to_loop_all_when_face_ratio_is_low(monkeypatch, cont
assert result.success is True
assert len(captured["boxes_used"]) == 4
assert sorted(captured["boxes_used"].keys()) == [30, 60, 90, 120]
assert context.devs.aerotech_omega == 60
assert context.deps.devs.aerotech_omega == 60
assert context._progress_events[-1]["running"] is False
@@ -196,16 +205,16 @@ def test_service_applies_centre_correction_when_target_is_far_from_beam(monkeypa
service._centre_correction(model, tolerance=0.2)
assert isinstance(context.devs.smargon_pos, SmargonCoordinate)
assert context.devs.smargon_pos.sh_mm.x == pytest.approx(1.0)
assert context.devs.smargon_pos.sh_mm.y == pytest.approx(1.8)
assert isinstance(context.deps.devs.smargon_pos, SmargonCoordinate)
assert context.deps.devs.smargon_pos.sh_mm.x == pytest.approx(1.0)
assert context.deps.devs.smargon_pos.sh_mm.y == pytest.approx(1.8)
def test_service_returns_failed_result_when_prediction_raises(monkeypatch, context, mock_logger):
service = FaceDetectionService(context=context, logger=mock_logger)
monkeypatch.setattr(
context.mlbox,
context.deps.mlbox,
"predict",
lambda **kwargs: (_ for _ in ()).throw(RuntimeError("prediction failed")),
)
@@ -216,5 +225,5 @@ def test_service_returns_failed_result_when_prediction_raises(monkeypatch, conte
assert isinstance(result.error, RuntimeError)
assert result.comment == "Face detection sequence failed"
assert result.payload["running"] is False
assert context.cfg.zoom_mode == ZoomModeEnum.User
assert context.deps.cfg.zoom_mode == ZoomModeEnum.User
assert context._progress_events[-1]["running"] is False
@@ -5,10 +5,16 @@ import pytest
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.models import MLBoxType, MLOutputModel
from aare.daq.mlbox import MLBoxPredictionsResult
from aare.daq.operations.common.runtime import DAQRuntimeState, TraceWriter, PredictionProvider, OperationServices
from aare.daq.operations.common.runtime import DAQRuntimeState
from aare.daq.operations.common.services import (
TraceWriter,
PredictionProvider,
OperationServices,
)
from aare.daq.operations.loop_centering.analyzer import LoopCenteringAnalyzer
from aare.daq.operations.loop_centering.models import (
LoopCenteringContext,
LoopCenteringDependencies,
LoopCenteringSettings,
)
@@ -43,29 +49,31 @@ def analyzer(mock_logger):
)
ctx = LoopCenteringContext(
cfg=types.SimpleNamespace(),
devs=types.SimpleNamespace(),
mlbox=DummyMlBox(),
settings=LoopCenteringSettings(),
deps=LoopCenteringDependencies(
cfg=types.SimpleNamespace(),
devs=types.SimpleNamespace(),
mlbox=DummyMlBox(),
),
runtime=runtime_state,
services=OperationServices(
screenshots=types.SimpleNamespace(
save_to_db=lambda *args, **kwargs: None
)
),
trace_writer=TraceWriter(
appender=types.SimpleNamespace(append_smargon_trace=lambda *args, **kwargs: None)
),
prediction_provider=PredictionProvider(
getter=types.SimpleNamespace(
get_predictions=lambda: MLBoxPredictionsResult(
predictions=None,
image=None,
target_point=None,
focus=None,
),
traces=TraceWriter(
appender=types.SimpleNamespace(append_smargon_trace=lambda *args, **kwargs: None)
),
predictions=PredictionProvider(
getter=types.SimpleNamespace(
get_predictions=lambda: MLBoxPredictionsResult(
predictions=None,
image=None,
target_point=None,
focus=None,
)
)
)
),
),
settings=LoopCenteringSettings(),
)
return LoopCenteringAnalyzer(context=ctx, logger=mock_logger)
@@ -142,7 +150,7 @@ def test_analyze_angle_with_crystal_prediction_returns_valid_target(analyzer):
boxes = MLOutputModel()
boxes.add_box(MLBoxType.CRYSTAL, (10, 20, 30, 40), conf=0.9)
analyzer.ctx.prediction_provider.getter.get_predictions = lambda: MLBoxPredictionsResult(
analyzer.ctx.services.predictions.getter.get_predictions = lambda: MLBoxPredictionsResult(
predictions=boxes,
image=None,
target_point=(25.0, 30.0),
@@ -3,10 +3,16 @@ import types
import pytest
from aare.common.models import LoopCenteringResult
from aare.daq.operations.common.runtime import DAQRuntimeState, OperationServices, TraceWriter, PredictionProvider
from aare.daq.operations.common.runtime import DAQRuntimeState
from aare.daq.operations.common.services import (
OperationServices,
TraceWriter,
PredictionProvider,
)
from aare.daq.operations.loop_centering.models import (
AngleAnalysis,
LoopCenteringContext,
LoopCenteringDependencies,
LoopCenteringSettings,
)
from aare.daq.operations.loop_centering.service import LoopCenteringService
@@ -31,24 +37,26 @@ def context():
)
ctx = LoopCenteringContext(
cfg=types.SimpleNamespace(),
devs=devs,
mlbox=types.SimpleNamespace(),
settings=LoopCenteringSettings(),
deps=LoopCenteringDependencies(
cfg=types.SimpleNamespace(),
devs=devs,
mlbox=types.SimpleNamespace(),
),
runtime=runtime_state,
services=OperationServices(
screenshots=types.SimpleNamespace(
save_to_db=lambda sample_id, filename, wait=0.0: screenshots.append((sample_id, filename, wait))
)
),
trace_writer=TraceWriter(
appender=types.SimpleNamespace(
append_smargon_trace=lambda **kwargs: traces.append(kwargs)
)
),
prediction_provider=PredictionProvider(
getter=types.SimpleNamespace(get_predictions=lambda: None)
),
traces=TraceWriter(
appender=types.SimpleNamespace(
append_smargon_trace=lambda **kwargs: traces.append(kwargs)
)
),
predictions=PredictionProvider(
getter=types.SimpleNamespace(get_predictions=lambda: None)
),
),
settings=LoopCenteringSettings(),
)
ctx._screenshots = screenshots
ctx._traces = traces
@@ -126,5 +134,5 @@ def test_service_uses_settings_values(monkeypatch, context, mock_logger):
result = service.run(sample_id=3)
assert context.devs.zoom == 250.0
assert context.deps.devs.zoom == 250.0
assert result.success is False
@@ -5,7 +5,7 @@ from aare.common.coordinate import AerotechCoordinate, Coordinate
from aare.common.exception_handler import CriticalTellException, MountingFailed
from aare.common.models import DewarAddress, SampleShortInfo
from aare.devices.tell_client import TellEventValueEnum
from aare.daq.operations.mounting.models import MountingContext, MountingResult
from aare.daq.operations.mounting.models import MountingContext, MountingResult, MountingDependencies, MountingSettings
from aare.daq.operations.mounting.service import MountingService
@@ -64,9 +64,13 @@ def _make_context(previous_sample=None):
)
return MountingContext(
cfg=cfg,
devs=devs,
mount_position=AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0),
deps=MountingDependencies(
cfg=cfg,
devs=devs,
),
settings=MountingSettings(
mount_position=AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0),
),
)
@@ -84,7 +88,7 @@ def test_execute_mount_success(mock_logger):
assert result.mounted_sample == target_sample
assert result.previous_sample == previous_sample
assert result.did_unmount_previous is True
assert ctx.cfg.current_sample == target_sample
assert ctx.deps.cfg.current_sample == target_sample
def test_execute_unmount_success(mock_logger):
@@ -99,13 +103,13 @@ def test_execute_unmount_success(mock_logger):
assert result.mounted_sample is None
assert result.previous_sample == previous_sample
assert result.did_unmount_previous is True
assert ctx.cfg.current_sample is None
assert ctx.deps.cfg.current_sample is None
def test_execute_mount_returns_failed_result_for_no_pin_in_gripper(mock_logger):
target_sample = _make_sample(2, "new")
ctx = _make_context()
ctx.devs.tell.mount = lambda **kwargs: TellEventValueEnum.NO_PIN_IN_GRIPPER
ctx.deps.devs.tell.mount = lambda **kwargs: TellEventValueEnum.NO_PIN_IN_GRIPPER
service = MountingService(context=ctx, logger=mock_logger)
@@ -119,7 +123,7 @@ def test_execute_mount_returns_failed_result_for_no_pin_in_gripper(mock_logger):
def test_execute_mount_returns_failed_result_for_unhandled_tell_response(mock_logger):
target_sample = _make_sample(2, "new")
ctx = _make_context()
ctx.devs.tell.mount = lambda **kwargs: TellEventValueEnum.UNKNOWN
ctx.deps.devs.tell.mount = lambda **kwargs: TellEventValueEnum.UNKNOWN
service = MountingService(context=ctx, logger=mock_logger)
@@ -136,14 +140,14 @@ def test_dry_unmounts_current_sample_before_drying(mock_logger):
unmount_calls = []
dry_calls = []
ctx.devs.tell.unmount = lambda **kwargs: unmount_calls.append(kwargs)
ctx.devs.tell.dry = lambda **kwargs: dry_calls.append(kwargs)
ctx.deps.devs.tell.unmount = lambda **kwargs: unmount_calls.append(kwargs)
ctx.deps.devs.tell.dry = lambda **kwargs: dry_calls.append(kwargs)
service = MountingService(context=ctx, logger=mock_logger)
service.dry(park=True, unmount=True)
assert len(unmount_calls) == 1
assert ctx.cfg.current_sample is None
assert ctx.deps.cfg.current_sample is None
assert dry_calls == [{"wait_cold": -1, "wait": True}]
@@ -152,10 +156,10 @@ def test_execute_mount_triggers_dry_on_third_consecutive_failure(mock_logger):
dry_calls = []
ctx = _make_context()
ctx.devs.tell.mount = lambda **kwargs: TellEventValueEnum.NO_PIN_IN_GRIPPER
ctx.devs.tell.dry = lambda **kwargs: dry_calls.append(kwargs)
ctx.cfg.increment_mount_failure_streak()
ctx.cfg.increment_mount_failure_streak()
ctx.deps.devs.tell.mount = lambda **kwargs: TellEventValueEnum.NO_PIN_IN_GRIPPER
ctx.deps.devs.tell.dry = lambda **kwargs: dry_calls.append(kwargs)
ctx.deps.cfg.increment_mount_failure_streak()
ctx.deps.cfg.increment_mount_failure_streak()
service = MountingService(context=ctx, logger=mock_logger)
@@ -172,12 +176,12 @@ def test_execute_mount_stops_automation_on_fifth_consecutive_failure(mock_logger
dry_calls = []
ctx = _make_context()
ctx.devs.tell.mount = lambda **kwargs: TellEventValueEnum.NO_PIN_IN_GRIPPER
ctx.devs.tell.dry = lambda **kwargs: dry_calls.append(kwargs)
ctx.cfg.increment_mount_failure_streak()
ctx.cfg.increment_mount_failure_streak()
ctx.cfg.increment_mount_failure_streak()
ctx.cfg.increment_mount_failure_streak()
ctx.deps.devs.tell.mount = lambda **kwargs: TellEventValueEnum.NO_PIN_IN_GRIPPER
ctx.deps.devs.tell.dry = lambda **kwargs: dry_calls.append(kwargs)
ctx.deps.cfg.increment_mount_failure_streak()
ctx.deps.cfg.increment_mount_failure_streak()
ctx.deps.cfg.increment_mount_failure_streak()
ctx.deps.cfg.increment_mount_failure_streak()
service = MountingService(context=ctx, logger=mock_logger)
@@ -194,15 +198,15 @@ def test_execute_mount_success_resets_failure_streak(mock_logger):
previous_sample = _make_sample(1, "old")
target_sample = _make_sample(2, "new")
ctx = _make_context(previous_sample=previous_sample)
ctx.cfg.increment_mount_failure_streak()
ctx.cfg.increment_mount_failure_streak()
ctx.deps.cfg.increment_mount_failure_streak()
ctx.deps.cfg.increment_mount_failure_streak()
service = MountingService(context=ctx, logger=mock_logger)
result = service.execute(target=target_sample)
assert result.success is True
assert ctx.cfg.get_mount_failure_streak() == 0
assert ctx.deps.cfg.get_mount_failure_streak() == 0
def test_execute_mount_critical_tell_error_does_not_increment_failure_streak(mock_logger):
@@ -212,7 +216,7 @@ def test_execute_mount_critical_tell_error_does_not_increment_failure_streak(moc
def mount(**kwargs):
raise CriticalTellException("critical tell problem")
ctx.devs.tell.mount = mount
ctx.deps.devs.tell.mount = mount
service = MountingService(context=ctx, logger=mock_logger)
@@ -220,4 +224,4 @@ def test_execute_mount_critical_tell_error_does_not_increment_failure_streak(moc
assert result.success is False
assert isinstance(result.error, CriticalTellException)
assert ctx.cfg.get_mount_failure_streak() == 0
assert ctx.deps.cfg.get_mount_failure_streak() == 0
+5 -6
View File
@@ -9,7 +9,7 @@ from aare.common.models import DewarAddress, SampleShortInfo, DAQOperation
from aare.common.tell_models import TellActivityEnum, TellPhaseEnum, TellStateModel
from aare.daq.config import ABR_POS_MOUNT, BeamlineStateEnum
from aare.daq.daq import AareDAQ
from aare.daq.operations.mounting.models import MountingResult
from aare.daq.operations.mounting.models import MountingResult, MountingContext
from aare.daq.operations.mounting.service import MountingService
from aare.daq.operations.screenshot.service import ScreenshotService
@@ -51,11 +51,10 @@ def bare_daq():
def test_create_mounting_service_builds_expected_context(mock_logger, bare_daq):
service = bare_daq._create_mounting_service()
assert isinstance(service, MountingService)
assert service.ctx.cfg is bare_daq._AareDAQ__cfg
assert service.ctx.devs is bare_daq._AareDAQ__devs
assert service.ctx.mount_position == ABR_POS_MOUNT
assert isinstance(service.ctx, MountingContext)
assert service.ctx.deps.cfg is bare_daq._AareDAQ__cfg
assert service.ctx.deps.devs is bare_daq._AareDAQ__devs
assert service.ctx.settings.mount_position == ABR_POS_MOUNT
def test_was_previous_sample_unmounted_since_prefers_tell_phase_confirmation():
previous_sample = _make_sample(1, "old")
+58 -23
View File
@@ -1,3 +1,5 @@
import types
import pytest
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -6,7 +8,13 @@ from jfjoch_client.exceptions import NotFoundException
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.raster_grid import RasterGridRequest
from aare.daq.operations.raster.models import RasterContext
from aare.daq.operations.common.runtime import DAQRuntimeState
from aare.daq.operations.common.services import OperationServices
from aare.daq.operations.raster.models import (
RasterContext,
RasterDependencies,
RasterSettings,
)
from aare.daq.operations.raster.service import RasterService
@@ -23,25 +31,52 @@ def make_request(n_x: int, n_y: int, cell_x: float = 0.01, cell_y: float = 0.02)
)
def make_service() -> RasterService:
runtime = SimpleNamespace(
sample=None,
sample_geometry=SimpleNamespace(),
status=SimpleNamespace(),
def _make_raster_context(*, jfjoch, aare, sample=None):
deps = RasterDependencies(
cfg=types.SimpleNamespace(
simulated_detector=False,
abr_meas_pos=types.SimpleNamespace(at_mm=Coordinate(x=0.0, y=0.0, z=0.0)),
get_beam_mark=lambda zoom: (0.0, 0.0),
),
devs=types.SimpleNamespace(
aerotech_omega=0.0,
zoom=100.0,
),
mlbox=types.SimpleNamespace(),
jfjoch=jfjoch,
aare=aare,
)
context = RasterContext(
cfg=MagicMock(),
devs=MagicMock(),
mlbox=MagicMock(),
jfjoch=MagicMock(),
aare=MagicMock(),
runtime=runtime,
state_controller=MagicMock(),
services=MagicMock(),
runtime = DAQRuntimeState(
sample_provider=types.SimpleNamespace(sample=sample),
sample_geometry_provider=types.SimpleNamespace(
sample_geometry=types.SimpleNamespace()
),
status_provider=types.SimpleNamespace(status=None),
)
services = OperationServices(
screenshots=types.SimpleNamespace(save_to_db=lambda *args, **kwargs: None),
)
settings = RasterSettings(
auto_raster_max_images=4500,
auto_raster_min_cell_size_mm=0.005,
auto_raster_skip_if_exceed_max_image_threshold=True,
)
return RasterContext(
deps=deps,
runtime=runtime,
services=services,
settings=settings,
)
def make_service() -> RasterService:
jfjoch = MagicMock()
aare = MagicMock()
context = _make_raster_context(jfjoch=jfjoch, aare=aare)
return RasterService(context=context, logger=MagicMock())
@@ -139,8 +174,8 @@ def test_upload_raster_diffraction_preview_skips_out_of_range_image_id():
request=request,
)
service.ctx.jfjoch.get_diffraction_image.assert_not_called()
service.ctx.aare.upload_jpg.assert_not_called()
service.ctx.deps.jfjoch.get_diffraction_image.assert_not_called()
service.ctx.deps.aare.upload_jpg.assert_not_called()
def test_upload_raster_diffraction_preview_ignores_not_found():
@@ -149,7 +184,7 @@ def test_upload_raster_diffraction_preview_ignores_not_found():
request = make_request(2, 2)
scan_result = MagicMock()
scan_result.images = [MagicMock(), MagicMock(), MagicMock(), MagicMock()]
service.ctx.jfjoch.get_diffraction_image.side_effect = NotFoundException()
service.ctx.deps.jfjoch.get_diffraction_image.side_effect = NotFoundException()
service._upload_raster_diffraction_preview(
sample_id=123,
@@ -159,8 +194,8 @@ def test_upload_raster_diffraction_preview_ignores_not_found():
request=request,
)
service.ctx.jfjoch.get_diffraction_image.assert_called_once_with(2)
service.ctx.aare.upload_jpg.assert_not_called()
service.ctx.deps.jfjoch.get_diffraction_image.assert_called_once_with(2)
service.ctx.deps.aare.upload_jpg.assert_not_called()
def test_upload_raster_diffraction_preview_uploads_when_present():
@@ -169,7 +204,7 @@ def test_upload_raster_diffraction_preview_uploads_when_present():
request = make_request(2, 2)
scan_result = MagicMock()
scan_result.images = [MagicMock(), MagicMock(), MagicMock(), MagicMock()]
service.ctx.jfjoch.get_diffraction_image.return_value = b"jpeg-bytes"
service.ctx.deps.jfjoch.get_diffraction_image.return_value = b"jpeg-bytes"
service._upload_raster_diffraction_preview(
sample_id=123,
@@ -179,5 +214,5 @@ def test_upload_raster_diffraction_preview_uploads_when_present():
request=request,
)
service.ctx.jfjoch.get_diffraction_image.assert_called_once_with(2)
service.ctx.aare.upload_jpg.assert_called_once_with(123, "preview", b"jpeg-bytes")
service.ctx.deps.jfjoch.get_diffraction_image.assert_called_once_with(2)
service.ctx.deps.aare.upload_jpg.assert_called_once_with(123, "preview", b"jpeg-bytes")
+3 -2
View File
@@ -1,7 +1,7 @@
import pytest
from unittest.mock import MagicMock, patch
from aare.daq.workflows import common_2rse, sa2se, sa2rse, sa2xtal_snapshot, dc2xtal_snapshot, xtal_snapshot2dc, xtal_snapshot2sa, dc2rse, se2sa, sa2dc, dc2sa, sa2xrf, sa2dh, dh2sa
from aare.daq.config import ABR_POS_MOUNT
from aare.daq.config import ABR_POS_MOUNT, ABR_OMEGA_MOUNT
from aare.common.models import StagePositionEnum
from aare.devices.area_detector import AutoEnum
from aare.devices.bec_worker import BeamlineState
@@ -99,7 +99,8 @@ def test_sa2se(mock_devs, mock_cfg):
mock_devs.samcam_auto.assert_any_call(AutoEnum.AUTO)
_assert_bec_moved(mock_devs, BeamlineState.MANUAL_SAMPLE_EXCHANGE)
assert mock_devs.aerotech_pos == mock_cfg.abr_meas_pos
assert mock_devs.aerotech_pos == ABR_POS_MOUNT
assert mock_devs.aerotech_omega == ABR_OMEGA_MOUNT
mock_devs.smargon_move_home.assert_called_once()
mock_devs.samcam_auto.assert_any_call(AutoEnum.ONCE)