diff --git a/src/simdb/cli/commands/simulation.py b/src/simdb/cli/commands/simulation.py index 59bd0bdd..d7327c2a 100644 --- a/src/simdb/cli/commands/simulation.py +++ b/src/simdb/cli/commands/simulation.py @@ -248,7 +248,7 @@ def simulation_push( schemas = api.get_validation_schemas() try: for schema in schemas: - Validator(schema).validate(simulation) + Validator(schema, config).validate(simulation) except ValidationError as err: raise click.ClickException(f"Simulation does not validate: {err}") from err @@ -471,7 +471,7 @@ def simulation_validate( click.echo("validating metadata ... ", nl=False) for schema in schemas: - Validator(schema).validate(simulation) + Validator(schema, config).validate(simulation) ids_list = [] for file in chain(simulation.inputs, simulation.outputs): diff --git a/src/simdb/cli/remote_api.py b/src/simdb/cli/remote_api.py index 99f87363..9b569b62 100644 --- a/src/simdb/cli/remote_api.py +++ b/src/simdb/cli/remote_api.py @@ -42,7 +42,7 @@ from .manifest import DataType if TYPE_CHECKING: - from simdb.database.models import File, Simulation, Watcher + from simdb.database.models import File, Watcher if TYPE_CHECKING or "sphinx" in sys.modules: # Only importing these for type checking and documentation generation in order to diff --git a/src/simdb/database/models/simulation.py b/src/simdb/database/models/simulation.py index ede09fdf..57db96f1 100644 --- a/src/simdb/database/models/simulation.py +++ b/src/simdb/database/models/simulation.py @@ -163,7 +163,7 @@ def meta(self) -> List[MetaDataWrapper]: Returns a list of MetaDataWrapper objects from the JSON metadata. """ meta_dict = self._get_metadata_dict() - return [MetaDataWrapper(k, v) for k, v in meta_dict.items()] + return list(itertools.starmap(MetaDataWrapper, meta_dict.items())) def _get_metadata_dict(self) -> Dict[str, Any]: if self._metadata is None: diff --git a/src/simdb/remote/apis/v1/simulations.py b/src/simdb/remote/apis/v1/simulations.py index 3a56c9cb..c22acfa4 100644 --- a/src/simdb/remote/apis/v1/simulations.py +++ b/src/simdb/remote/apis/v1/simulations.py @@ -52,7 +52,7 @@ def _update_simulation_status( def _validate(simulation, user) -> Dict: schema = Validator.validation_schema() try: - Validator(schema).validate(simulation) + Validator(schema, current_app.simdb_config).validate(simulation) _update_simulation_status(simulation, models_sim.Simulation.Status.PASSED, user) return { "passed": True, diff --git a/src/simdb/remote/apis/v1_1/simulations.py b/src/simdb/remote/apis/v1_1/simulations.py index 03c956b9..105b92e9 100644 --- a/src/simdb/remote/apis/v1_1/simulations.py +++ b/src/simdb/remote/apis/v1_1/simulations.py @@ -53,7 +53,7 @@ def _validate(simulation, user) -> Dict: schemas = Validator.validation_schemas(current_app.simdb_config, simulation) try: for schema in schemas: - Validator(schema).validate(simulation) + Validator(schema, current_app.simdb_config).validate(simulation) _update_simulation_status( simulation, models_sim.Simulation.Status.PASSED, user ) diff --git a/src/simdb/remote/apis/v1_2/simulations.py b/src/simdb/remote/apis/v1_2/simulations.py index a8594490..b55719af 100644 --- a/src/simdb/remote/apis/v1_2/simulations.py +++ b/src/simdb/remote/apis/v1_2/simulations.py @@ -82,7 +82,7 @@ def _validate(simulation, user) -> ValidationResult: schemas = Validator.validation_schemas(current_app.simdb_config, simulation) try: for schema in schemas: - Validator(schema).validate(simulation) + Validator(schema, current_app.simdb_config).validate(simulation) _update_simulation_status( simulation, models_sim.Simulation.Status.PASSED, user ) diff --git a/src/simdb/validation/file/ids_validator.py b/src/simdb/validation/file/ids_validator.py index 07452a41..9a7714bf 100644 --- a/src/simdb/validation/file/ids_validator.py +++ b/src/simdb/validation/file/ids_validator.py @@ -43,7 +43,7 @@ def configure(self, arguments: dict): .split(",") ] - ### Define logic for rule_filter + # Define logic for rule_filter list_of_filter_names = ( arguments.get("rule_filter_name", "").strip('"').split(",") ) diff --git a/src/simdb/validation/validator.py b/src/simdb/validation/validator.py index f9a3bb2b..3decf052 100644 --- a/src/simdb/validation/validator.py +++ b/src/simdb/validation/validator.py @@ -1,5 +1,6 @@ import re import warnings +from importlib import import_module from pathlib import Path from typing import Any, Dict, List, Optional, Union, cast @@ -200,9 +201,45 @@ def validation_schemas( return schemas - def __init__(self, schema: Dict): + def _custom_validation_ext(self, config: Config): + + module_path = config.get_option("validation.custom_validator", default=None) + if module_path is None: + return CustomValidator + + if not isinstance(module_path, str): + raise TypeError( + "Expected 'custom_validator config value' to be a string, " + f"got {type(module_path).__name__}" + ) + + if "." not in module_path: + raise ValueError( + f"Invalid validator path '{module_path}'." + "Expected format: 'package.module.ClassName'" + ) + + module_name, class_name = module_path.rsplit(".", 1) + + try: + module = import_module(module_name) + except ModuleNotFoundError as err: + raise ImportError( + f"Unable to import module '{module_name}': {err}. " + "Please ensure the necessary validation package is installed" + ) from err + try: + validation_cls = getattr(module, class_name) + except AttributeError as err: + raise AttributeError( + f"Module '{module_name}' does not have classor attribute '{class_name}'" + ) from err + return validation_cls + + def __init__(self, schema: Dict, config: Config): try: - self._validator = CustomValidator(schema) + validation_cls = self._custom_validation_ext(config) + self._validator = validation_cls(schema) self._validator.allow_unknown = True except cerberus.SchemaError as err: raise LoadError("Failed to parse validation schema") from err diff --git a/tests/remote/api/conftest.py b/tests/remote/api/conftest.py index 4310350c..da2c7b17 100644 --- a/tests/remote/api/conftest.py +++ b/tests/remote/api/conftest.py @@ -1,5 +1,4 @@ import base64 -import importlib import importlib.util import os import shutil diff --git a/tests/remote/test_authentication.py b/tests/remote/test_authentication.py index d62f56a3..b5522092 100644 --- a/tests/remote/test_authentication.py +++ b/tests/remote/test_authentication.py @@ -1,4 +1,3 @@ -import importlib import importlib.util from typing import TYPE_CHECKING, ClassVar, cast from unittest import mock