Source code for semi_cr.core.lab.utils.dataclass_base

from __future__ import annotations

import sys
import typing
from collections.abc import Iterable, Mapping
from dataclasses import fields, is_dataclass
from pathlib import Path
from typing import Any, Self, TypeVar, Union, get_args, get_origin

import numpy as np
import yaml

T = TypeVar("T")

DEFAULT_SKIP_NAMES = {
    # backrefs / cycle magnets
    "parent", "_parent", "owner",
    # common caches / internals (edit to taste)
    "__weakref__", "__dict__",
}

DEFAULT_SKIP_PREFIXES = ("_",)  # skip private attrs from __dict__ by default


def _strip_optional(tp: Any) -> Any:
    origin = get_origin(tp)
    if origin is Union:
        args = [a for a in get_args(tp) if a is not type(None)]
        if len(args) == 1:
            return args[0]
    return tp


def _safe_get_type_hints(cls) -> dict[str, Any]:
    # Avoid exploding on forward refs / circular refs in annotations
    mod = sys.modules.get(cls.__module__)
    globalns = vars(mod) if mod else {}
    localns = dict(vars(cls))
    try:
        return typing.get_type_hints(cls, globalns=globalns, localns=localns)
    except NameError:
        return {}


def _convert_value(target_type: Any, value: Any) -> Any:
    if value is None:
        return None

    target_type = _strip_optional(target_type)
    origin = get_origin(target_type)

    if origin is list:
        (item_type,) = get_args(target_type)
        return [_convert_value(item_type, v) for v in value]

    if origin is dict:
        key_type, val_type = get_args(target_type)
        return {
            _convert_value(key_type, k): _convert_value(val_type, v)
            for k, v in value.items()
        }

    # value is dict and target_type has from_dict
    if isinstance(value, dict) and isinstance(target_type, type):
        from_dict = getattr(target_type, "from_dict", None)
        if callable(from_dict):
            return from_dict(value)

        if is_dataclass(target_type):
            return _from_dict_dataclass(target_type, value)

    return value


SKIP_FIELD_NAMES = {"owner", "parent"}


def _from_dict_dataclass(cls, data):
    type_hints = _safe_get_type_hints(cls)
    kwargs = {}
    for f in fields(cls):
        if not f.init:
            continue
        if f.name in SKIP_FIELD_NAMES:   # <-- IMPORTANT
            continue
        if f.name not in data:
            continue
        field_type = type_hints.get(f.name, f.type)
        kwargs[f.name] = _convert_value(field_type, data[f.name])
    return cls(**kwargs)

[docs] class Dataclass:
[docs] @classmethod def from_dict(cls, data: dict) -> Self: if not is_dataclass(cls): raise ValueError(f"{cls.__name__} must be a dataclass") return _from_dict_dataclass(cls, data)
[docs] @classmethod def from_yaml(cls, path: str | Path) -> Self: with open(path, encoding="utf-8") as f: data = yaml.safe_load(f) return cls.from_dict(data)
[docs] def to_dict_full( cls, *, include_dataclass_init_false: bool = True, # include init=False fields too include_object_dict: bool = True, # expand __dict__ for non-dataclass objects include_properties: bool = False, # risky: may compute a lot skip_names: set[str] | None = None, skip_prefixes: tuple[str, ...] = DEFAULT_SKIP_PREFIXES, drop_none: bool = False, max_depth: int = 50, add_class_tag: bool = True, # include "__class__" ) -> Any: """ Snapshot serializer: produces a YAML-friendly structure representing *everything* on the object graph, including runtime-attached objects like ff_params / center_coords / shape_params. Cycle-safe, depth-limited. """ if skip_names is None: skip_names = set(DEFAULT_SKIP_NAMES) memo: dict[int, Any] = {} in_progress: set[int] = set() def _iter_properties_in_mro(cls: type) -> Iterable[str]: seen = set() for base in cls.__mro__: for name, attr in vars(base).items(): if name in seen: continue if isinstance(attr, property): seen.add(name) yield name def _should_skip_attr(name: str) -> bool: if name in skip_names: return True if skip_prefixes and name.startswith(skip_prefixes): return True return False def _convert(x: Any, depth: int) -> Any: if depth > max_depth: return f"<max_depth {max_depth} reached>" # numpy scalars -> python scalars if np is not None and isinstance(x, np.generic): return x.item() # primitives if x is None or isinstance(x, (bool, int, float, str)): return x oid = id(x) if oid in memo: return memo[oid] if oid in in_progress: # cycle detected return f"<cycle {type(x).__name__}>" # dataclass objects if is_dataclass(x): in_progress.add(oid) dataclass_out: dict[str, Any] = {} if add_class_tag: dataclass_out["__class__"] = type(x).__name__ memo[oid] = dataclass_out for f in fields(x): if (not f.init) and (not include_dataclass_init_false): continue name = f.name if _should_skip_attr(name): continue try: val = getattr(x, name) except Exception as e: dataclass_out[name] = f"<error reading field: {e!r}>" continue if drop_none and val is None: continue dataclass_out[name] = _convert(val, depth + 1) if include_properties: for name in _iter_properties_in_mro(type(x)): if _should_skip_attr(name): continue try: val = getattr(x, name) if drop_none and val is None: continue dataclass_out[name] = _convert(val, depth + 1) except Exception as e: dataclass_out[name] = f"<error property: {e!r}>" in_progress.remove(oid) return dataclass_out # mapping if isinstance(x, Mapping): in_progress.add(oid) mapping_out: dict[Any, Any] = {} memo[oid] = mapping_out for k, v in x.items(): kk = k if isinstance(k, str) else str(k) vv = _convert(v, depth + 1) if drop_none and vv is None: continue mapping_out[kk] = vv in_progress.remove(oid) return mapping_out # sequence if isinstance(x, (list, tuple, set)): in_progress.add(oid) out_list: list[Any] = [] memo[oid] = out_list for v in x: vv = _convert(v, depth + 1) if drop_none and vv is None: continue out_list.append(vv) in_progress.remove(oid) return out_list if not isinstance(x, tuple) else list(out_list) # plain object: expand __dict__ if include_object_dict and hasattr(x, "__dict__"): in_progress.add(oid) object_out: dict[str, Any] = {} if add_class_tag: object_out["__class__"] = type(x).__name__ memo[oid] = object_out for name, val in vars(x).items(): if _should_skip_attr(name): continue if drop_none and val is None: continue object_out[name] = _convert(val, depth + 1) if include_properties: for name in _iter_properties_in_mro(type(x)): if _should_skip_attr(name): continue try: val = getattr(x, name) if drop_none and val is None: continue object_out[name] = _convert(val, depth + 1) except Exception as e: object_out[name] = f"<error property: {e!r}>" in_progress.remove(oid) return object_out # fallback return repr(x) return _convert(cls, 0)
[docs] @classmethod def to_yaml_full(self, path: str | Path, **kwargs) -> None: data = self.to_dict_full(**kwargs) with Path(path).open("w", encoding="utf-8") as f: yaml.safe_dump(data, f, sort_keys=False)