Skip to content
Merged
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
23 changes: 17 additions & 6 deletions src/objectstate/context_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,19 +66,31 @@ def _merge_nested_dataclass(base, override, mask_with_none: bool = False):
if not is_dataclass(base) or not is_dataclass(override):
return override

from objectstate.lazy_factory import get_base_type_for_lazy, replace_raw

base_type = get_base_type_for_lazy(type(base)) or type(base)
override_type = get_base_type_for_lazy(type(override)) or type(override)
# Lazy overlays of the same owner inherit on the base instance. An authored
# nominal change instead owns the result, including subtype-only fields.
result_owner = base if base_type is override_type else override
base_fields = {field_info.name for field_info in fields(base)}
merge_values = {}
for field_info in fields(override):
field_name = field_info.name
override_value = object.__getattribute__(override, field_name)
if field_name not in base_fields:
# This field belongs only to the authored nominal owner.
continue
base_value = object.__getattribute__(base, field_name)

if override_value is None:
if mask_with_none:
# None overrides base value (masking mode)
merge_values[field_name] = None
else:
# None means "don't override" - keep base value
continue
# None means inherit, including when the result owner changes.
if result_owner is override:
merge_values[field_name] = base_value
elif is_dataclass(override_value):
# Recursively merge nested dataclass
if base_value is not None and is_dataclass(base_value):
Expand All @@ -89,13 +101,12 @@ def _merge_nested_dataclass(base, override, mask_with_none: bool = False):
# Concrete value - use override
merge_values[field_name] = override_value

# Merge with base using replace_raw to preserve None values
# Merge on the nominal owner using replace_raw to preserve None values
# (dataclasses.replace triggers lazy resolution, baking in resolved values)
if merge_values:
from objectstate.lazy_factory import replace_raw
return replace_raw(base, **merge_values)
return replace_raw(result_owner, **merge_values)
else:
return base
return result_owner


@contextmanager
Expand Down
52 changes: 52 additions & 0 deletions tests/test_context_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
clear_current_temp_global,
merge_configs,
extract_all_configs,
set_base_config_type,
)


Expand All @@ -23,6 +24,57 @@ def test_config_context_basic(global_config):
assert current.num_workers == 8


@pytest.mark.parametrize("mask_with_none", [False, True])
def test_nested_context_preserves_authored_subtype_and_sibling_fields(mask_with_none):
@dataclass
class SpatialConfig:
width: int | None = 12

@dataclass
class VolumeConfig(SpatialConfig):
depth: int = 60

@dataclass
class SourceConfig:
names: tuple[str, ...] = ()
spatial: SpatialConfig = field(default_factory=SpatialConfig)

@dataclass
class GlobalConfig:
sources: SourceConfig = field(default_factory=SourceConfig)

set_base_config_type(GlobalConfig)
authored = GlobalConfig(SourceConfig(("DNA", "Mito"), VolumeConfig(width=None)))
with config_context(authored, mask_with_none=mask_with_none, use_live_global=False):
current = get_current_temp_global()
assert current.sources.names == ("DNA", "Mito")
assert type(current.sources.spatial) is VolumeConfig
assert current.sources.spatial.depth == 60
assert current.sources.spatial.width == (None if mask_with_none else 12)
assert authored.sources.spatial.width is None


def test_nested_context_respects_explicit_subtype_downgrade():
@dataclass
class SpatialConfig:
width: int | None = 12

@dataclass
class VolumeConfig(SpatialConfig):
depth: int = 60

@dataclass
class GlobalConfig:
spatial: SpatialConfig = field(default_factory=SpatialConfig)

set_base_config_type(GlobalConfig)
with config_context(GlobalConfig(VolumeConfig(width=24)), use_live_global=False):
with config_context(GlobalConfig(SpatialConfig(width=None))):
current = get_current_temp_global()
assert type(current.spatial) is SpatialConfig
assert current.spatial.width == 24


def test_config_context_nested(global_config, pipeline_config):
"""Test nested config_context."""
with config_context(global_config):
Expand Down
Loading