diff --git a/src/openflexure_microscope_server/server/__init__.py b/src/openflexure_microscope_server/server/__init__.py index 1f14fbd3..93a675d8 100644 --- a/src/openflexure_microscope_server/server/__init__.py +++ b/src/openflexure_microscope_server/server/__init__.py @@ -6,10 +6,13 @@ from typing import Optional, Callable, Any from functools import wraps from copy import copy import logging +from argparse import Namespace import labthings_fastapi as lt import uvicorn from uvicorn.main import Server + +from openflexure_microscope_server.utilities import load_patched_config from .serve_static_files import add_static_files from .legacy_api import add_v2_endpoints from ..logging import configure_logging, retrieve_log, retrieve_log_from_file @@ -81,7 +84,7 @@ def serve_from_cli(argv: Optional[list[str]] = None) -> None: config = None server = None try: - config = lt.cli.config_from_args(args) + config = _full_config_from_args(args) log_folder = config.get("log_folder", "./openflexure/logs") scans_folder = _get_scans_dir(config) server = lt.cli.server_from_config(config) @@ -128,3 +131,16 @@ def serve_from_cli(argv: Optional[list[str]] = None) -> None: ) else: raise e + + +def _full_config_from_args(args: Namespace) -> dict: + """Load configuration from LabThings args allowing patching. + + This provides similar functionarlity to lt.cli.config_from_args except allows the + configuration file to specify a base config, and optionally patches. + """ + # If no config file specified let LabThings handle it. + if not args.config: + return lt.cli.config_from_args(args) + + return load_patched_config(args.config) diff --git a/src/openflexure_microscope_server/utilities.py b/src/openflexure_microscope_server/utilities.py index 8dfee823..d387c2ed 100644 --- a/src/openflexure_microscope_server/utilities.py +++ b/src/openflexure_microscope_server/utilities.py @@ -20,6 +20,7 @@ import logging from importlib.metadata import version import tomllib from functools import wraps +import json from pydantic import BaseModel import numpy as np @@ -430,3 +431,64 @@ def merge_patch( # Else, use the patch value result[key] = patch_val return result + + +def load_patched_config(config_path: str) -> dict: + """Load a json configuration file, if it a patch, return the full config. + + Load a json file. If it does not contains the "base_config_file" key then return + the data as read. + + If the configuration specifies a "base_config_file" but no "patch" then return the + data from reading base_config_file as json. + + If the configuration specifies a "base_config_file" and "patch" then apply the + patch to the data in base_config_file and return the result. The patching method + is merge_patch() + + :param config_path: The path of the configuration file. + :return: The contents of the configuration file after loading and patching the + specified base configuration file if applicable. + + """ + try: + with open(config_path, "r", encoding="utf-8") as f: + config = json.load(f) + except IOError as e: + raise type(e)(f"Couldn't load configuration file {config_path}") from e + + # If base_config_file not specified then this is a normal LabThings config file + # Just return it. + if "base_config_file" not in config: + return config + + config_dir = os.path.dirname(os.path.abspath(config_path)) + base_conf_path = config["base_config_file"] + + base_conf_path = resolve_path_from_dir(base_conf_path, config_dir) + + try: + with open(base_conf_path, "r", encoding="utf-8") as f: + base_config = json.load(f) + except IOError as e: + raise type(e)(f"Couldn't load configuration file {base_conf_path}") from e + + if "patch" in config: + return merge_patch(base_config, config["patch"], enforce_dict=True) + + return base_config + + +def resolve_path_from_dir(path: str, directory: str) -> str: + """Convert a relative or abs path specified in one dir to the working dir. + + This also expands any environment variables. + + :param path: The specified path (could be absolute or relative) + :param directory: The directory from which the path was specified + :return: The normalised path. + """ + path = os.path.expandvars(os.path.expanduser(path)) + if not os.path.isabs(path): + path = os.path.join(directory, path) + return os.path.normpath(path)