diff --git a/src/objectstate/context_manager.py b/src/objectstate/context_manager.py index f1a7728..6f63e9a 100644 --- a/src/objectstate/context_manager.py +++ b/src/objectstate/context_manager.py @@ -66,10 +66,21 @@ 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: @@ -77,8 +88,9 @@ def _merge_nested_dataclass(base, override, mask_with_none: bool = False): # 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): @@ -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 diff --git a/tests/test_context_manager.py b/tests/test_context_manager.py index 4edc5f9..83cc92d 100644 --- a/tests/test_context_manager.py +++ b/tests/test_context_manager.py @@ -11,6 +11,7 @@ clear_current_temp_global, merge_configs, extract_all_configs, + set_base_config_type, ) @@ -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):