From 83ceb82ea854dba1b4189f2222c801fcf2dc2504 Mon Sep 17 00:00:00 2001 From: Julian Stirling Date: Tue, 2 Jun 2026 12:57:46 +0100 Subject: [PATCH] Use a path object that supports setting a Thing for the data directory. This allows us to ensure that arbitrary paths are not selected via the API. --- .../things/__init__.py | 87 +++++++++++++++++++ .../things/autofocus.py | 5 +- .../things/camera/__init__.py | 44 +++++----- .../things/scan_workflows.py | 53 ++++++----- tests/unit_tests/test_stack.py | 2 + 5 files changed, 146 insertions(+), 45 deletions(-) diff --git a/src/openflexure_microscope_server/things/__init__.py b/src/openflexure_microscope_server/things/__init__.py index b0c5768f..c05a34ee 100644 --- a/src/openflexure_microscope_server/things/__init__.py +++ b/src/openflexure_microscope_server/things/__init__.py @@ -5,9 +5,13 @@ with other Things and including them in the LabThings-FastAPI config file. """ import os +from pathlib import PurePath +from tempfile import TemporaryDirectory from types import TracebackType from typing import Optional, Self +from pydantic import PrivateAttr, RootModel, model_validator + import labthings_fastapi as lt @@ -59,3 +63,86 @@ class OFMThing(lt.Thing): "No data directory set. Has the LabThings server been started?" ) return self._data_dir + + def create_data_path(self, path: str) -> "RelativeDataPath": + """Create a ``RelativeDataPath`` object with this Thing set as the saving Thing. + + :param path: The relative path within the data directory of this Thing's data + dir that the data shudl be saved. + + :return: A ``RelativeDataPath`` object with the saving Thing already set. + """ + rel_data_path = RelativeDataPath(path) + rel_data_path.set_saving_thing(self) + return rel_data_path + + +class RelativeDataPath(RootModel[str]): + """A relative path that is validated, and can have a Thing assigned to it. + + Use ``set_saving_thing`` or ``set_saving_thing_if_unset`` to set the Thing whose + data directory will be used for the final save. + + Use the ``abs_data_path`` property to get the final path for saving. + """ + + _saving_thing: Optional[OFMThing | TemporaryDirectory] = PrivateAttr(default=None) + + @model_validator(mode="before") + @classmethod + def validate_relative_path(cls, value: str) -> str: + """Validate the relative path is relative and has no parent dir references.""" + p = PurePath(value) + + if p.is_absolute(): + raise ValueError("Absolute paths are not allowed") + + if ".." in p.parts: + raise ValueError("Parent directory references are not allowed") + + return os.path.normpath(value) + + @property + def save_location_set(self) -> bool: + """Return True if the saving thing is set.""" + return self._saving_thing is not None + + def set_saving_thing(self, thing: OFMThing) -> None: + """Set the Thing that is saving the data. This will set the data directory.""" + if self.save_location_set: + raise RuntimeError("The saving Thing for the relative path is already set") + self._saving_thing = thing + + def set_saving_thing_if_unset(self, thing: OFMThing) -> None: + """Set the Thing that is saving the data if it is not already set. + + Use this in an action to set the Thing for paths set via the API. + """ + if not self.save_location_set: + self._saving_thing = thing + + def save_to_tempdir(self) -> TemporaryDirectory: + """Use a temporary directory to save raher than an ``OFMThing``. + + :returns: the ``TemporaryDirectory`` object. + """ + if self.save_location_set: + raise RuntimeError("The saving Thing for the relative path is already set") + self._saving_thing = TemporaryDirectory() + return self._saving_thing + + def join(self, sub_path: str) -> "RelativeDataPath": + """Join a path to the end of this path. + + :return: A new ``RelativeDataPath`` object with the path appended. + """ + return RelativeDataPath(os.path.join(self.root, sub_path)) + + @property + def abs_data_path(self) -> str: + """The absolute data directory to save to.""" + if self.save_location_set: + raise RuntimeError("The saving Thing for the relative path was never set") + if isinstance(self._saving_thing, TemporaryDirectory): + return os.path.join(self._saving_thing.name, self.root) + return os.path.join(self._saving_thing.data_dir, self.root) diff --git a/src/openflexure_microscope_server/things/autofocus.py b/src/openflexure_microscope_server/things/autofocus.py index 3d1f599c..bff91e81 100644 --- a/src/openflexure_microscope_server/things/autofocus.py +++ b/src/openflexure_microscope_server/things/autofocus.py @@ -9,7 +9,6 @@ See repository root for licensing information. import enum import logging -import os import time from dataclasses import dataclass from types import TracebackType @@ -737,7 +736,7 @@ class AutofocusThing(lt.Thing): # Loop through the range, saving each capture to disk for capture in captures[slice_to_save]: - path = os.path.join(capture_parameters.images_dir, capture.filename) + path = capture_parameters.images_dir.join(capture.filename) self._cam.save_from_memory(path=path, buffer_id=capture.buffer_id) self._cam.clear_buffers() @@ -956,7 +955,7 @@ class AutofocusThing(lt.Thing): # Save all captures for capture in captures: - path = os.path.join(capture_parameters.images_dir, capture.filename) + path = capture_parameters.images_dir.join(capture.filename) self._cam.save_from_memory(path=path, buffer_id=capture.buffer_id) self._cam.clear_buffers() diff --git a/src/openflexure_microscope_server/things/camera/__init__.py b/src/openflexure_microscope_server/things/camera/__init__.py index 9f00a1f2..6da91133 100644 --- a/src/openflexure_microscope_server/things/camera/__init__.py +++ b/src/openflexure_microscope_server/things/camera/__init__.py @@ -11,23 +11,22 @@ from __future__ import annotations import io import json import os -import tempfile import time from abc import ABC, abstractmethod from copy import deepcopy from datetime import datetime from types import TracebackType -from typing import Annotated, Any, Literal, Mapping, Optional, Self +from typing import Any, Literal, Mapping, Optional, Self import numpy as np import piexif from PIL import Image -from pydantic import BaseModel, Field +from pydantic import BaseModel import labthings_fastapi as lt from labthings_fastapi.types.numpy import NDArray -from openflexure_microscope_server.things import OFMThing +from openflexure_microscope_server.things import OFMThing, RelativeDataPath from openflexure_microscope_server.things.background_detect import ( BackgroundDetectAlgorithm, ) @@ -62,15 +61,11 @@ class CaptureError(RuntimeError): """An error trying to capture from a CameraThing.""" -PositiveInt = Annotated[int, Field(ge=1)] -NonEmptyString = Annotated[str, Field(min_length=1)] - - class CaptureParams(BaseModel): """A class for capturing at least a single image.""" - images_dir: NonEmptyString - save_resolution: tuple[PositiveInt, PositiveInt] + images_dir: RelativeDataPath + capture_mode: str class NoImageInMemoryError(RuntimeError): @@ -576,7 +571,7 @@ class BaseCamera(OFMThing, ABC): the specific camera being used. :param capture_mode: The mode to use, must be one of ``capture_modes``. - :param image_format: The image fromat to use, must be one of + :param image_format: The image format to use, must be one of ``supported_image_formats`` :param retain_image: (Default True) True to save image to the microscope, False to only save temporarily for transfer. @@ -586,9 +581,12 @@ class BaseCamera(OFMThing, ABC): format_info = self.supported_image_formats[image_format] fname = datetime.now().strftime("%Y-%m-%d-%H%M%S") + format_info.extension - tmpdir = None if retain_image else tempfile.TemporaryDirectory() - dir_path = self.data_dir if tmpdir is None else tmpdir.name - path = os.path.join(dir_path, fname) + path = RelativeDataPath(fname) + tmpdir = None + if retain_image: + path.set_saving_thing(self) + else: + tmpdir = path.save_to_tempdir() self.capture_and_save_to_path(path, capture_mode) @@ -601,7 +599,7 @@ class BaseCamera(OFMThing, ABC): def capture_and_save_to_path( self, - path: str, + path: RelativeDataPath, capture_mode: str = "standard", ) -> None: """Capture an image and save it to disk. @@ -654,7 +652,7 @@ class BaseCamera(OFMThing, ABC): def save_from_memory( self, - path: str, + path: RelativeDataPath, buffer_id: Optional[int] = None, ) -> None: """Save an image that has been captured to memory. @@ -674,12 +672,16 @@ class BaseCamera(OFMThing, ABC): mode_info = self.capture_modes[mode] save_resolution = mode_info.save_resolution + path.set_saving_thing_if_unset(self) + resolved_path = path.abs_data_path + if save_resolution is not None and image.size != save_resolution: image = image.resize(save_resolution, Image.Resampling.BOX) try: save_kwargs: dict[str, Any] = {} # TODO: Test that the save_kwargs are called as expected for different formats. - if path.lower().endswith(BASE_IMAGE_FORMATS["jpeg"].supported_extensions): + jpeg_exts = BASE_IMAGE_FORMATS["jpeg"].supported_extensions + if resolved_path.lower().endswith(jpeg_exts): # Per PIL documentation, # (https://pillow.readthedocs.io/en/stable/handbook/image-file-formats.html#jpeg) # there are two factors when saving a JPEG. Subsampling affects the colour, @@ -688,15 +690,15 @@ class BaseCamera(OFMThing, ABC): # quality = 95 is the maximum recommended - above this, JPEG compression is # disabled, file size increases and quality is barely or not affected save_kwargs = {"quality": 95, "subsampling": 0} - image.save(path, **save_kwargs) + image.save(resolved_path, **save_kwargs) try: - self._add_metadata_to_capture(path, dict(metadata)) + self._add_metadata_to_capture(resolved_path, dict(metadata)) except Exception: # We need to capture any exception as there are many reasons metadata # might not be added. We warn rather than log the error. - self.logger.exception(f"Failed to add metadata to {path}") + self.logger.exception(f"Failed to add metadata to {resolved_path}") except Exception as e: - raise IOError(f"An error occurred while saving {path}") from e + raise IOError(f"An error occurred while saving {resolved_path}") from e @abstractmethod def _capture_image(self, capture_mode: str = "standard") -> Image.Image: diff --git a/src/openflexure_microscope_server/things/scan_workflows.py b/src/openflexure_microscope_server/things/scan_workflows.py index 4eeb419b..ec3f20a2 100644 --- a/src/openflexure_microscope_server/things/scan_workflows.py +++ b/src/openflexure_microscope_server/things/scan_workflows.py @@ -4,7 +4,6 @@ This module contains the base ``ScanWorkflow`` class that all workflows should s as well as specific workflows. """ -import os from typing import ( Generic, Literal, @@ -27,6 +26,7 @@ from openflexure_microscope_server.stitching import ( TARGET_STITCHING_DIMENSION, StitchingSettings, ) +from openflexure_microscope_server.things import RelativeDataPath from openflexure_microscope_server.things.autofocus import ( MAX_TEST_IMAGE_COUNT, MIN_TEST_IMAGE_COUNT, @@ -55,6 +55,10 @@ from openflexure_microscope_server.ui import ( SettingModelType = TypeVar("SettingModelType", bound=BaseModel) +class WorkflowStartError(lt.exceptions.InvocationError): + """The scan workflow cannot start, as the requested configuration is invalid.""" + + class ScanWorkflow(Generic[SettingModelType], lt.Thing): """A base class for all Scanworkflows. @@ -75,7 +79,7 @@ class ScanWorkflow(Generic[SettingModelType], lt.Thing): _planner_cls: type[ScanPlanner] # All workflows set a save resolution - save_resolution: tuple[int, int] = lt.setting(default=(1640, 1232)) + capture_mode: str = lt.setting(default="standard") """A tuple of the image resolution to capture.""" # CSM may not be set, and isn't required for a workflow. Allow for it to exist or be None @@ -85,15 +89,19 @@ class ScanWorkflow(Generic[SettingModelType], lt.Thing): _stage: BaseStage = lt.thing_slot() _autofocus: AutofocusThing = lt.thing_slot() - def check_before_start(self, scan_name: str) -> None: + # The noqa statement is because scan_name is unused but is needed for equivalence + # with other workflows that may want to validate the scan name. + def check_before_start(self, scan_name: str) -> None: # noqa: ARG002 """Check before the scan starts. Throw an error if the scan shouldn't start. The scan_name is passed to this function to enable workflows to validate the scan name if needed. """ - raise NotImplementedError( - "Each specific ScanWorkflow must implement a check_before_start." - ) + if self.capture_mode not in self._cam.capture_modes: + cam_name = type(self._cam).__name__ + raise WorkflowStartError( + f"{cam_name} has no capure mode {self.capture_mode}" + ) @lt.property def ready(self) -> bool: @@ -148,14 +156,14 @@ class ScanWorkflow(Generic[SettingModelType], lt.Thing): self, xyz_pos: tuple[int, int, int], dz: int, - images_dir: str, - save_resolution: tuple[int, int], + images_dir: RelativeDataPath, + capture_mode: str, ) -> tuple[bool, Optional[int]]: """Autofocus and then capture, this can be used as an acquisition routine. :param dz: The dz for autofocus. - :param images_dir: The path to the directory for saving images.. - :param save_resolution: The resolution to save images at. + :param images_dir: The path to the directory for saving images. + :param capture mode: The name of the camera capture mode. :return: A tuple ready to pass out of acquisition routine. In this method, image is always taken, so first return is True. @@ -165,8 +173,8 @@ class ScanWorkflow(Generic[SettingModelType], lt.Thing): focus_height = self._stage.get_xyz_position()[2] filename = f"img_{xyz_pos[0]}_{xyz_pos[1]}_{focus_height}.jpeg" self._cam.capture_and_save_to_path( - path=os.path.join(images_dir, filename), - capture_mode="standard", + path=images_dir.join(filename), + capture_mode=capture_mode, ) return True, focus_height @@ -217,16 +225,15 @@ class RectGridWorkflow( must be above this. 3000 is a sensible limit for 20x objectives. """ - # The noqa statement is because scan_name is unused but is needed for equivalence - # with other workflows that may want to validate the scan name. - def check_before_start(self, scan_name: str) -> None: # noqa: ARG002 + def check_before_start(self, scan_name: str) -> None: """Before starting a scan, check that camera-stage-mapping is set. Raise error if: - camera stage mapping is not set """ + super().check_before_start(scan_name) if self._csm.calibration_required: - raise RuntimeError("Camera Stage Mapping is not calibrated.") + raise WorkflowStartError("Camera Stage Mapping is not calibrated.") def _calc_displacement_from_overlap(self, overlap: float) -> tuple[int, int]: """Use camera stage mapping to calculate x and y displacement from given overlap. @@ -265,6 +272,7 @@ class RectGridWorkflow( """Return a stitching settings model based on current settings.""" # Use the save resolution and target stitch resolution to choose a unit fraction, # which makes correlating faster + # TODO Calculate this using a camera method. width, height = self.save_resolution # Target area in pixels target_area = TARGET_STITCHING_DIMENSION**2 @@ -302,8 +310,9 @@ class RectGridWorkflow( "overlap": self.overlap, "dx": dx, "dy": dy, + # TODO set images dir correctly as a RelDataPath "capture_params": CaptureParams( - images_dir=images_dir, save_resolution=self.save_resolution + images_dir=images_dir, capture_mode=self.capture_mode ), "autofocus_params": AutofocusParams(dz=self.autofocus_dz), } @@ -510,20 +519,22 @@ class HistoScanWorkflow(RectGridWorkflow[HistoScanSettingsModel], SmartStackMixi # The noqa statement is because scan_name is unused but is needed for equivalence # with other workflows that may want to validate the scan name. def check_before_start(self, scan_name: str) -> None: # noqa: ARG002 - """Before starting a scan, check that background and camera-stage-mapping are set. + """Before starting a scan, check that background and CSM are set. Raise error if: - background is to be skipped but is not set - camera stage mapping is not set - Raise warning if not using background detect that scan will go on until max steps reached + Raise warning if not using background detect that scan will go on until max + steps reached. """ + super().check_before_start(scan_name) if self._csm.calibration_required: - raise RuntimeError("Camera Stage Mapping is not calibrated.") + raise WorkflowStartError("Camera Stage Mapping is not calibrated.") if self.skip_background: if not self._background_detector.ready: - raise RuntimeError( + raise WorkflowStartError( "Background is not set: you need to calibrate background detection." ) else: diff --git a/tests/unit_tests/test_stack.py b/tests/unit_tests/test_stack.py index ced295e6..1c965d2e 100644 --- a/tests/unit_tests/test_stack.py +++ b/tests/unit_tests/test_stack.py @@ -888,6 +888,7 @@ def test_invalid_stack_settling_raises(): def test_invalid_capture_dir_raises(bad_path, match_err): """Test basic stack raises expected error for bad image dir paths.""" with pytest.raises(ValueError, match=match_err): + # TODO set images dir correctly as a RelDataPath CaptureParams(images_dir=bad_path, save_resolution=(20, 20)) @@ -909,4 +910,5 @@ def test_invalid_capture_dir_raises(bad_path, match_err): def test_invalid_capture_res_raises(bad_res, match_err): """Test basic stack raises expected error for invalid save resolutions.""" with pytest.raises(ValueError, match=match_err): + # TODO set images dir correctly as a RelDataPath CaptureParams(images_dir="dummy", save_resolution=bad_res)