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
11 changes: 9 additions & 2 deletions src/griptape_nodes/exe_types/elements/parameter.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,7 +295,10 @@ def to_dict(self) -> dict[str, Any]:
return our_dict

def trait_states(self) -> list[dict[str, Any]]:
"""Return save-only trait identity and plain-data state."""
"""Return save-only trait identity and plain-data state.

``NodeManager`` stabilizes dynamic library module names before saving.
"""
states: list[dict[str, Any]] = []
for trait in self.find_elements_by_type(Trait):
trait_state: dict[str, Any] = {}
Expand All @@ -311,7 +314,11 @@ def trait_states(self) -> list[dict[str, Any]]:
)
continue
trait_state[key] = saved.value
entry = TraitStateEntry(trait_name=type(trait).__name__, trait_state=trait_state)
entry = TraitStateEntry(
trait_name=type(trait).__name__,
trait_module=type(trait).__module__,
trait_state=trait_state,
)
states.append(entry.to_dict())
return states

Expand Down
14 changes: 11 additions & 3 deletions src/griptape_nodes/exe_types/trait_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,9 +56,14 @@ def _as_saved_sequence(value: list | tuple | set | frozenset) -> SavedStateValue

@dataclass(frozen=True)
class TraitStateEntry:
"""Saved trait identity and state."""
"""Saved trait identity and state.

``trait_module`` distinguishes same-named traits from different libraries, and is what
finds the class when no instance is attached to carry the state.
"""

trait_name: str
trait_module: str | None = None
trait_state: dict[str, Any] = field(default_factory=dict)

@classmethod
Expand All @@ -67,10 +72,13 @@ def from_dict(cls, entry: dict[str, Any]) -> Self | None:
trait_name = entry.get("trait_name")
if not isinstance(trait_name, str):
return None
trait_module = entry.get("trait_module")
if not isinstance(trait_module, str):
trait_module = None
trait_state = entry.get("trait_state")
if not isinstance(trait_state, dict):
trait_state = {}
return cls(trait_name=trait_name, trait_state=trait_state)
return cls(trait_name=trait_name, trait_module=trait_module, trait_state=trait_state)

def to_dict(self) -> dict[str, Any]:
return {"trait_name": self.trait_name, "trait_state": self.trait_state}
return {"trait_name": self.trait_name, "trait_module": self.trait_module, "trait_state": self.trait_state}
146 changes: 116 additions & 30 deletions src/griptape_nodes/retained_mode/managers/node_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -244,6 +244,7 @@
)
from griptape_nodes.retained_mode.managers.library_manager import LibraryManager
from griptape_nodes.retained_mode.retained_mode import RetainedMode
from griptape_nodes.traits.trait_resolver import resolve_trait
from griptape_nodes.utils.exception_utils import readable_exception_message

logger = logging.getLogger("griptape_nodes")
Expand Down Expand Up @@ -4016,7 +4017,7 @@ def on_serialize_node_to_commands(self, request: SerializeNodeToCommandsRequest)
# Create the parameter, or alter it on the existing node
if parameter.user_defined:
# Always serialize user-defined parameters regardless of node type
add_param_request = AddParameterToNodeRequest.create(**parameter.save_dict(), initial_setup=True)
add_param_request = AddParameterToNodeRequest.create(**self._parameter_save_dict(parameter))
element_modification_commands.append(add_param_request)
elif isinstance(node, ErrorProxyNode):
# For ErrorProxyNode, replay all recorded initialization requests for this parameter
Expand All @@ -4032,7 +4033,7 @@ def on_serialize_node_to_commands(self, request: SerializeNodeToCommandsRequest)
element_modification_commands.extend(matching_requests)
elif reference_node is None:
# Normal node with no reference - treat all parameters as needing serialization
add_param_request = AddParameterToNodeRequest.create(**parameter.save_dict(), initial_setup=True)
add_param_request = AddParameterToNodeRequest.create(**self._parameter_save_dict(parameter))
element_modification_commands.append(add_param_request)
else:
# Normal node - compare against reference node
Expand All @@ -4045,6 +4046,8 @@ def on_serialize_node_to_commands(self, request: SerializeNodeToCommandsRequest)
if relevant:
diff["parameter_name"] = parameter.name
diff["initial_setup"] = True
if "traits" in diff:
diff["traits"] = self._stabilize_trait_modules(diff["traits"])
alter_param_request = AlterParameterDetailsRequest.create(**diff)
element_modification_commands.append(alter_param_request)

Expand All @@ -4062,6 +4065,8 @@ def on_serialize_node_to_commands(self, request: SerializeNodeToCommandsRequest)
if relevant:
diff["group_name"] = group.name
diff["initial_setup"] = True
if "traits" in diff:
diff["traits"] = self._stabilize_trait_modules(diff["traits"])
alter_group_request = AlterParameterGroupDetailsRequest(**diff)
element_modification_commands.append(alter_group_request)

