Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions ccflow/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 *
Expand Down
12 changes: 6 additions & 6 deletions ccflow/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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:
Expand Down
72 changes: 72 additions & 0 deletions ccflow/config.py
Original file line number Diff line number Diff line change
@@ -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
5 changes: 2 additions & 3 deletions ccflow/examples/calculator/__main__.py
Original file line number Diff line number Diff line change
@@ -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)

Expand Down
5 changes: 2 additions & 3 deletions ccflow/examples/etl/__main__.py
Original file line number Diff line number Diff line change
@@ -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)

Expand Down
6 changes: 0 additions & 6 deletions ccflow/tests/config_user/sample2.yml

This file was deleted.

7 changes: 7 additions & 0 deletions ccflow/tests/enums/test_enums.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
43 changes: 38 additions & 5 deletions ccflow/tests/test_base_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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(
{
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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")
Expand All @@ -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)]
Expand Down Expand Up @@ -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(
Expand All @@ -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):
Expand Down
Loading
Loading