diff --git a/ccflow/__init__.py b/ccflow/__init__.py index 04c5a94f..2e33fb00 100644 --- a/ccflow/__init__.py +++ b/ccflow/__init__.py @@ -5,6 +5,7 @@ # which, in turn, import `ccflow`). from .exttypes import * # noqa: I001 +from . import config as config from .arrow import * from .base import * from .compose import * diff --git a/ccflow/base.py b/ccflow/base.py index e59a01a7..e6509e38 100644 --- a/ccflow/base.py +++ b/ccflow/base.py @@ -699,14 +699,14 @@ def create_config_from_path( Returns: The instance of the model registry, with the configs loaded. """ - import hydra # Heavy import, only import if used. + from .config import compose, initialize_config_dir overrides = overrides or [] path = pathlib.Path(path).absolute() # Hydra requires absolute paths if not path.parent.exists(): raise OSError(f"Path does not exist: {path.parent}") - with hydra.initialize_config_dir(version_base=version_base, config_dir=str(path.parent)): - cfg = hydra.compose(config_name=path.name, overrides=overrides) + with initialize_config_dir(version_base=version_base, config_dir=str(path.parent)): + cfg = compose(config_name=path.name, overrides=overrides) return cfg def load_config_from_path( @@ -882,7 +882,7 @@ def _add_pending(self, name: str, cfg: Any, overwrite: bool, lookup_registries: self._pending_lookup_registries[name] = list(lookup_registries) def _materialize(self, name: str) -> BaseModel: - from hydra.utils import instantiate + from .config import instantiate key = (id(self), name) stack = _LAZY_LOADING_STACK.get() @@ -1112,10 +1112,10 @@ def load_config( # This also allows for nested attributes on the model itself to # be constructed, even if they are not themselves of BaseModel type, # or if they are of a specific subclass of the parent. - from hydra.errors import InstantiationException - from hydra.utils import instantiate from omegaconf import OmegaConf, UnsupportedValueType + from .config import InstantiationException, instantiate + if resolve_from is not None and resolve_from is not registry: initial_chain = [resolve_from, registry] else: diff --git a/ccflow/config.py b/ccflow/config.py new file mode 100644 index 00000000..8119ed7c --- /dev/null +++ b/ccflow/config.py @@ -0,0 +1,72 @@ +"""Expose the selected configuration framework through one API. + +Lerna is preferred when available unless ``CCFLOW_CONFIG_FRAMEWORK=hydra``. +Applications that select Hydra must install ``hydra-core``. +""" + +import importlib +import os +from types import ModuleType + +__all__ = ( + "CONFIG_FRAMEWORK", + "HydraConfig", + "InstantiationException", + "compose", + "initialize", + "initialize_config_dir", + "instantiate", + "main", +) + +_ENV_VAR = "CCFLOW_CONFIG_FRAMEWORK" +_requested = os.environ.get(_ENV_VAR, "").strip().lower() + +if _requested not in ("", "hydra", "lerna"): + raise ValueError(f"{_ENV_VAR} must be 'hydra', 'lerna', or unset, got {_requested!r}") + + +def _import_framework(name: str) -> ModuleType: + package = "hydra-core" if name == "hydra" else name + try: + return importlib.import_module(name) + except ModuleNotFoundError as error: + if error.name != name: + raise + raise ImportError(f"{_ENV_VAR}={name} requires the '{package}' package") from error + + +if _requested == "hydra": + _framework = _import_framework("hydra") +elif _requested == "lerna": + _framework = _import_framework("lerna") +else: + try: + _framework = importlib.import_module("lerna") + except ModuleNotFoundError as error: + if error.name != "lerna": + raise + _framework = _import_framework("hydra") + +CONFIG_FRAMEWORK = _framework.__name__ +main = _framework.main +compose = _framework.compose +initialize = _framework.initialize +initialize_config_dir = _framework.initialize_config_dir + +if CONFIG_FRAMEWORK == "lerna": + DefaultsList = importlib.import_module("lerna._internal.defaults_list").DefaultsList + GlobalHydra = importlib.import_module("lerna.core.global_hydra").GlobalHydra + HydraConfig = importlib.import_module("lerna.core.hydra_config").HydraConfig + InstantiationException = importlib.import_module("lerna.errors").InstantiationException + ObjectType = importlib.import_module("lerna.core.object_type").ObjectType + RunMode = importlib.import_module("lerna.types").RunMode + instantiate = importlib.import_module("lerna.utils").instantiate +else: + DefaultsList = importlib.import_module("hydra._internal.defaults_list").DefaultsList + GlobalHydra = importlib.import_module("hydra.core.global_hydra").GlobalHydra + HydraConfig = importlib.import_module("hydra.core.hydra_config").HydraConfig + InstantiationException = importlib.import_module("hydra.errors").InstantiationException + ObjectType = importlib.import_module("hydra.core.object_type").ObjectType + RunMode = importlib.import_module("hydra.types").RunMode + instantiate = importlib.import_module("hydra.utils").instantiate diff --git a/ccflow/examples/calculator/__main__.py b/ccflow/examples/calculator/__main__.py index 19210cfa..89181aa7 100644 --- a/ccflow/examples/calculator/__main__.py +++ b/ccflow/examples/calculator/__main__.py @@ -1,11 +1,10 @@ -import hydra - +from ccflow import config from ccflow.utils.hydra import cfg_run __all__ = ("main",) -@hydra.main(config_path="config", config_name="base", version_base=None) +@config.main(config_path="config", config_name="base", version_base=None) def main(cfg): cfg_run(cfg) diff --git a/ccflow/examples/etl/__main__.py b/ccflow/examples/etl/__main__.py index 73f13c9f..b280ae33 100644 --- a/ccflow/examples/etl/__main__.py +++ b/ccflow/examples/etl/__main__.py @@ -1,11 +1,10 @@ -import hydra - +from ccflow import config from ccflow.utils.hydra import cfg_run __all__ = ("main",) -@hydra.main(config_path="config", config_name="base", version_base=None) +@config.main(config_path="config", config_name="base", version_base=None) def main(cfg): cfg_run(cfg) diff --git a/ccflow/tests/config_user/sample2.yml b/ccflow/tests/config_user/sample2.yml deleted file mode 100644 index 9a3e99e7..00000000 --- a/ccflow/tests/config_user/sample2.yml +++ /dev/null @@ -1,6 +0,0 @@ -user_bar: - _target_: ccflow.tests.test_base_registry.MyNestedModel - x: foo - y: # Note that when type is defined on parent model, no need to specify _target_ - a: test2 - b: 2.0 \ No newline at end of file diff --git a/ccflow/tests/enums/test_enums.py b/ccflow/tests/enums/test_enums.py index f799ece4..55711031 100644 --- a/ccflow/tests/enums/test_enums.py +++ b/ccflow/tests/enums/test_enums.py @@ -12,11 +12,18 @@ def auto(): class TestEnum(TestCase): + def setUp(self) -> None: + self._csp_modules = {name: module for name, module in sys.modules.items() if name == "csp" or name.startswith("csp.")} + def tearDown(self) -> None: # Because test_init_parent and test_init_parent_csp muck around with imports # Make sure we always rest the imports at the end of each test so that other # tests are unaffected os.environ.pop("CCFLOW_NO_CSP", None) + for name in list(sys.modules): + if name == "csp" or name.startswith("csp."): + sys.modules.pop(name) + sys.modules.update(self._csp_modules) importlib.invalidate_caches() import ccflow.enums diff --git a/ccflow/tests/test_base_registry.py b/ccflow/tests/test_base_registry.py index 0d20b399..8e81135f 100644 --- a/ccflow/tests/test_base_registry.py +++ b/ccflow/tests/test_base_registry.py @@ -9,13 +9,13 @@ from unittest import TestCase, mock import pytest -from hydra.errors import InstantiationException from omegaconf import OmegaConf from omegaconf.errors import InterpolationKeyError from pydantic import ConfigDict, Field from ccflow import BaseModel, LazyRegistry, ModelRegistry, RegistryLookupContext, RootModelRegistry, model_alias from ccflow.base import RegistryKeyError, resolve_str +from ccflow.config import InstantiationException class MyTestModel(BaseModel): @@ -461,6 +461,25 @@ def test_load_config(self): r.load_config(cfg, overwrite=True) self.assertEqual(r["foo"], m) + def test_load_config_uses_selected_framework(self): + from ccflow import config + + cfg = OmegaConf.create( + { + "foo": { + "_target_": "ccflow.tests.test_base_registry.MyTestModel", + "a": "test", + "b": 0.0, + } + } + ) + registry = ModelRegistry(name="test") + + with mock.patch("ccflow.config.instantiate", wraps=config.instantiate) as instantiate: + registry.load_config(cfg) + + instantiate.assert_called_once() + def test_load_config_with_function(self): cfg = OmegaConf.create( { @@ -754,6 +773,17 @@ def test_defers_models_until_access(self): self.assertIs(root["/lazy/group/source"], source) self.assertEqual(LazyTestModel.constructions, 1) + def test_materialization_uses_selected_framework(self): + from ccflow import config + + root = self._load_registry() + + with mock.patch("ccflow.config.instantiate", wraps=config.instantiate) as instantiate: + source = root["/lazy/group/source"] + + self.assertIsInstance(source, LazyTestModel) + instantiate.assert_called_once() + def test_materializes_dependency_closure(self): root = self._load_registry() @@ -870,7 +900,9 @@ def test_cross_registry_cycle_reports_full_path(self): left["model"] def test_concurrent_cross_registry_cycle_does_not_deadlock(self): - from hydra.utils import instantiate as hydra_instantiate + from ccflow import config + + selected_instantiate = config.instantiate root = ModelRegistry.root() left = LazyRegistry(name="left") @@ -890,10 +922,10 @@ def synchronized_instantiate(*args, **kwargs): if not getattr(local, "started", False): local.started = True barrier.wait(timeout=2) - return hydra_instantiate(*args, **kwargs) + return selected_instantiate(*args, **kwargs) with ( - mock.patch("hydra.utils.instantiate", side_effect=synchronized_instantiate), + mock.patch("ccflow.config.instantiate", side_effect=synchronized_instantiate), concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor, ): futures = [executor.submit(registry.__getitem__, "model") for registry in (left, right)] @@ -979,6 +1011,7 @@ def test_pending_config_is_read_only_copy(self): self.assertEqual(registry["pending"].a, "pending") def test_recursive_hydra_instantiation_is_rejected(self): + from hydra.errors import InstantiationException as HydraInstantiationException from hydra.utils import instantiate cfg = OmegaConf.create( @@ -993,7 +1026,7 @@ def test_recursive_hydra_instantiation_is_rejected(self): } ) - with self.assertRaisesRegex(InstantiationException, "Set '_recursive_: false'"): + with self.assertRaisesRegex(HydraInstantiationException, "Set '_recursive_: false'"): instantiate(cfg, _convert_="all") def test_resolve_from_applies_to_direct_pending_entry(self): diff --git a/ccflow/tests/test_config.py b/ccflow/tests/test_config.py new file mode 100644 index 00000000..f1a83496 --- /dev/null +++ b/ccflow/tests/test_config.py @@ -0,0 +1,165 @@ +import os +import subprocess +import sys + +import pytest + + +def _config_framework( + env_value: str | None, + *, + hide_hydra: bool = False, + hide_lerna: bool = False, +) -> subprocess.CompletedProcess[str]: + env = os.environ.copy() + if env_value is None: + env.pop("CCFLOW_CONFIG_FRAMEWORK", None) + else: + env["CCFLOW_CONFIG_FRAMEWORK"] = env_value + hidden_modules = [] + if hide_hydra: + hidden_modules.append("hydra") + if hide_lerna: + hidden_modules.append("lerna") + hide_modules_statement = "".join(f'sys.modules["{name}"] = None;' for name in hidden_modules) + return subprocess.run( + [ + sys.executable, + "-c", + f"import sys; {hide_modules_statement} from ccflow import config; print(config.CONFIG_FRAMEWORK)", + ], + capture_output=True, + text=True, + env=env, + check=False, + ) + + +def _compose_config(config_dir, framework: str, overrides: list[str] | None = None) -> subprocess.CompletedProcess[str]: + env = os.environ.copy() + env["CCFLOW_CONFIG_FRAMEWORK"] = framework + return subprocess.run( + [ + sys.executable, + "-c", + ( + "from ccflow import config; " + f"context = config.initialize_config_dir(config_dir={str(config_dir)!r}, version_base=None); " + f"context.__enter__(); cfg = config.compose(config_name='config', overrides={overrides or []!r}); " + "print(cfg); context.__exit__(None, None, None)" + ), + ], + capture_output=True, + text=True, + env=env, + check=False, + ) + + +def test_lerna_is_preferred_when_available(): + pytest.importorskip("lerna") + + result = _config_framework(None) + + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "lerna" + + +def test_hydra_can_be_selected_explicitly(): + pytest.importorskip("hydra") + + result = _config_framework("hydra") + + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "hydra" + + +def test_lerna_can_be_selected_explicitly(): + pytest.importorskip("lerna") + + result = _config_framework("lerna") + + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "lerna" + + +def test_hydra_is_used_when_lerna_is_unavailable(): + pytest.importorskip("hydra") + + result = _config_framework(None, hide_lerna=True) + + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "hydra" + + +def test_explicit_unavailable_lerna_fails_clearly(): + result = _config_framework("lerna", hide_lerna=True) + + assert result.returncode != 0 + assert "CCFLOW_CONFIG_FRAMEWORK=lerna requires the 'lerna' package" in result.stderr + + +def test_explicit_unavailable_hydra_fails_clearly(): + result = _config_framework("hydra", hide_hydra=True) + + assert result.returncode != 0 + assert "CCFLOW_CONFIG_FRAMEWORK=hydra requires the 'hydra-core' package" in result.stderr + + +def test_unknown_framework_is_rejected(): + result = _config_framework("hydraa") + + assert result.returncode != 0 + assert "CCFLOW_CONFIG_FRAMEWORK must be 'hydra', 'lerna', or unset" in result.stderr + + +def test_public_api_used_by_downstream_applications(): + from ccflow import config + + assert config.CONFIG_FRAMEWORK in ("hydra", "lerna") + assert callable(config.main) + assert callable(config.compose) + assert callable(config.initialize) + assert callable(config.initialize_config_dir) + assert callable(config.instantiate) + assert hasattr(config.HydraConfig, "get") + assert issubclass(config.InstantiationException, Exception) + + +def test_lerna_patch_is_available_through_facade(tmp_path): + pytest.importorskip("hydra") + group = tmp_path / "group" + group.mkdir() + (group / "base.yaml").write_text("drop_me: true\nkeep: 42\n") + (tmp_path / "config.yaml").write_text( + """defaults: + - group/base@_here_ + - _self_ + - _patch_: + - ~drop_me +""" + ) + + lerna_result = _compose_config(tmp_path, "lerna") + hydra_result = _compose_config(tmp_path, "hydra") + + assert lerna_result.returncode == 0, lerna_result.stderr + assert "drop_me" not in lerna_result.stdout + assert "'keep': 42" in lerna_result.stdout + assert hydra_result.returncode != 0 + assert "Could not load '_patch_/~drop_me'" in hydra_result.stderr + + +def test_lerna_list_operations_are_available_through_facade(tmp_path): + pytest.importorskip("hydra") + (tmp_path / "config.yaml").write_text("tags: [one, two]\n") + + lerna_result = _compose_config(tmp_path, "lerna", ["tags=append(three)"]) + hydra_result = _compose_config(tmp_path, "hydra", ["tags=append(three)"]) + + assert lerna_result.returncode == 0, lerna_result.stderr + assert "one" in lerna_result.stdout + assert "two" in lerna_result.stdout + assert "three" in lerna_result.stdout + assert hydra_result.returncode != 0 + assert "Unknown function 'append'" in hydra_result.stderr diff --git a/ccflow/tests/utils/test_hydra.py b/ccflow/tests/utils/test_hydra.py index 1b2f59b5..888e2d55 100644 --- a/ccflow/tests/utils/test_hydra.py +++ b/ccflow/tests/utils/test_hydra.py @@ -4,8 +4,8 @@ from unittest.mock import MagicMock import pytest -from hydra import compose, initialize +from ccflow.config import compose, initialize from ccflow.utils.hydra import ( add_hydra_config_args, add_panel_server_args, @@ -265,14 +265,16 @@ def test_config_dir_with_overrides(basepath): assert "user_foo" in result.cfg["config_user"] -def test_config_name_yml_not_yaml(basepath): +def test_config_name_yml_not_yaml(basepath, tmp_path): root_config_dir = str(Path(__file__).resolve().parent.parent / "config") - config_dir = str(Path(__file__).resolve().parent.parent / "config_user") + config_dir = tmp_path / "config_user" + config_dir.mkdir() + (config_dir / "sample2.yml").write_text("foo: bar") with pytest.raises(ValueError): load_config( root_config_dir=root_config_dir, root_config_name="conf", - config_dir=config_dir, + config_dir=str(config_dir), config_name="sample2", basepath=basepath, ) @@ -338,7 +340,7 @@ def test_debug(basepath): assert "hydra/job_logging" in result.group_options assert len(result.group_options["hydra/job_logging"]) > 1 assert "config_user" in result.group_options - assert result.group_options["config_user"] == ["sample"] + assert "sample" in result.group_options["config_user"] # Arguable whether these should be here assert "conf_out_of_order" in result.group_options[""] @@ -347,7 +349,7 @@ def test_debug(basepath): assert merged assert "foo" in merged assert "config_user" in merged - assert merged["config_user"]["__options__"] == ["sample"] + assert "sample" in merged["config_user"]["__options__"] assert merged["config_user"]["__parent__"] == "conf" # Maybe this should be a path to a file assert merged["config_user"]["__selected__"] == "sample" assert "user_foo" in merged["config_user"] diff --git a/ccflow/utils/hydra.py b/ccflow/utils/hydra.py index bc07dea6..69325773 100644 --- a/ccflow/utils/hydra.py +++ b/ccflow/utils/hydra.py @@ -1,6 +1,7 @@ import argparse import inspect import os +import sys from collections.abc import Callable from dataclasses import dataclass from logging import getLogger @@ -9,11 +10,11 @@ from textwrap import dedent from typing import Any -from hydra._internal.defaults_list import DefaultsList from omegaconf import DictConfig, ListConfig, OmegaConf from ..base import ModelRegistry from ..callable import FlowOptions, FlowOptionsOverride +from ..config import DefaultsList, GlobalHydra, ObjectType, RunMode, compose, initialize_config_dir _log = getLogger(__name__) @@ -111,8 +112,6 @@ def _find_group_options(config_loader, path, config_name, overrides, results): Note that it will pick up config files that are not intended to be used as config group options, but that exist to provide common config options to other files in the group (i.e. to default) """ - from hydra.core.object_type import ObjectType - groups = config_loader.get_group_options(path, ObjectType.GROUP, config_name, overrides) options = config_loader.get_group_options(path, ObjectType.CONFIG, config_name, overrides) if options: @@ -179,11 +178,6 @@ def load_config( basepath: The base path to start searching for the `config_dir`. This is useful when you want to load from an absolute (rather than relative) path. debug: (Experimental) Whether to enable debug mode. This will return more information about the configs on ConfigLoadResult. """ - # Heavy import, only import if used - import os - - from hydra import compose, initialize_config_dir - if return_hydra_config and debug: raise ValueError("Cannot return hydra config and debug=True at the same time. Please set return_hydra_config=False.") @@ -205,22 +199,37 @@ def load_config( result = ConfigLoadResult(root_config_dir=root_config_dir, root_config_name=root_config_name, cfg=cfg) if debug: import yaml - from hydra.core.global_hydra import GlobalHydra - from hydra.types import RunMode # To track the source file for each config value, we need to monkey patch the yaml loader original_yaml_load = yaml.load + # Lerna may use a Rust YAML parser that bypasses yaml.load entirely. + # Temporarily disable it so our monkey patch can intercept all YAML loading. + _rust_patches = {} + for _mod_name in ( + "lerna._internal.core_plugins.file_config_source", + "lerna._internal.core_plugins.importlib_resources_config_source", + ): + _mod = sys.modules.get(_mod_name) + if _mod and getattr(_mod, "_RUST_AVAILABLE", False): + _rust_patches[_mod] = True + _mod._RUST_AVAILABLE = False + try: def yaml_load(*args, **kwargs): res = original_yaml_load(*args, **kwargs) - return _dict_add_source(res, args[0].name) + # hydra passes file objects (with .name) to yaml.load; + # lerna passes strings. Use "unknown" as fallback source. + source = getattr(args[0], "name", "unknown") if args else "unknown" + return _dict_add_source(res, source) yaml.load = yaml_load # We can't load the hydra config after monkey patching yaml loading, so skip that step result.cfg_sources = compose(config_name=root_config_name, overrides=overrides, return_hydra_config=False) finally: yaml.load = original_yaml_load + for _mod, _val in _rust_patches.items(): + _mod._RUST_AVAILABLE = _val config_loader = GlobalHydra.instance().config_loader() # Load defaults list using the standard hydra function diff --git a/pyproject.toml b/pyproject.toml index f3de296f..76e3f9f6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,8 +39,8 @@ dependencies = [ "cloudpathlib", "cloudpickle", "deprecated", - "hydra-core", "jinja2", + "lerna>=2.1.0", "narwhals", "numpy<3", "orjson", @@ -89,6 +89,7 @@ develop = [ "csp>=0.8.0,<1; python_version < '3.14'", # TODO: remove when when 3.14 wheels disted "duckdb", "IPython", + "hydra-core", "panel", "panel_material_ui", "plotly",