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)