operations refactor: standardised operations, added rotation operation, updated tests.
This commit is contained in:
+215
-93
@@ -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(
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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})"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user