diff --git a/src/furu/_pytree.py b/src/furu/_pytree.py new file mode 100644 index 00000000..2b12479e --- /dev/null +++ b/src/furu/_pytree.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import copy +import dataclasses +from collections.abc import Callable, Iterable +from dataclasses import dataclass +from typing import Literal, cast + +from pydantic import BaseModel + +type IsLeaf = Callable[[object], bool] +type NodeKind = Literal[ + "list", "tuple", "set", "frozenset", "dict", "dataclass", "pydantic" +] + + +@dataclass(frozen=True) +class PyTreeNode: + kind: NodeKind + value: object + entries: tuple[object, ...] + children: tuple[object, ...] + + def unflatten(self, children: Iterable[object]) -> object: + values = tuple(children) + match self.kind: + case "list" | "set" | "frozenset": + return type(self.value)(values) + case "tuple": + return ( + type(self.value)(*values) + if hasattr(type(self.value), "_fields") + else type(self.value)(values) + ) + case "dict": + return type(self.value)(zip(self.entries, values, strict=True)) + case "dataclass": + result = copy.copy(self.value) + for name, child in zip(self.entries, values, strict=True): + object.__setattr__(result, cast(str, name), child) + return result + case "pydantic": + assert isinstance(self.value, BaseModel) + return self.value.model_copy( + update=dict( + zip(cast(tuple[str, ...], self.entries), values, strict=True) + ) + ) + + +def tree_node(value: object) -> PyTreeNode | None: + if isinstance(value, BaseModel): + entries = tuple(type(value).model_fields) + kind: NodeKind = "pydantic" + children = tuple(getattr(value, name) for name in entries) + elif dataclasses.is_dataclass(value) and not isinstance(value, type): + entries = tuple(field.name for field in dataclasses.fields(value)) + kind = "dataclass" + children = tuple(getattr(value, name) for name in entries) + elif isinstance(value, (list, tuple)): + children = tuple(value) + entries = tuple(range(len(children))) + kind = "list" if isinstance(value, list) else "tuple" + elif isinstance(value, (set, frozenset)): + children = tuple( + sorted( + value, + key=lambda item: ( + type(item).__module__, type(item).__qualname__, repr(item) + ), + ) + ) + entries = tuple(range(len(children))) + kind = "frozenset" if isinstance(value, frozenset) else "set" + elif isinstance(value, dict): + mapping = cast(dict[object, object], value) + entries = tuple(mapping) + children = tuple(mapping[key] for key in entries) + kind = "dict" + else: + return None + return PyTreeNode(kind, value, entries, children) + + +def tree_map( + function: Callable[[object], object], + value: object, + *, + is_leaf: IsLeaf | None = None, +) -> object: + node = None if is_leaf is not None and is_leaf(value) else tree_node(value) + if node is None: + return function(value) + return node.unflatten( + tree_map(function, child, is_leaf=is_leaf) for child in node.children + ) + + +def tree_leaves(value: object, *, is_leaf: IsLeaf | None = None) -> list[object]: + leaves: list[object] = [] + + def collect(leaf: object) -> object: + leaves.append(leaf) + return leaf + + tree_map(collect, value, is_leaf=is_leaf) + return leaves diff --git a/src/furu/dag.py b/src/furu/dag.py index 975893bd..f3b8775f 100644 --- a/src/furu/dag.py +++ b/src/furu/dag.py @@ -3,8 +3,9 @@ from collections.abc import Sequence from dataclasses import dataclass, field from functools import cached_property -from typing import TYPE_CHECKING, assert_never +from typing import TYPE_CHECKING, assert_never, cast +from furu._pytree import tree_leaves from furu.core import Spec from furu.dependencies import collect_declared_refs from furu.metadata import ArtifactSpec @@ -27,11 +28,11 @@ def batch_group(self) -> tuple[object, int] | None: return _batch_group(self.obj) -def _add_to_dag(coordinator: ExecutionCoordinator, objs: Sequence[Spec]) -> None: - if any(not isinstance(obj, Spec) for obj in objs): - # TODO: accept pytrees of Spec objects (e.g. nested lists/dicts/dataclasses) - # and flatten them before walking dependencies. - raise TypeError("expected Spec objects") +def _add_to_dag(coordinator: ExecutionCoordinator, obj_tree: object) -> list[Spec]: + leaves = tree_leaves(obj_tree, is_leaf=lambda value: isinstance(value, Spec)) + if any(not isinstance(obj, Spec) for obj in leaves): + raise TypeError("expected Spec objects in every PyTree leaf") + objs = cast(list[Spec], leaves) refs_by_id: dict[str, tuple[Spec, ...]] = {} newly_added: list[DagNode] = [] @@ -75,6 +76,7 @@ def _add_to_dag(coordinator: ExecutionCoordinator, objs: Sequence[Spec]) -> None coordinator.blocked[node.obj.object_id] = node else: coordinator.ready[node.obj.object_id] = node + return objs def _update_dag_blocking_dependencies( diff --git a/src/furu/dependencies.py b/src/furu/dependencies.py index e1f6b22b..dba7f0db 100644 --- a/src/furu/dependencies.py +++ b/src/furu/dependencies.py @@ -3,11 +3,11 @@ from collections.abc import Callable, Iterator from contextlib import contextmanager from contextvars import ContextVar -from dataclasses import fields, is_dataclass +from dataclasses import fields from functools import cached_property from typing import TYPE_CHECKING, Any, Self, overload -from pydantic import BaseModel as PydanticBaseModel +from furu._pytree import tree_leaves if TYPE_CHECKING: from furu.core import Spec @@ -47,21 +47,11 @@ def dependency[TSpec: Spec[Any], T]( def find_nested_furu_objects(value: object) -> Iterator[Spec]: from furu.core import Spec - match value: - case Spec(): - yield value - case _ if is_dataclass(value) and not isinstance(value, type): - for field in fields(value): - yield from find_nested_furu_objects(getattr(value, field.name)) - case PydanticBaseModel(): - for name in type(value).model_fields: - yield from find_nested_furu_objects(getattr(value, name)) - case tuple() | list() | set() | frozenset(): - for item in value: - yield from find_nested_furu_objects(item) - case dict(): - for item in value.values(): - yield from find_nested_furu_objects(item) + yield from ( + leaf + for leaf in tree_leaves(value, is_leaf=lambda item: isinstance(item, Spec)) + if isinstance(leaf, Spec) + ) def collect_declared_refs(obj: Spec) -> tuple[Spec, ...]: diff --git a/src/furu/execution/execution_coordinator.py b/src/furu/execution/execution_coordinator.py index 21d77874..a96ef6dd 100644 --- a/src/furu/execution/execution_coordinator.py +++ b/src/furu/execution/execution_coordinator.py @@ -3,7 +3,7 @@ import hashlib import threading import time -from collections.abc import Iterator, Sequence +from collections.abc import Iterator from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from dataclasses import dataclass, field @@ -89,9 +89,9 @@ def _counts_detail(self) -> dict[str, object]: } @classmethod - def run[ObjsT: Sequence[Spec]]( + def run[ObjsT]( cls, - objs: ObjsT, # TODO: support pytrees + objs: ObjsT, *, max_retries_per_object: int | None = None, worker_backends: tuple[WorkerBackend, ...], @@ -100,9 +100,9 @@ def run[ObjsT: Sequence[Spec]]( if max_retries_per_object is None: max_retries_per_object = get_config().worker.max_retries_per_object coordinator = cls(max_retries_per_object=max_retries_per_object) - _add_to_dag(coordinator, objs) + root_objs = _add_to_dag(coordinator, objs) digest = hashlib.blake2s(digest_size=16) - for obj in objs: + for obj in root_objs: digest.update(obj.object_id.encode("utf-8")) digest.update(b"\0") coordinator.executor_id = digest.hexdigest() diff --git a/src/furu/execution/load_or_create.py b/src/furu/execution/load_or_create.py index 3873529f..bd1297db 100644 --- a/src/furu/execution/load_or_create.py +++ b/src/furu/execution/load_or_create.py @@ -13,6 +13,7 @@ from furu._batched import _BatchedHook from furu._declared_types import declared_result_type +from furu._pytree import tree_leaves, tree_map from furu.config import get_config from furu.core import Missing, Spec from furu.dependencies import dependency_recorder, record_dependency_call @@ -127,17 +128,23 @@ def _store_result[T]( def _load_or_create[T](obj: Spec[T], *, use_lock: bool = True) -> T: ... +@overload +def _load_or_create[T]( + objs: tuple[Spec[T], ...], *, use_lock: bool = True +) -> tuple[T, ...]: ... @overload def _load_or_create[T]( objs: Sequence[Spec[T]], *, use_lock: bool = True ) -> list[T]: ... +@overload +def _load_or_create(obj_tree: object, *, use_lock: bool = True) -> object: ... def _load_or_create[T]( - obj_or_objs: Spec[T] | Sequence[Spec[T]], + obj_or_objs: object, *, use_lock: bool = True, -) -> T | list[T]: +) -> object: _require_uv() if _in_worker_execution.get(): return _load_or_create_worker(obj_or_objs) @@ -173,45 +180,69 @@ def _ensure_group_result[T]( def _normalize_load_or_create_input[T]( - obj_or_objs: Spec[T] | Sequence[Spec[T]], + obj_or_objs: object, + *, + log_single_create: bool, ) -> tuple[list[Spec[T]], bool]: - match obj_or_objs: - case Spec() as obj: - assert not isinstance(obj, Sequence) - record_dependency_call(obj) - obj.logger.debug(".create called for %s", obj) - return [obj], True - case Sequence() as objs: - for obj in objs: - record_dependency_call(obj) - return list(objs), False + leaves = tree_leaves(obj_or_objs, is_leaf=lambda value: isinstance(value, Spec)) + if any(not isinstance(obj, Spec) for obj in leaves): + raise TypeError("expected a PyTree of Spec objects") + objs = cast(list[Spec[T]], leaves) + for obj in objs: + record_dependency_call(obj) + unwrap = isinstance(obj_or_objs, Spec) + if log_single_create and unwrap: + objs[0].logger.debug(".create called for %s", objs[0]) + return objs, unwrap + + +def _rebuild_spec_tree(obj_tree: object, values: list[object]) -> object: + values_iter = iter(values) + return tree_map( + lambda _: next(values_iter), + obj_tree, + is_leaf=lambda value: isinstance(value, Spec), + ) @overload def create[T](obj: Spec[T], *, on: Sequence[WorkerBackend] | None = None) -> T: ... @overload +def create[T]( + objs: tuple[Spec[T], ...], *, on: Sequence[WorkerBackend] | None = None +) -> tuple[T, ...]: ... +@overload def create[T]( objs: Sequence[Spec[T]], *, on: Sequence[WorkerBackend] | None = None ) -> list[T]: ... +@overload +def create( + obj_tree: object, *, on: Sequence[WorkerBackend] | None = None +) -> object: ... def create[T]( - obj_or_objs: Spec[T] | Sequence[Spec[T]], + obj_or_objs: object, *, on: Sequence[WorkerBackend] | None = None, -) -> T | list[T]: +) -> object: if on is not None: from furu.execution.execution_coordinator import ExecutionCoordinator - objs = [obj_or_objs] if isinstance(obj_or_objs, Spec) else list(obj_or_objs) - ExecutionCoordinator.run(objs, worker_backends=tuple(on)) + ExecutionCoordinator.run(obj_or_objs, worker_backends=tuple(on)) return _load_or_create(obj_or_objs) -def load_existing[T](objs: Sequence[Spec[T]]) -> list[T]: - if not isinstance(objs, Sequence): - raise TypeError("load_existing() expected a sequence of Spec objects") - objs = list(objs) - if any(not isinstance(obj, Spec) for obj in objs): - raise TypeError("load_existing() expected Spec objects") +@overload +def load_existing[T](obj: Spec[T]) -> T: ... +@overload +def load_existing[T](objs: tuple[Spec[T], ...]) -> tuple[T, ...]: ... +@overload +def load_existing[T](objs: Sequence[Spec[T]]) -> list[T]: ... +@overload +def load_existing(obj_tree: object) -> object: ... +def load_existing[T](obj_tree: object) -> object: + objs, _ = _normalize_load_or_create_input( + obj_tree, log_single_create=False + ) loaded: list[T] = [] missing: list[Spec[T]] = [] for obj in objs: @@ -245,7 +276,7 @@ def load_existing[T](objs: Sequence[Spec[T]]) -> list[T]: ) else: get_logger().info("loaded 0 furu objects") - return loaded + return _rebuild_spec_tree(obj_tree, cast(list[object], loaded)) def _cached_to_build_msg(cached: list[Spec[Any]], to_build: list[Spec[Any]]) -> str: @@ -259,9 +290,11 @@ def fmt(objs: list[Spec[Any]]) -> str: def _load_or_create_worker[T]( - obj_or_objs: Spec[T] | Sequence[Spec[T]], -) -> T | list[T]: - objs, unwrap = _normalize_load_or_create_input(obj_or_objs) + obj_or_objs: object, +) -> object: + objs, _ = _normalize_load_or_create_input( + obj_or_objs, log_single_create=True + ) loaded: list[T] = [] cached: list[Spec[T]] = [] @@ -293,21 +326,20 @@ def _load_or_create_worker[T]( call_kind="create", ) - if unwrap: - (result,) = loaded - return result - return loaded + return _rebuild_spec_tree(obj_or_objs, cast(list[object], loaded)) def _load_or_create_local[T]( - obj_or_objs: Spec[T] | Sequence[Spec[T]], + obj_or_objs: object, *, use_lock: bool = True, -) -> T | list[T]: - objs, unwrap = _normalize_load_or_create_input(obj_or_objs) +) -> object: + objs, unwrap = _normalize_load_or_create_input( + obj_or_objs, log_single_create=True + ) if not objs: - return [] + return _rebuild_spec_tree(obj_or_objs, []) unique_by_object_id: dict[str, Spec[T]] = {} for obj in objs: @@ -389,15 +421,13 @@ def _load_or_create_local[T]( if unwrap: (obj,) = objs - (output,) = outputs if direct_create_started: obj.logger.info( "finished %s ok ยท %s", obj._log_label, format_duration(time.monotonic() - create_started_at), ) - return output - return outputs + return _rebuild_spec_tree(obj_or_objs, cast(list[object], outputs)) def _batch_group(obj: Spec[Any]) -> tuple[object, int] | None: diff --git a/src/furu/explain.py b/src/furu/explain.py index c076b1b1..87b8a314 100644 --- a/src/furu/explain.py +++ b/src/furu/explain.py @@ -1,11 +1,10 @@ from __future__ import annotations import json -from dataclasses import fields, is_dataclass from pathlib import Path from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel as PydanticBaseModel +from furu._pytree import tree_node if TYPE_CHECKING: from furu.core import Spec @@ -63,16 +62,9 @@ def _rows(children: list[tuple[str, object]], depth: ExplainDepth) -> list[str]: def _children(value: object) -> list[tuple[str, object]] | None: - if isinstance(value, PydanticBaseModel): - return [(name, getattr(value, name)) for name in type(value).model_fields] - if is_dataclass(value) and not isinstance(value, type): - return [(field.name, getattr(value, field.name)) for field in fields(value)] - if isinstance(value, (tuple, list)): - return [(str(index), item) for index, item in enumerate(value)] - if isinstance(value, (set, frozenset)): - return [ - (str(index), item) for index, item in enumerate(sorted(value, key=repr)) - ] - if isinstance(value, dict): - return [(str(key), item) for key, item in value.items()] - return None + if (node := tree_node(value)) is None: + return None + return [ + (str(entry), child) + for entry, child in zip(node.entries, node.children, strict=True) + ] diff --git a/src/furu/result/bundle.py b/src/furu/result/bundle.py index ad5b6732..b55f2a1b 100644 --- a/src/furu/result/bundle.py +++ b/src/furu/result/bundle.py @@ -19,6 +19,7 @@ import pydantic from furu._declared_types import child_declared_type, strip_annotated +from furu._pytree import tree_node from furu.constants import FIELDSMARKER, KINDMARKER, TYPEMARKER from furu.result.codec import Codec, CodecMeta from furu.result.ref import Ref @@ -127,40 +128,17 @@ def _dump_value( match value: case None | bool() | int() | float() | str(): return value - case list(): - width = len(str(len(value))) - return [ - _dump_value( - item, - declared_type=child_declared_type(declared_type, i), - value_path=(*value_path, f"{i:0{width}d}"), - bundle_dir=bundle_dir, - result_codecs=result_codecs, - dump_state=dump_state, - ) - for i, item in enumerate(value) - ] - case tuple(): - width = len(str(len(value))) + case Path(): return { WRAPPER_KEY: { - KINDMARKER: "tuple", - "items": [ - _dump_value( - item, - declared_type=child_declared_type(declared_type, i), - value_path=(*value_path, f"{i:0{width}d}"), - bundle_dir=bundle_dir, - result_codecs=result_codecs, - dump_state=dump_state, - ) - for i, item in enumerate(value) - ], + KINDMARKER: "path", + "value": str(value), } } - case set() | frozenset(): - kind = "frozenset" if isinstance(value, frozenset) else "set" - for item in value: + + if node := tree_node(value): + if node.kind in ("set", "frozenset"): + for item in node.children: if type(item).__repr__ is object.__repr__: raise ValueError( f"Unsupported result value at {_value_path_display(value_path)}:\n" @@ -168,107 +146,69 @@ def _dump_value( "value-based repr, so their order cannot be made " "deterministic; use a list or implement __repr__." ) - items = sorted( - value, - key=lambda item: ( - type(item).__module__, - type(item).__qualname__, - repr(item), - ), - ) - width = len(str(len(items))) - return { - WRAPPER_KEY: { - KINDMARKER: kind, - "items": [ - _dump_value( - item, - declared_type=child_declared_type(declared_type, i), - value_path=(*value_path, f"{i:0{width}d}"), - bundle_dir=bundle_dir, - result_codecs=result_codecs, - dump_state=dump_state, - ) - for i, item in enumerate(items) - ], - } - } - case dict(): - out: dict[str, JsonValue] = {} - for raw_key, child in value.items(): - key = _validate_result_path_segment( - raw_key, - parent_value_path=value_path, - ) - out[key] = _dump_value( - child, - declared_type=child_declared_type(declared_type, raw_key), - value_path=(*value_path, key), - bundle_dir=bundle_dir, - result_codecs=result_codecs, - dump_state=dump_state, - ) - return out - case Path(): - return { - WRAPPER_KEY: { - KINDMARKER: "path", - "value": str(value), - } - } - case pydantic.BaseModel(): - fields_out: dict[str, JsonValue] = {} - field_types = get_type_hints(value.__class__, include_extras=True) - for raw_name in value.__class__.model_fields: - name = _validate_result_path_segment( - raw_name, parent_value_path=value_path + + field_types = ( + get_type_hints(type(value), include_extras=True) + if node.kind in ("dataclass", "pydantic") + else {} + ) + width = len(str(len(node.children))) + dumped: list[tuple[object, JsonValue]] = [] + for entry, child in zip(node.entries, node.children, strict=True): + if node.kind in ("dict", "dataclass", "pydantic"): + segment = _validate_result_path_segment( + entry, parent_value_path=value_path ) - fields_out[name] = _dump_value( - getattr(value, name), - declared_type=field_types.get(name, Any), - value_path=(*value_path, name), - bundle_dir=bundle_dir, - result_codecs=result_codecs, - dump_state=dump_state, + else: + assert isinstance(entry, int) + segment = f"{entry:0{width}d}" + dumped.append( + ( + entry, + _dump_value( + child, + declared_type=( + field_types.get(cast(str, entry), Any) + if node.kind in ("dataclass", "pydantic") + else child_declared_type(declared_type, entry) + ), + value_path=(*value_path, segment), + bundle_dir=bundle_dir, + result_codecs=result_codecs, + dump_state=dump_state, + ), ) + ) + + if node.kind == "list": + return [child for _, child in dumped] + if node.kind == "dict": + return {cast(str, key): child for key, child in dumped} + if node.kind in ("tuple", "set", "frozenset"): return { WRAPPER_KEY: { - KINDMARKER: "pydantic", - TYPEMARKER: fully_qualified_name(type(value)), - FIELDSMARKER: fields_out, + KINDMARKER: node.kind, + "items": [child for _, child in dumped], } } - case _ if dataclasses.is_dataclass(value) and not isinstance(value, type): - fields_out: dict[str, JsonValue] = {} - field_types = get_type_hints(type(value), include_extras=True) - for field in dataclasses.fields(cast(Any, value)): - name = _validate_result_path_segment( - field.name, parent_value_path=value_path - ) - fields_out[name] = _dump_value( - getattr(value, name), - declared_type=field_types.get(field.name, Any), - value_path=(*value_path, name), - bundle_dir=bundle_dir, - result_codecs=result_codecs, - dump_state=dump_state, - ) - return { - WRAPPER_KEY: { - KINDMARKER: "dataclass", - TYPEMARKER: fully_qualified_name(type(value)), - FIELDSMARKER: fields_out, - } + return { + WRAPPER_KEY: { + KINDMARKER: node.kind, + TYPEMARKER: fully_qualified_name(type(value)), + FIELDSMARKER: { + cast(str, name): child for name, child in dumped + }, } - case _: - if codec := CodecMeta.find_codec(value, result_codecs): - return _dump_artifact( - value, - codec=codec, - value_path=value_path, - bundle_dir=bundle_dir, - dump_state=dump_state, - ) + } + + if codec := CodecMeta.find_codec(value, result_codecs): + return _dump_artifact( + value, + codec=codec, + value_path=value_path, + bundle_dir=bundle_dir, + dump_state=dump_state, + ) raise ValueError( f"Unsupported result value at {_value_path_display(value_path)}:\n" diff --git a/src/furu/serializer/artifact.py b/src/furu/serializer/artifact.py index 21edafa1..a061e8a7 100644 --- a/src/furu/serializer/artifact.py +++ b/src/furu/serializer/artifact.py @@ -1,16 +1,14 @@ import enum -from dataclasses import fields, is_dataclass from datetime import datetime from pathlib import Path from typing import TYPE_CHECKING, Any, cast, get_type_hints -from pydantic import BaseModel as PydanticBaseModel - from furu._declared_types import ( child_declared_type, has_skip_hash, strip_annotated, ) +from furu._pytree import tree_node from furu.constants import ( CLASSMARKER, FIELDSMARKER, @@ -88,100 +86,78 @@ def assert_correct_dict_key(x: Any) -> str: return {KINDMARKER: "path", VALUEMARKER: str(obj)} case datetime(): return {KINDMARKER: "datetime", VALUEMARKER: obj.isoformat()} - case list(): - return [ + + if node := tree_node(obj): + hints = ( + get_type_hints(type(obj), include_extras=True) + if node.kind in ("dataclass", "pydantic") + else {} + ) + entries_and_children = [ + (entry, child) + for entry, child in zip(node.entries, node.children, strict=True) + if not ( + for_hash + and node.kind in ("dataclass", "pydantic") + and has_skip_hash(hints.get(cast(str, entry), Any)) + ) + ] + encoded = [ + ( + entry, to_json( - x, - declared_type=child_declared_type(declared_type, i), - artifact_serializers=artifact_serializers, - for_hash=for_hash, - ) - for i, x in enumerate(obj) - ] - case tuple(): - return { - KINDMARKER: "tuple", - VALUEMARKER: [ - to_json( - x, - declared_type=child_declared_type(declared_type, i), - artifact_serializers=artifact_serializers, - for_hash=for_hash, - ) - for i, x in enumerate(obj) - ], - } - case set() | frozenset(): - element_type = child_declared_type(declared_type, 0) - return { - KINDMARKER: "set" if isinstance(obj, set) else "frozenset", - VALUEMARKER: sorted( - ( - to_json( - x, - declared_type=element_type, - artifact_serializers=artifact_serializers, - for_hash=for_hash, - ) - for x in obj + child, + declared_type=( + hints.get(cast(str, entry), Any) + if node.kind in ("dataclass", "pydantic") + else child_declared_type(declared_type, entry) ), - key=_stable_json_dump, - ), - } - case dict(): - return { - assert_correct_dict_key(k): to_json( - v, - declared_type=child_declared_type(declared_type, k), artifact_serializers=artifact_serializers, for_hash=for_hash, - ) - for k, v in obj.items() - } - case x if is_dataclass(x): - hints = get_type_hints(type(x), include_extras=True) - return { - KINDMARKER: "instance", - CLASSMARKER: fully_qualified_name(type(x)), - FIELDSMARKER: { - f.name: to_json( - getattr(x, f.name), - declared_type=hints.get(f.name, Any), - artifact_serializers=artifact_serializers, - for_hash=for_hash, - ) - for f in fields(x) - if not (for_hash and has_skip_hash(hints.get(f.name, Any))) - }, - } - case PydanticBaseModel(): - model_cls = type(obj) - model_fields = model_cls.model_fields - hints = get_type_hints(model_cls, include_extras=True) - return { - KINDMARKER: "instance", - CLASSMARKER: fully_qualified_name(model_cls), - FIELDSMARKER: { - k: to_json( - getattr(obj, k), - declared_type=hints.get(k, Any), - artifact_serializers=artifact_serializers, - for_hash=for_hash, - ) - for k in model_fields - if not (for_hash and has_skip_hash(hints.get(k, Any))) - }, - } - case enum.Enum(): - raise TypeError( - f"Cannot serialize enum value {obj!r}: enums are not supported " - "in furu artifacts yet" - ) - case _: - raise TypeError( - f"Cannot serialize value {obj!r} of type {type(obj).__name__!r} " - "into a furu artifact; register a Serializer for this type" + ), ) + for entry, child in entries_and_children + ] + match node.kind: + case "list": + return [child for _, child in encoded] + case "tuple": + return { + KINDMARKER: "tuple", + VALUEMARKER: [child for _, child in encoded], + } + case "set" | "frozenset": + return { + KINDMARKER: node.kind, + VALUEMARKER: sorted( + (child for _, child in encoded), key=_stable_json_dump + ), + } + case "dict": + return { + assert_correct_dict_key(key): child + for key, child in encoded + } + case "dataclass" | "pydantic": + cls = type(obj) + return { + KINDMARKER: "instance", + CLASSMARKER: fully_qualified_name(cls), + FIELDSMARKER: { + cast(str, name): child + for name, child in encoded + }, + } + + if isinstance(obj, enum.Enum): + raise TypeError( + f"Cannot serialize enum value {obj!r}: enums are not supported " + "in furu artifacts yet" + ) + raise TypeError( + f"Cannot serialize value {obj!r} of type {type(obj).__name__!r} " + "into a furu artifact; register a Serializer for this type" + ) def _from_json_field(value: JsonValue, expected_type: Any) -> Any: diff --git a/tests/test_core.py b/tests/test_core.py index a5c7712a..d93c8945 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -1419,13 +1419,12 @@ def test_top_level_create_accepts_single_spec_and_sequence() -> None: assert furu.create(nodes) == ["Node(top-create-a)", "Node(top-create-b)"] -def test_top_level_load_existing_rejects_single_furu_object() -> None: +def test_top_level_load_existing_accepts_single_furu_object_pytree() -> None: node = Node(name="single-load") assert node.create() == "Node(single-load)" - with pytest.raises(TypeError, match="expected a sequence of Spec objects"): - furu.load_existing(node) # ty:ignore[invalid-argument-type] + assert furu.load_existing(node) == "Node(single-load)" def test_top_level_load_existing_accepts_list_and_logs_once(tmp_path: Path) -> None: diff --git a/tests/test_pytree.py b/tests/test_pytree.py new file mode 100644 index 00000000..75cc5048 --- /dev/null +++ b/tests/test_pytree.py @@ -0,0 +1,82 @@ +from dataclasses import dataclass +from typing import cast + +import pytest +from pydantic import BaseModel + +import furu +from furu._pytree import tree_leaves, tree_map +from furu.worker.backends.local import LocalThreadWorkerBackend + + +class Number(furu.Spec[int]): + value: int + + def create(self) -> int: + return self.value * 10 + + +@dataclass(frozen=True) +class Pair: + left: object + right: object + + +class Box(BaseModel): + item: object + + +def test_tree_leaves_and_identity_map_preserve_the_tree() -> None: + tree = {"list": [1, Pair(2, 3)], "tuple": (4,), "set": {6, 5}} + + assert tree_leaves(tree) == [1, 2, 3, 4, 5, 6] + assert tree_map(lambda leaf: leaf, tree) == tree + + +def test_tree_map_is_composed_from_flatten_and_unflatten() -> None: + tree = Pair(left=[1, 2], right=Box(item=(3, 4))) + + mapped = tree_map(lambda value: cast(int, value) * 10, tree) + + assert mapped == Pair(left=[10, 20], right=Box(item=(30, 40))) + + +def test_create_preserves_an_arbitrary_pytree_shape() -> None: + tree = { + "pair": Pair(Number(value=1), Number(value=2)), + "box": Box(item=Number(value=3)), + "tuple": (Number(value=4),), + } + + created = furu.create(tree) + + assert created == { + "pair": Pair(10, 20), + "box": Box(item=30), + "tuple": (40,), + } + assert furu.load_existing(tree) == created + + +def test_create_treats_a_spec_as_a_leaf_not_as_its_dataclass_fields() -> None: + assert furu.create(Number(value=7)) == 70 + + +def test_create_on_worker_backend_accepts_pytree_roots() -> None: + tree = {"left": Number(value=8), "right": (Number(value=9),)} + + assert furu.create(tree, on=(LocalThreadWorkerBackend(),)) == { + "left": 80, + "right": (90,), + } + + +@pytest.mark.parametrize("tree", [{"bad": 1}, [Number(value=1), "bad"]]) +def test_create_rejects_non_spec_leaves(tree: object) -> None: + with pytest.raises(TypeError, match="PyTree of Spec objects"): + furu.create(tree) + + +@pytest.mark.parametrize("tree", [[], (), {}, Pair([], {})]) +def test_create_preserves_empty_pytrees(tree: object) -> None: + assert furu.create(tree) == tree