diff --git a/src/griptape_nodes/exe_types/elements/parameter.py b/src/griptape_nodes/exe_types/elements/parameter.py index 6bf12271ca..cc80457b05 100644 --- a/src/griptape_nodes/exe_types/elements/parameter.py +++ b/src/griptape_nodes/exe_types/elements/parameter.py @@ -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] = {} @@ -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 diff --git a/src/griptape_nodes/exe_types/trait_state.py b/src/griptape_nodes/exe_types/trait_state.py index 43ff0ed524..ddaa5facc9 100644 --- a/src/griptape_nodes/exe_types/trait_state.py +++ b/src/griptape_nodes/exe_types/trait_state.py @@ -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 @@ -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} diff --git a/src/griptape_nodes/retained_mode/managers/node_manager.py b/src/griptape_nodes/retained_mode/managers/node_manager.py index 91d2cdb926..f6cb6090d2 100644 --- a/src/griptape_nodes/retained_mode/managers/node_manager.py +++ b/src/griptape_nodes/retained_mode/managers/node_manager.py @@ -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") @@ -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 @@ -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 @@ -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) @@ -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) @@ -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: @@ -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: diff --git a/src/griptape_nodes/traits/trait_resolver.py b/src/griptape_nodes/traits/trait_resolver.py new file mode 100644 index 0000000000..c740bbbbfd --- /dev/null +++ b/src/griptape_nodes/traits/trait_resolver.py @@ -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}.") diff --git a/tests/unit/exe_types/test_core_types_serialization.py b/tests/unit/exe_types/test_core_types_serialization.py index 0395aafac9..12fedcbe31 100644 --- a/tests/unit/exe_types/test_core_types_serialization.py +++ b/tests/unit/exe_types/test_core_types_serialization.py @@ -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, diff --git a/tests/unit/exe_types/test_trait_state.py b/tests/unit/exe_types/test_trait_state.py index a276e40e7e..78b45f1a75 100644 --- a/tests/unit/exe_types/test_trait_state.py +++ b/tests/unit/exe_types/test_trait_state.py @@ -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 == {} diff --git a/tests/unit/retained_mode/managers/test_trait_state_pairing.py b/tests/unit/retained_mode/managers/test_trait_state_pairing.py index 4b5053301d..b07e5afa89 100644 --- a/tests/unit/retained_mode/managers/test_trait_state_pairing.py +++ b/tests/unit/retained_mode/managers/test_trait_state_pairing.py @@ -42,6 +42,7 @@ class Ranged(Trait): def __init__(self, level: int = 1) -> None: super().__init__() + self._validate_level(level) self.level = level @classmethod @@ -55,10 +56,14 @@ def apply_state(self, state: dict) -> None: if "level" not in state: return level = state["level"] + self._validate_level(level) + self.level = level + + @staticmethod + def _validate_level(level: int) -> None: if level not in (1, _ALLOWED_LEVEL, 3): msg = "level must be between 1 and 3" raise ValueError(msg) - self.level = level def ui_options_for_trait(self) -> dict: return {} @@ -88,10 +93,12 @@ def test_a_trait_is_paired_with_the_entry_that_names_its_class(self) -> None: [ { "trait_name": "Slider", + "trait_module": "griptape_nodes.traits.slider", "trait_state": {"min_val": 2, "max_val": 8}, }, { "trait_name": "Options", + "trait_module": "griptape_nodes.traits.options", "trait_state": {"choices": ["b"]}, }, ], @@ -101,6 +108,33 @@ def test_a_trait_is_paired_with_the_entry_that_names_its_class(self) -> None: assert (slider.min, slider.max) == (2, 8) assert parameter.find_elements_by_type(Options)[0].choices == ["b"] + def test_an_entry_without_a_module_still_reaches_an_attached_trait(self) -> None: + slider = Slider(min_val=0, max_val=1) + parameter = Parameter(name="p", tooltip="t", traits={slider}) + + NodeManager._apply_trait_states( + parameter, + [{"trait_name": "Slider", "trait_state": {"min_val": 2, "max_val": 8}}], + ) + + assert (slider.min, slider.max) == (2, 8) + + def test_a_same_named_trait_from_another_library_is_not_mistaken_for_it(self, foreign_twin: ModuleType) -> None: + local = Twin(tag="local") + parameter = Parameter(name="p", tooltip="t", traits={local}) + + # The saved entry names the other library's Twin, so it describes a trait this + # parameter does not carry. Matching on the name alone would overwrite the local one. + NodeManager._apply_trait_states( + parameter, + [{"trait_name": "Twin", "trait_module": foreign_twin.__name__, "trait_state": {"tag": "foreign"}}], + ) + + assert local.tag == "local" + attached = parameter.find_elements_by_type(Twin) + assert len(attached) == 2 # noqa: PLR2004 + assert {type(trait).__module__ for trait in attached} == {Twin.__module__, foreign_twin.__name__} + class TestPairingIsOneToOne: """A parameter can carry two traits of one class, and each entry describes one of them.""" @@ -116,10 +150,12 @@ def test_both_traits_of_a_class_get_their_own_state(self) -> None: [ { "trait_name": "Options", + "trait_module": "griptape_nodes.traits.options", "trait_state": {"choices": ["first"]}, }, { "trait_name": "Options", + "trait_module": "griptape_nodes.traits.options", "trait_state": {"choices": ["second"]}, }, ], @@ -133,8 +169,8 @@ def test_no_extra_trait_is_built_for_the_second_entry(self) -> None: NodeManager._apply_trait_states( parameter, [ - {"trait_name": "Options", "trait_state": {}}, - {"trait_name": "Options", "trait_state": {}}, + {"trait_name": "Options", "trait_module": "griptape_nodes.traits.options", "trait_state": {}}, + {"trait_name": "Options", "trait_module": "griptape_nodes.traits.options", "trait_state": {}}, ], ) @@ -148,7 +184,7 @@ def test_an_existing_trait_keeps_values_the_state_omits(self) -> None: NodeManager._apply_trait_states( parameter, - [{"trait_name": "Slider", "trait_state": {"min_val": 2}}], + [{"trait_name": "Slider", "trait_module": "griptape_nodes.traits.slider", "trait_state": {"min_val": 2}}], ) assert (slider.min, slider.max) == (2, 1) @@ -167,7 +203,7 @@ def test_an_existing_trait_keeps_what_init_built(self) -> None: NodeManager._apply_trait_states( parameter, - [{"trait_name": "Ranged", "trait_state": {"level": 99}}], + [{"trait_name": "Ranged", "trait_module": __name__, "trait_state": {"level": 99}}], ) assert ranged.level == _ALLOWED_LEVEL @@ -178,7 +214,7 @@ def test_a_warning_names_the_trait_when_one_is_already_attached(self, caplog: py NodeManager._apply_trait_states( parameter, - [{"trait_name": "Ranged", "trait_state": {"level": 99}}], + [{"trait_name": "Ranged", "trait_module": __name__, "trait_state": {"level": 99}}], ) assert any("Ranged" in record.getMessage() for record in caplog.records) @@ -188,7 +224,7 @@ def test_no_trait_is_built_when_none_is_already_attached(self) -> None: NodeManager._apply_trait_states( parameter, - [{"trait_name": "Ranged", "trait_state": {"level": 99}}], + [{"trait_name": "Ranged", "trait_module": __name__, "trait_state": {"level": 99}}], ) assert parameter.find_elements_by_type(Ranged) == [] @@ -199,7 +235,7 @@ def test_a_warning_names_the_trait_when_none_is_already_attached(self, caplog: p NodeManager._apply_trait_states( parameter, - [{"trait_name": "Ranged", "trait_state": {"level": 99}}], + [{"trait_name": "Ranged", "trait_module": __name__, "trait_state": {"level": 99}}], ) assert any("Ranged" in record.getMessage() for record in caplog.records) diff --git a/tests/unit/retained_mode/managers/test_trait_state_serialization.py b/tests/unit/retained_mode/managers/test_trait_state_serialization.py index 2c72613fc2..cf938e12e0 100644 --- a/tests/unit/retained_mode/managers/test_trait_state_serialization.py +++ b/tests/unit/retained_mode/managers/test_trait_state_serialization.py @@ -170,6 +170,7 @@ def test_emitted_commands_carry_trait_state(self, engine: Engine) -> None: assert model_command.traits == [ { "trait_name": "Options", + "trait_module": "griptape_nodes.traits.options", "trait_state": { "choices": ["sdxl", "sd3", "flux"], "show_search": True, @@ -183,6 +184,7 @@ def test_emitted_commands_carry_trait_state(self, engine: Engine) -> None: assert reload_command.traits == [ { "trait_name": "Button", + "trait_module": "griptape_nodes.traits.button", "trait_state": { "label": "Reload", "variant": "secondary", @@ -201,6 +203,77 @@ def test_emitted_commands_carry_trait_state(self, engine: Engine) -> None: } ] + def test_a_dynamic_trait_module_is_saved_by_stable_name( + self, engine: Engine, monkeypatch: pytest.MonkeyPatch + ) -> None: + node = _add_node(engine, "picker") + node.discover() + monkeypatch.setattr(Options, "__module__", "gtn_dynamic_module_options_test") + monkeypatch.setattr( + engine.library_manager, + "get_stable_namespace_for_dynamic_module", + lambda _module: "stable_library.options", + ) + + result = engine.node_manager.on_serialize_node_to_commands(SerializeNodeToCommandsRequest(node_name=node.name)) + + assert isinstance(result, SerializeNodeToCommandsResultSuccess) + commands = _added_parameter_commands(result.serialized_node_commands.element_modification_commands) + traits = commands["model"].traits + assert traits is not None + assert traits[0]["trait_module"] == "stable_library.options" + + def test_a_dynamic_trait_without_a_stable_name_omits_its_module( + self, engine: Engine, monkeypatch: pytest.MonkeyPatch + ) -> None: + node = _add_node(engine, "picker") + node.discover() + monkeypatch.setattr(Options, "__module__", "gtn_dynamic_module_options_test") + monkeypatch.setattr(engine.library_manager, "get_stable_namespace_for_dynamic_module", lambda _module: None) + + result = engine.node_manager.on_serialize_node_to_commands(SerializeNodeToCommandsRequest(node_name=node.name)) + + assert isinstance(result, SerializeNodeToCommandsResultSuccess) + commands = _added_parameter_commands(result.serialized_node_commands.element_modification_commands) + traits = commands["model"].traits + assert traits is not None + assert traits[0]["trait_module"] is None + + def test_replaying_the_commands_restores_state(self, engine: Engine) -> None: + node = _add_node(engine, "picker") + node.discover() + node.set_parameter_value("model", "flux") + + result = engine.node_manager.on_serialize_node_to_commands(SerializeNodeToCommandsRequest(node_name=node.name)) + assert isinstance(result, SerializeNodeToCommandsResultSuccess) + commands = _added_parameter_commands(result.serialized_node_commands.element_modification_commands) + + target = _add_node(engine, "reloaded") + for command in commands.values(): + command.node_name = target.name + command.initial_setup = True + replay_result = engine.handle_request(command) + assert isinstance(replay_result, AddParameterToNodeResultSuccess) + + model = target.get_parameter_by_name("model") + assert model is not None + assert model.ui_options["simple_dropdown"] == ["sdxl", "sd3", "flux"] + converted = "not-a-model" + for converter in model.converters: + converted = converter(converted) + assert converted == "sdxl" # Options snaps an invalid value to its first choice. + + # A replayed command carries state, not behavior. Behavior comes from the node's own + # code, which is why a declared parameter's button still fires: its trait is built by + # __init__ and only updated from the save. A bare replay onto a node whose code never + # built this button has nothing to supply the handler. + reload_button = next( + trait + for trait in target.get_parameter_by_name("reload").find_elements_by_type(Button) # type: ignore[union-attr] + ) + assert reload_button.label == "Reload" + assert reload_button.on_click_callback is None + class _UnsaveableValueTrait(Trait): """Stands in for a trait returning something no saved workflow can express.""" @@ -273,6 +346,7 @@ def test_serializing_still_succeeds(self, engine: Engine, caplog: pytest.LogCapt assert broken_command.traits == [ { "trait_name": "_UnsaveableValueTrait", + "trait_module": _UnsaveableValueTrait.__module__, "trait_state": {}, } ] diff --git a/tests/unit/traits/test_trait_resolver.py b/tests/unit/traits/test_trait_resolver.py new file mode 100644 index 0000000000..e8f74bec10 --- /dev/null +++ b/tests/unit/traits/test_trait_resolver.py @@ -0,0 +1,123 @@ +"""resolve_trait finds a saved trait through its recorded module.""" + +from __future__ import annotations + +import logging +import sys +from types import ModuleType +from typing import TYPE_CHECKING + +import pytest + +from griptape_nodes.exe_types.core_types import Trait +from griptape_nodes.traits.slider import Slider +from griptape_nodes.traits.trait_resolver import resolve_trait + +if TYPE_CHECKING: + from collections.abc import Generator + from pathlib import Path + + +class _CollisionTraitA(Trait): + """One of two distinct trait classes renamed below to share a class name.""" + + def __init__(self) -> None: + super().__init__(element_id="_CollisionTraitA") + + def ui_options_for_trait(self) -> dict: + return {} + + +class _CollisionTraitB(Trait): + """The other of the two, standing in for a second library's trait of the same name.""" + + def __init__(self) -> None: + super().__init__(element_id="_CollisionTraitB") + + def ui_options_for_trait(self) -> dict: + return {} + + +# Renamed to collide: both now report the class name "SharedTraitName", from different +# modules, the way two libraries shipping a same-named trait would. +_CollisionTraitA.__name__ = "SharedTraitName" +_CollisionTraitA.__qualname__ = "SharedTraitName" +_CollisionTraitA.__module__ = "tests.fake.collision_module_a" +_CollisionTraitB.__name__ = "SharedTraitName" +_CollisionTraitB.__qualname__ = "SharedTraitName" +_CollisionTraitB.__module__ = "tests.fake.collision_module_b" + + +@pytest.fixture(autouse=True) +def _fake_collision_modules() -> Generator[None]: + """Register the two fake modules the collision classes claim to live in.""" + module_a = ModuleType(_CollisionTraitA.__module__) + setattr(module_a, "SharedTraitName", _CollisionTraitA) # noqa: B010 + module_b = ModuleType(_CollisionTraitB.__module__) + setattr(module_b, "SharedTraitName", _CollisionTraitB) # noqa: B010 + sys.modules[_CollisionTraitA.__module__] = module_a + sys.modules[_CollisionTraitB.__module__] = module_b + yield + sys.modules.pop(_CollisionTraitA.__module__, None) + sys.modules.pop(_CollisionTraitB.__module__, None) + + +class TestResolveByModule: + def test_resolves_an_in_tree_trait_through_its_module(self) -> None: + resolved = resolve_trait("Slider", "griptape_nodes.traits.slider") + + assert resolved is Slider + + def test_the_module_disambiguates_two_classes_sharing_a_name(self) -> None: + assert resolve_trait("SharedTraitName", _CollisionTraitA.__module__) is _CollisionTraitA + assert resolve_trait("SharedTraitName", _CollisionTraitB.__module__) is _CollisionTraitB + + +class TestABrokenTraitModuleDoesNotFailTheLoad: + """Resolving executes a library's module code, which must not take the workflow down.""" + + def test_a_module_that_raises_on_import_warns_instead_of_propagating( + self, caplog: pytest.LogCaptureFixture, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + (tmp_path / "exploding_trait_module.py").write_text('raise RuntimeError("library blew up on import")') + monkeypatch.syspath_prepend(str(tmp_path)) + + with caplog.at_level(logging.WARNING, logger="griptape_nodes"): + resolved = resolve_trait("Slider", "exploding_trait_module") + + assert resolved is None + assert "library blew up on import" in caplog.text + + def test_a_missing_dependency_is_reported( + self, caplog: pytest.LogCaptureFixture, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + module_name = "trait_module_with_missing_dependency" + dependency_name = "dependency_that_does_not_exist" + (tmp_path / f"{module_name}.py").write_text(f"import {dependency_name}\n") + monkeypatch.syspath_prepend(str(tmp_path)) + + with caplog.at_level(logging.WARNING, logger="griptape_nodes"): + resolved = resolve_trait("Slider", module_name) + + assert resolved is None + assert dependency_name in caplog.text + + def test_a_missing_module_resolves_to_nothing_quietly(self, caplog: pytest.LogCaptureFixture) -> None: + """The caller reports the trait it could not restore, so this stays quiet.""" + with caplog.at_level(logging.WARNING, logger="griptape_nodes"): + resolved = resolve_trait("Slider", "griptape_nodes.traits.never_existed") + + assert resolved is None + assert caplog.text == "" + + +class TestUnresolvableTraitDegradesInsteadOfRaising: + def test_a_name_the_module_does_not_hold_returns_none(self) -> None: + resolved = resolve_trait("NoSuchTraitAnywhere", "griptape_nodes.traits.slider") + + assert resolved is None + + def test_a_name_that_holds_something_other_than_a_trait_returns_none(self) -> None: + resolved = resolve_trait("Parameter", "griptape_nodes.exe_types.core_types") + + assert resolved is None