Expand Down Expand Up @@ -4625,14 +4630,55 @@ def on_duplicate_selected_nodes(self, request: DuplicateSelectedNodesRequest) ->
result_details=f"Successfully duplicated {len(serialize_result.node_names_serialized)} nodes.",
)

def _parameter_save_dict(self, parameter: Parameter) -> dict[str, Any]:
"""Build the fields that recreate a parameter."""
param_dict = parameter.save_dict()
param_dict["initial_setup"] = True
param_dict["traits"] = self._stabilize_trait_modules(param_dict["traits"])
return param_dict

def _stabilize_trait_modules(self, trait_states: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Replace process-local library module names with stable namespaces."""
library_manager = self.engine.library_manager
for entry in trait_states:
trait_module = entry.get("trait_module")
if trait_module is None:
continue
if not library_manager.is_dynamic_module(trait_module):
continue
stable_namespace = library_manager.get_stable_namespace_for_dynamic_module(trait_module)
if stable_namespace is None:
entry["trait_module"] = None
logger.warning(
"Attempted to save the '%s' control, but the library providing it has no stable name to "
"record. Its state can only be restored when the node rebuilds that control.",
entry.get("trait_name"),
)
continue
entry["trait_module"] = stable_namespace
return trait_states

@staticmethod
def _apply_trait_states(parameter: Parameter, trait_states: list[dict[str, Any]]) -> None:
"""Hand saved state to the trait the node's own code built.
"""Hand saved state to the trait the node built, or build one it did not.

Updating the attached instance rather than replacing it is what keeps a callback the
node's ``__init__`` supplied, along with everything else the constructor wired.
"""
unmatched = parameter.find_elements_by_type(Trait)
entries = NodeManager._parse_trait_entries(parameter, trait_states)
paired = NodeManager._pair_saved_traits(parameter, entries)
for entry, existing in zip(entries, paired, strict=True):
if existing is None:
NodeManager._build_saved_trait(parameter, entry)
continue
try:
existing.apply_state(entry.trait_state)
except (TypeError, ValueError):
NodeManager._warn_unsatisfiable_trait_state(parameter, entry.trait_name)

@staticmethod
def _parse_trait_entries(parameter: Parameter, trait_states: list[dict[str, Any]]) -> list[TraitStateEntry]:
entries: list[TraitStateEntry] = []
for state in trait_states:
entry = TraitStateEntry.from_dict(state)
if entry is None:
Expand All @@ -4642,36 +4688,76 @@ def _apply_trait_states(parameter: Parameter, trait_states: list[dict[str, Any]]
parameter.name,
)
continue
trait = NodeManager._take_trait_named(unmatched, entry.trait_name)
if trait is None:
logger.warning(
"Parameter '%s' was saved with a '%s' control, but nothing on this node builds one, "
"so the parameter loads without it. This usually means the node's library changed.",
parameter.name,
entry.trait_name,
)
entries.append(entry)
return entries

@staticmethod
def _pair_saved_traits(parameter: Parameter, entries: list[TraitStateEntry]) -> list[Trait | None]:
"""Match by resolved class, consuming each attached trait at most once."""
unmatched = parameter.find_elements_by_type(Trait)
paired: list[Trait | None] = []
for entry in entries:
trait_class = NodeManager._resolve_saved_trait(entry)
if trait_class is None and entry.trait_module is not None:
paired.append(None)
continue
try:
trait.apply_state(entry.trait_state)
except (TypeError, ValueError):
logger.warning(
"Parameter '%s' was saved with state for its '%s' control, but the control did not "
"accept it. The parameter keeps the state its node supplied.",
parameter.name,
entry.trait_name,
)

match = None
for candidate in unmatched:
matches_resolved_class = trait_class is not None and type(candidate) is trait_class
matches_legacy_name = trait_class is None and type(candidate).__name__ == entry.trait_name
if matches_resolved_class or matches_legacy_name:
match = candidate
break
if match is not None:
unmatched.remove(match)
paired.append(match)
return paired

@staticmethod
def _take_trait_named(unmatched: list[Trait], trait_name: str) -> Trait | None:
"""Take an attached trait by class name, consuming it so two entries cannot share one.
def _build_saved_trait(parameter: Parameter, entry: TraitStateEntry) -> Trait | None:
"""Construct a saved trait no attached instance accounts for."""
if entry.trait_module is None:
logger.warning(
"Parameter '%s' was saved with a '%s' control, but no module was recorded and the node did "
"not rebuild it. The parameter loads without that control.",
parameter.name,
entry.trait_name,
)
return None

A name is enough because the candidates are the traits already on this one parameter.
"""
for candidate in unmatched:
if type(candidate).__name__ == trait_name:
unmatched.remove(candidate)
return candidate
return None
trait_class = NodeManager._resolve_saved_trait(entry)
if trait_class is None:
logger.warning(
"Parameter '%s' was saved with the '%s' trait from '%s', but that trait could not be loaded. "
"The parameter will load without it. Check that the library providing it is installed.",
parameter.name,
entry.trait_name,
entry.trait_module,
)
return None
try:
trait = trait_class(**entry.trait_state)
except (TypeError, ValueError):
NodeManager._warn_unsatisfiable_trait_state(parameter, entry.trait_name)
return None
parameter.add_trait(trait)
return trait

@staticmethod
def _resolve_saved_trait(entry: TraitStateEntry) -> type[Trait] | None:
if entry.trait_module is None:
return None
return resolve_trait(entry.trait_name, entry.trait_module)

@staticmethod
def _warn_unsatisfiable_trait_state(parameter: Parameter, trait_name: str) -> None:
logger.warning(
"Parameter '%s' was saved with the '%s' trait, but its saved state could not build that control. "
"The parameter loads without it. Check that the library providing it is up to date.",
parameter.name,
trait_name,
)

@staticmethod
def _manage_alter_details(parameter: Parameter, base_node_obj: BaseNode) -> dict:
Expand Down
50 changes: 50 additions & 0 deletions src/griptape_nodes/traits/trait_resolver.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
from __future__ import annotations

import importlib
import logging
import sys
from typing import TYPE_CHECKING

from griptape_nodes.exe_types.core_types import Trait

if TYPE_CHECKING:
from types import ModuleType

logger = logging.getLogger("griptape_nodes")


def resolve_trait(trait_name: str, trait_module: str) -> type[Trait] | None:
"""Resolve a saved trait by module and class name."""
module = sys.modules.get(trait_module)
if module is None:
module = _import_trait_module(trait_module, trait_name)
if module is None:
return None
candidate = getattr(module, trait_name, None)
if not isinstance(candidate, type) or not issubclass(candidate, Trait):
return None
return candidate


def _import_trait_module(trait_module: str, trait_name: str) -> ModuleType | None:
"""Import a trait module without failing the workflow load."""
try:
return importlib.import_module(trait_module)
except Exception as error:
if isinstance(error, ModuleNotFoundError) and _requested_module_is_missing(trait_module, error.name):
return None
logger.warning(
"Attempted to restore the '%s' trait from '%s'. Loading that module failed (%s), "
"so the parameter will load without that trait. This usually means the library "
"providing it is broken or partly installed.",
trait_name,
trait_module,
error,
)
return None


def _requested_module_is_missing(trait_module: str, missing_module: str | None) -> bool:
if missing_module is None:
return False
return trait_module == missing_module or trait_module.startswith(f"{missing_module}.")
1 change: 1 addition & 0 deletions tests/unit/exe_types/test_core_types_serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -353,6 +353,7 @@ def test_a_dynamic_choices_update_is_carried_by_trait_state(self) -> None:
assert param.trait_states() == [
{
"trait_name": "Options",
"trait_module": "griptape_nodes.traits.options",
"trait_state": {
"choices": ["updated 1", "updated 2"],
"show_search": True,
Expand Down
12 changes: 10 additions & 2 deletions tests/unit/exe_types/test_trait_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,25 +64,33 @@ class TestTheSavedShape:
def test_a_full_entry_round_trips(self) -> None:
entry = TraitStateEntry(
trait_name="Button",
trait_module="griptape_nodes.traits.button",
trait_state={"label": "Refresh"},
)

assert TraitStateEntry.from_dict(entry.to_dict()) == entry

def test_an_entry_with_no_state_still_names_its_trait(self) -> None:
entry = TraitStateEntry(trait_name="Options")
entry = TraitStateEntry(trait_name="Options", trait_module="griptape_nodes.traits.options")

assert entry.to_dict() == {
"trait_name": "Options",
"trait_module": "griptape_nodes.traits.options",
"trait_state": {},
}

def test_an_entry_naming_no_trait_reads_as_nothing(self) -> None:
assert TraitStateEntry.from_dict({"trait_state": {"label": "Refresh"}}) is None

def test_a_missing_module_reads_as_absent_rather_than_guessed(self) -> None:
entry = TraitStateEntry.from_dict({"trait_name": "Options"})

assert entry is not None
assert entry.trait_module is None

def test_a_missing_state_reads_as_empty(self) -> None:
"""An entry that says nothing about state asks for whatever the node's code builds."""
entry = TraitStateEntry.from_dict({"trait_name": "Options"})
entry = TraitStateEntry.from_dict({"trait_name": "Options", "trait_module": "m"})

assert entry is not None
assert entry.trait_state == {}
Loading
Loading