diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index f71d3381..7c6781bd 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -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( diff --git a/src/aare/daq/operations/common/models.py b/src/aare/daq/operations/common/models.py new file mode 100644 index 00000000..22d1d305 --- /dev/null +++ b/src/aare/daq/operations/common/models.py @@ -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 \ No newline at end of file diff --git a/src/aare/daq/operations/common/runtime.py b/src/aare/daq/operations/common/runtime.py index aa2aee38..1f911b42 100644 --- a/src/aare/daq/operations/common/runtime.py +++ b/src/aare/daq/operations/common/runtime.py @@ -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 \ No newline at end of file + return self.status_provider.status \ No newline at end of file diff --git a/src/aare/daq/operations/common/services.py b/src/aare/daq/operations/common/services.py new file mode 100644 index 00000000..701e32e7 --- /dev/null +++ b/src/aare/daq/operations/common/services.py @@ -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 \ No newline at end of file diff --git a/src/aare/daq/operations/face_detection/models.py b/src/aare/daq/operations/face_detection/models.py index ec0cfbf2..a837c492 100644 --- a/src/aare/daq/operations/face_detection/models.py +++ b/src/aare/daq/operations/face_detection/models.py @@ -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 diff --git a/src/aare/daq/operations/face_detection/service.py b/src/aare/daq/operations/face_detection/service.py index e85fcde3..dcc75d5b 100644 --- a/src/aare/daq/operations/face_detection/service.py +++ b/src/aare/daq/operations/face_detection/service.py @@ -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 \ No newline at end of file + self.ctx.deps.cfg.zoom_mode = ZoomModeEnum.User \ No newline at end of file diff --git a/src/aare/daq/operations/loop_centering/analyzer.py b/src/aare/daq/operations/loop_centering/analyzer.py index 50060bdc..66d6eada 100644 --- a/src/aare/daq/operations/loop_centering/analyzer.py +++ b/src/aare/daq/operations/loop_centering/analyzer.py @@ -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}", diff --git a/src/aare/daq/operations/loop_centering/models.py b/src/aare/daq/operations/loop_centering/models.py index 4313b81a..dc2d8f51 100644 --- a/src/aare/daq/operations/loop_centering/models.py +++ b/src/aare/daq/operations/loop_centering/models.py @@ -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 \ No newline at end of file +class LoopCenteringContext( + BaseOperationContext[LoopCenteringDependencies, LoopCenteringSettings] +): + pass \ No newline at end of file diff --git a/src/aare/daq/operations/loop_centering/service.py b/src/aare/daq/operations/loop_centering/service.py index 0fe50966..c4d0f3a7 100644 --- a/src/aare/daq/operations/loop_centering/service.py +++ b/src/aare/daq/operations/loop_centering/service.py @@ -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})" diff --git a/src/aare/daq/operations/mounting/models.py b/src/aare/daq/operations/mounting/models.py index a29783df..18fc8ad9 100644 --- a/src/aare/daq/operations/mounting/models.py +++ b/src/aare/daq/operations/mounting/models.py @@ -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 diff --git a/src/aare/daq/operations/mounting/service.py b/src/aare/daq/operations/mounting/service.py index fa33d6cc..3dae5552 100644 --- a/src/aare/daq/operations/mounting/service.py +++ b/src/aare/daq/operations/mounting/service.py @@ -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, diff --git a/src/aare/daq/operations/raster/models.py b/src/aare/daq/operations/raster/models.py index 6aa9488e..389c26e2 100644 --- a/src/aare/daq/operations/raster/models.py +++ b/src/aare/daq/operations/raster/models.py @@ -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 diff --git a/src/aare/daq/operations/raster/service.py b/src/aare/daq/operations/raster/service.py index 3e83a929..3573f926 100644 --- a/src/aare/daq/operations/raster/service.py +++ b/src/aare/daq/operations/raster/service.py @@ -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", diff --git a/src/aare/daq/operations/rotation/models.py b/src/aare/daq/operations/rotation/models.py index e69de29b..42f4bc6b 100644 --- a/src/aare/daq/operations/rotation/models.py +++ b/src/aare/daq/operations/rotation/models.py @@ -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 \ No newline at end of file diff --git a/src/aare/daq/operations/rotation/service.py b/src/aare/daq/operations/rotation/service.py index e69de29b..64f540ed 100644 --- a/src/aare/daq/operations/rotation/service.py +++ b/src/aare/daq/operations/rotation/service.py @@ -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 \ No newline at end of file diff --git a/tests/unit/daq/operations/face_detection/test_face_detection_service.py b/tests/unit/daq/operations/face_detection/test_face_detection_service.py index 8b3babca..b06c012c 100644 --- a/tests/unit/daq/operations/face_detection/test_face_detection_service.py +++ b/tests/unit/daq/operations/face_detection/test_face_detection_service.py @@ -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 \ No newline at end of file diff --git a/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py b/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py index 93d8f9aa..76ac5684 100644 --- a/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py +++ b/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py @@ -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), diff --git a/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py b/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py index cd7664d4..41588c87 100644 --- a/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py +++ b/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py @@ -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 \ No newline at end of file diff --git a/tests/unit/daq/operations/mounting/test_mounting_service.py b/tests/unit/daq/operations/mounting/test_mounting_service.py index 6419e545..5f981d88 100644 --- a/tests/unit/daq/operations/mounting/test_mounting_service.py +++ b/tests/unit/daq/operations/mounting/test_mounting_service.py @@ -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 \ No newline at end of file diff --git a/tests/unit/daq/test_mount.py b/tests/unit/daq/test_mount.py index 4ab09cfd..8737705a 100644 --- a/tests/unit/daq/test_mount.py +++ b/tests/unit/daq/test_mount.py @@ -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") diff --git a/tests/unit/daq/test_raster_logic.py b/tests/unit/daq/test_raster_logic.py index 88916f44..900a3ad8 100644 --- a/tests/unit/daq/test_raster_logic.py +++ b/tests/unit/daq/test_raster_logic.py @@ -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") \ No newline at end of file + 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") \ No newline at end of file diff --git a/tests/unit/daq/test_workflows.py b/tests/unit/daq/test_workflows.py index 067fb655..142667d6 100644 --- a/tests/unit/daq/test_workflows.py +++ b/tests/unit/daq/test_workflows.py @@ -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)