diff --git a/pyproject.toml b/pyproject.toml index 067a940..c86e1d3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "retrack" -version = "3.5.0" +version = "3.6.0" description = "A business rules engine" authors = ["Gabriel Guarisa "] license = "MIT" diff --git a/retrack/nodes/dynamic/__init__.py b/retrack/nodes/dynamic/__init__.py index 3af5565..37cb03b 100644 --- a/retrack/nodes/dynamic/__init__.py +++ b/retrack/nodes/dynamic/__init__.py @@ -4,6 +4,7 @@ from retrack.nodes.dynamic.flow_connector import flow_connector_factory from retrack.nodes.dynamic.conditional_connector import conditional_connector_factory from retrack.nodes.dynamic.glm import glm_factory +from retrack.nodes.dynamic.multiple_outputs import multiple_outputs_factory from retrack.utils.registry import Registry @@ -21,6 +22,7 @@ def registry() -> Registry: _registry.register("ModelConnector", conditional_connector_factory) _registry.register("FeatureConnector", conditional_connector_factory) _registry.register("GLM", glm_factory) + _registry.register("MultipleOutputs", multiple_outputs_factory) return _registry diff --git a/retrack/nodes/dynamic/multiple_outputs.py b/retrack/nodes/dynamic/multiple_outputs.py new file mode 100644 index 0000000..b2689e3 --- /dev/null +++ b/retrack/nodes/dynamic/multiple_outputs.py @@ -0,0 +1,65 @@ +import typing + +import pandas as pd +import pydantic + +from retrack.nodes.base import ( + InputConnectionModel, + NodeKind, + OptionalCastedToNoneStringType, +) +from retrack.nodes.dynamic.base import BaseDynamicIOModel, BaseDynamicNode +from retrack.utils import constants + + +class MultipleOutputsMetadataModel(pydantic.BaseModel): + headers_map: typing.List[str] + message: OptionalCastedToNoneStringType = None + name: typing.Optional[str] = None + + @pydantic.field_validator("headers_map") + @classmethod + def validate_headers_map(cls, v: typing.List[str]) -> typing.List[str]: + if not v: + raise ValueError("headers_map must not be empty") + if len(v) != len(set(v)): + raise ValueError("headers_map must not contain duplicate keys") + return v + + +def multiple_outputs_factory( + inputs: typing.Dict[str, typing.Any], **kwargs +) -> typing.Type[BaseDynamicNode]: + input_fields = { + name: BaseDynamicNode.create_sub_field(InputConnectionModel) + for name in inputs.keys() + } + inputs_model = BaseDynamicIOModel.with_fields( + "MultipleOutputsInputsModel", **input_fields + ) + + models = { + "inputs": BaseDynamicNode.create_sub_field(inputs_model), + "data": BaseDynamicNode.create_sub_field(MultipleOutputsMetadataModel), + } + + BaseModel = BaseDynamicNode.with_fields("MultipleOutputsBaseModel", **models) + + class MultipleOutputs(BaseModel): + def kind(self) -> NodeKind: + return NodeKind.OUTPUT + + async def run(self, **kwargs) -> typing.Dict[str, typing.Any]: + keys = self.data.headers_map + rows = zip(*[kwargs[k].values for k in keys]) + output_series = pd.Series( + [[{"key": k, "value": v} for k, v in zip(keys, row)] for row in rows], + index=kwargs[keys[0]].index, + dtype=object, + ) + return { + constants.OUTPUT_REFERENCE_COLUMN: output_series, + constants.OUTPUT_MESSAGE_REFERENCE_COLUMN: self.data.message, + } + + return MultipleOutputs diff --git a/retrack/utils/constants.py b/retrack/utils/constants.py index c27a0a6..6f7943d 100644 --- a/retrack/utils/constants.py +++ b/retrack/utils/constants.py @@ -4,4 +4,4 @@ FILTER_SUFFIX = "_filter" INPUT_OUTPUT_VALUE_CONNECTOR_NAME = "output_value" FLOW_NODE_NAME = "flowv0" -EXCLUDED_NODE_TYPES = {"Input", "Output", "Start"} +EXCLUDED_NODE_TYPES = {"Input", "Output", "MultipleOutputs", "Start"} diff --git a/retrack/utils/transformers.py b/retrack/utils/transformers.py index 1223b98..6c050bf 100644 --- a/retrack/utils/transformers.py +++ b/retrack/utils/transformers.py @@ -7,6 +7,9 @@ from retrack.utils.constants import EXCLUDED_NODE_TYPES, FILTER_SUFFIX, NULL_SUFFIX +_TERMINAL_NODE_TYPES = frozenset({"Output", "MultipleOutputs"}) + + def is_excluded_node(node_type: str) -> bool: return node_type in EXCLUDED_NODE_TYPES @@ -251,25 +254,44 @@ def normalize_execution_for_debug_iter( outputs = [] for node in nodes_at_index: - if node.get("type") == "Output": - inputs_list = node.get("inputs", []) - node_name = node.get("name") - if inputs_list: - first_input = inputs_list[0] - value = first_input.get("value") - - message = None - for item in node.get("data", []): - if item.get("name") == "message": - message = item.get("value") - break - - if value is not None and not ( - isinstance(value, float) and pd.isna(value) + node_type = node.get("type") + if node_type not in _TERMINAL_NODE_TYPES: + continue + + inputs_list = node.get("inputs", []) + message = None + for item in node.get("data", []): + if item.get("name") == "message": + message = item.get("value") + break + + if node_type == "Output": + if not inputs_list: + continue + value = inputs_list[0].get("value") + if value is not None and not ( + isinstance(value, float) and pd.isna(value) + ): + outputs.append( + { + "name": "output", + "value": value, + "message": message, + } + ) + + elif node_type == "MultipleOutputs": + for inp in inputs_list: + key = inp.get("target_name") + value = inp.get("value") + if ( + key + and value is not None + and not (isinstance(value, float) and pd.isna(value)) ): outputs.append( { - "name": "output", + "name": key, "value": value, "message": message, } diff --git a/retrack/validators/__init__.py b/retrack/validators/__init__.py index f39e52c..f9a5c6b 100644 --- a/retrack/validators/__init__.py +++ b/retrack/validators/__init__.py @@ -3,6 +3,7 @@ from retrack.validators.check_is_dag import CheckIsDAG from retrack.validators.node_exists import NodeExistsValidator from retrack.validators.node_validator import IntervalCatV0Validator +from retrack.validators.single_terminal_node import SingleTerminalNodeValidator def registry() -> Registry: @@ -13,6 +14,7 @@ def registry() -> Registry: ) _registry.register("check_is_dag", CheckIsDAG()) _registry.register("interval_cat_v0", IntervalCatV0Validator()) + _registry.register("single_terminal_node", SingleTerminalNodeValidator()) return _registry diff --git a/retrack/validators/single_terminal_node.py b/retrack/validators/single_terminal_node.py new file mode 100644 index 0000000..05a59d9 --- /dev/null +++ b/retrack/validators/single_terminal_node.py @@ -0,0 +1,30 @@ +from typing import Optional + +from retrack.validators.base import BaseValidator + + +class SingleTerminalNodeValidator(BaseValidator): + """Prevents mixing Output and MultipleOutputs nodes in the same graph. + + Multiple Output nodes are valid (used in conditional branches). Multiple + MultipleOutputs nodes are also valid. What is not valid is having both + types simultaneously, because their output contracts are incompatible. + """ + + def validate(self, graph_data: dict, **kwargs) -> tuple[bool, Optional[str]]: + nodes = graph_data.get("nodes", {}) + + has_output = any( + node.get("name", "").lower() == "output" for _, node in nodes.items() + ) + has_multiple_outputs = any( + node.get("name", "").lower() == "multipleoutputs" + for _, node in nodes.items() + ) + + if has_output and has_multiple_outputs: + return ( + False, + "Graph cannot mix Output and MultipleOutputs nodes in the same flow", + ) + return True, None diff --git a/tests/resources/multiple-outputs.json b/tests/resources/multiple-outputs.json new file mode 100644 index 0000000..e86cb65 --- /dev/null +++ b/tests/resources/multiple-outputs.json @@ -0,0 +1,82 @@ +{ + "nodes": { + "1": { + "id": 1, + "data": {}, + "inputs": {}, + "outputs": { + "output_up_void": { + "connections": [ + {"node": 2, "input": "input_void", "data": {}}, + {"node": 4, "input": "input_void", "data": {}} + ] + }, + "output_down_void": { + "connections": [ + {"node": 3, "input": "input_void", "data": {}} + ] + } + }, + "name": "Start" + }, + "2": { + "id": 2, + "data": {"name": "value_a", "default": null}, + "inputs": { + "input_void": { + "connections": [{"node": 1, "output": "output_up_void", "data": {}}] + } + }, + "outputs": { + "output_value": { + "connections": [{"node": 10, "input": "basic_7", "data": {}}] + } + }, + "name": "Input" + }, + "3": { + "id": 3, + "data": {"name": "value_b", "default": null}, + "inputs": { + "input_void": { + "connections": [{"node": 1, "output": "output_down_void", "data": {}}] + } + }, + "outputs": { + "output_value": { + "connections": [{"node": 10, "input": "basic_15", "data": {}}] + } + }, + "name": "Input" + }, + "4": { + "id": 4, + "data": {"name": "value_c", "default": null}, + "inputs": { + "input_void": { + "connections": [{"node": 1, "output": "output_up_void", "data": {}}] + } + }, + "outputs": { + "output_value": { + "connections": [{"node": 10, "input": "basic_30", "data": {}}] + } + }, + "name": "Input" + }, + "10": { + "id": 10, + "data": { + "headers_map": ["basic_7", "basic_15", "basic_30"], + "message": "cobertura" + }, + "inputs": { + "basic_7": {"connections": [{"node": 2, "output": "output_value", "data": {}}]}, + "basic_15": {"connections": [{"node": 3, "output": "output_value", "data": {}}]}, + "basic_30": {"connections": [{"node": 4, "output": "output_value", "data": {}}]} + }, + "outputs": {}, + "name": "MultipleOutputs" + } + } +} diff --git a/tests/test_engine/test_executor.py b/tests/test_engine/test_executor.py index 0202d0f..8ac662d 100644 --- a/tests/test_engine/test_executor.py +++ b/tests/test_engine/test_executor.py @@ -282,6 +282,99 @@ def test_create_from_json_with_invalid_type(): from_json(1) +@pytest.mark.asyncio +async def test_multiple_outputs_single_row(): + with open("tests/resources/multiple-outputs.json", "r") as f: + graph_data = json.load(f) + + executor = Rule.create( + graph_data, + nodes_registry=nodes.registry(), + dynamic_nodes_registry=nodes.dynamic_nodes_registry(), + ).executor + + result = await executor.execute( + pd.DataFrame([{"value_a": "3028", "value_b": "7194", "value_c": "15720"}]) + ) + + assert isinstance(result, pd.DataFrame) + records = result.to_dict(orient="records") + assert len(records) == 1 + assert records[0]["message"] == "cobertura" + output = records[0]["output"] + assert isinstance(output, list) + assert len(output) == 3 + assert output[0]["key"] == "basic_7" + assert output[1]["key"] == "basic_15" + assert output[2]["key"] == "basic_30" + assert output[0]["value"] == "3028" + assert output[1]["value"] == "7194" + assert output[2]["value"] == "15720" + + +@pytest.mark.asyncio +async def test_multiple_outputs_batch(): + with open("tests/resources/multiple-outputs.json", "r") as f: + graph_data = json.load(f) + + executor = Rule.create( + graph_data, + nodes_registry=nodes.registry(), + dynamic_nodes_registry=nodes.dynamic_nodes_registry(), + ).executor + + payload = pd.DataFrame( + [ + {"value_a": "100", "value_b": "200", "value_c": "300"}, + {"value_a": "10", "value_b": "20", "value_c": "30"}, + ] + ) + result = await executor.execute(payload) + + assert isinstance(result, pd.DataFrame) + assert len(result) == 2 + + records = result.to_dict(orient="records") + + assert records[0]["output"][0] == {"key": "basic_7", "value": "100"} + assert records[0]["output"][1] == {"key": "basic_15", "value": "200"} + assert records[0]["output"][2] == {"key": "basic_30", "value": "300"} + assert records[0]["message"] == "cobertura" + + assert records[1]["output"][0] == {"key": "basic_7", "value": "10"} + assert records[1]["output"][1] == {"key": "basic_15", "value": "20"} + assert records[1]["output"][2] == {"key": "basic_30", "value": "30"} + assert records[1]["message"] == "cobertura" + + +def test_multiple_outputs_rejects_mixed_terminal_nodes(): + """Rule.create must reject a graph mixing Output and MultipleOutputs.""" + import json as _json + + with open("tests/resources/multiple-outputs.json", "r") as f: + graph_data = _json.load(f) + + # inject a conflicting Output node + graph_data["nodes"]["99"] = { + "id": 99, + "name": "Output", + "data": {"message": None}, + "inputs": { + "input_value": { + "connections": [{"node": 2, "output": "output_value", "data": {}}] + } + }, + "outputs": {}, + } + + with pytest.raises(ValueError, match="mix"): + Rule.create( + graph_data, + nodes_registry=nodes.registry(), + dynamic_nodes_registry=nodes.dynamic_nodes_registry(), + ) + + @pytest.mark.asyncio async def test_subflow_with_connector(): _input_df = pd.DataFrame( diff --git a/tests/test_nodes/test_multiple_outputs.py b/tests/test_nodes/test_multiple_outputs.py new file mode 100644 index 0000000..7dbaf28 --- /dev/null +++ b/tests/test_nodes/test_multiple_outputs.py @@ -0,0 +1,208 @@ +import pandas as pd +import pydantic +import pytest + +from retrack.nodes import dynamic_nodes_registry +from retrack.nodes.base import NodeKind +from retrack.nodes.dynamic.multiple_outputs import ( + MultipleOutputsMetadataModel, + multiple_outputs_factory, +) +from retrack.utils import constants + + +# --------------------------------------------------------------------------- +# MultipleOutputsMetadataModel — unit tests +# --------------------------------------------------------------------------- + + +def test_metadata_valid(): + m = MultipleOutputsMetadataModel(headers_map=["basic_7", "basic_15"]) + assert m.headers_map == ["basic_7", "basic_15"] + assert m.message is None + + +def test_metadata_with_message(): + m = MultipleOutputsMetadataModel(headers_map=["x"], message="cobertura") + assert m.message == "cobertura" + + +def test_metadata_empty_string_message_becomes_none(): + m = MultipleOutputsMetadataModel(headers_map=["x"], message="") + assert m.message is None + + +def test_metadata_rejects_empty_headers_map(): + with pytest.raises(pydantic.ValidationError, match="headers_map must not be empty"): + MultipleOutputsMetadataModel(headers_map=[]) + + +def test_metadata_rejects_duplicate_keys(): + with pytest.raises(pydantic.ValidationError, match="duplicate"): + MultipleOutputsMetadataModel(headers_map=["a", "a"]) + + +# --------------------------------------------------------------------------- +# Factory — registration and construction +# --------------------------------------------------------------------------- + + +@pytest.fixture +def node_metadata(): + return { + "id": 10, + "data": { + "headers_map": ["basic_7", "basic_15", "basic_30"], + "message": "cobertura", + }, + "inputs": { + "basic_7": { + "connections": [{"node": 2, "output": "output_value", "data": {}}] + }, + "basic_15": { + "connections": [{"node": 3, "output": "output_value", "data": {}}] + }, + "basic_30": { + "connections": [{"node": 4, "output": "output_value", "data": {}}] + }, + }, + "name": "MultipleOutputs", + } + + +def test_factory_is_registered_in_dynamic_registry(): + factory = dynamic_nodes_registry().get("MultipleOutputs") + assert callable(factory) + + +def test_factory_creates_valid_pydantic_model(node_metadata): + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + + assert issubclass(NodeClass, pydantic.BaseModel) + + node = NodeClass(**node_metadata) + assert hasattr(node, "run") + assert node.kind() == NodeKind.OUTPUT + + +def test_factory_creates_inputs_dynamically(node_metadata): + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + node = NodeClass(**node_metadata) + + input_fields = node.inputs.model_fields + assert set(input_fields.keys()) == {"basic_7", "basic_15", "basic_30"} + + +def test_factory_node_has_no_outputs(node_metadata): + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + node = NodeClass(**node_metadata) + + assert node.outputs is None + + +# --------------------------------------------------------------------------- +# run() — output contract +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_run_returns_series_of_lists(node_metadata): + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + node = NodeClass(**node_metadata) + + result = await node.run( + basic_7=pd.Series([3028, 2000]), + basic_15=pd.Series([7194, 5000]), + basic_30=pd.Series([15720, 8000]), + ) + + output = result[constants.OUTPUT_REFERENCE_COLUMN] + assert isinstance(output, pd.Series) + assert len(output) == 2 + + assert output.iloc[0] == [ + {"key": "basic_7", "value": 3028}, + {"key": "basic_15", "value": 7194}, + {"key": "basic_30", "value": 15720}, + ] + assert output.iloc[1] == [ + {"key": "basic_7", "value": 2000}, + {"key": "basic_15", "value": 5000}, + {"key": "basic_30", "value": 8000}, + ] + + +@pytest.mark.asyncio +async def test_run_preserves_index(node_metadata): + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + node = NodeClass(**node_metadata) + + idx = pd.Index([10, 20, 30]) + result = await node.run( + basic_7=pd.Series([1, 2, 3], index=idx), + basic_15=pd.Series([4, 5, 6], index=idx), + basic_30=pd.Series([7, 8, 9], index=idx), + ) + + assert list(result[constants.OUTPUT_REFERENCE_COLUMN].index) == [10, 20, 30] + + +@pytest.mark.asyncio +async def test_run_message_is_replicated(node_metadata): + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + node = NodeClass(**node_metadata) + + result = await node.run( + basic_7=pd.Series([1]), + basic_15=pd.Series([2]), + basic_30=pd.Series([3]), + ) + + assert result[constants.OUTPUT_MESSAGE_REFERENCE_COLUMN] == "cobertura" + + +@pytest.mark.asyncio +async def test_run_single_key(): + single_key_metadata = { + "id": 99, + "data": {"headers_map": ["only_key"], "message": None}, + "inputs": { + "only_key": { + "connections": [{"node": 1, "output": "output_value", "data": {}}] + }, + }, + "name": "MultipleOutputs", + } + factory = multiple_outputs_factory(**single_key_metadata) + node = factory(**single_key_metadata) + + result = await node.run(only_key=pd.Series([42])) + + output = result[constants.OUTPUT_REFERENCE_COLUMN] + assert output.iloc[0] == [{"key": "only_key", "value": 42}] + + +@pytest.mark.asyncio +async def test_run_output_column_is_not_nan(node_metadata): + """has_ended() depends on output not being NaN after run.""" + + factory = dynamic_nodes_registry().get("MultipleOutputs") + NodeClass = factory(**node_metadata) + node = NodeClass(**node_metadata) + + result = await node.run( + basic_7=pd.Series([100]), + basic_15=pd.Series([200]), + basic_30=pd.Series([300]), + ) + + output_series = result[constants.OUTPUT_REFERENCE_COLUMN] + assert not output_series.isna().any(), ( + "output must not be NaN so has_ended() returns True" + ) diff --git a/tests/test_transformers/__init__.py b/tests/test_transformers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_transformers/test_normalize_debug.py b/tests/test_transformers/test_normalize_debug.py new file mode 100644 index 0000000..aca343c --- /dev/null +++ b/tests/test_transformers/test_normalize_debug.py @@ -0,0 +1,280 @@ +from retrack.utils.transformers import normalize_execution_for_debug_iter + + +def _make_node(node_type, inputs, data=None): + return { + "id": "node-1", + "name": node_type, + "type": node_type, + "inputs": inputs, + "outputs": [], + "default": None, + "data": data or [], + } + + +def _normalize(nodes_at_index): + return list(normalize_execution_for_debug_iter([nodes_at_index])) + + +# --------------------------------------------------------------------------- +# node Output +# --------------------------------------------------------------------------- + + +def test_output_single_value(): + node = _make_node( + "Output", + inputs=[ + { + "node_id": "n0", + "target_name": "input_value", + "value": 3028, + "source_name": "output", + } + ], + data=[{"name": "message", "value": "basic"}], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [ + {"name": "output", "value": 3028, "message": "basic"} + ] + + +def test_output_message_none(): + node = _make_node( + "Output", + inputs=[ + { + "node_id": "n0", + "target_name": "input_value", + "value": 1.5, + "source_name": "output", + } + ], + data=[{"name": "message", "value": None}], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [{"name": "output", "value": 1.5, "message": None}] + + +def test_output_filters_none_value(): + node = _make_node( + "Output", + inputs=[ + { + "node_id": "n0", + "target_name": "input_value", + "value": None, + "source_name": "output", + } + ], + data=[], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [] + + +def test_output_filters_nan_value(): + node = _make_node( + "Output", + inputs=[ + { + "node_id": "n0", + "target_name": "input_value", + "value": float("nan"), + "source_name": "output", + } + ], + data=[], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [] + + +def test_output_empty_inputs(): + node = _make_node("Output", inputs=[], data=[{"name": "message", "value": "x"}]) + result = _normalize([node]) + assert result[0]["outputs"] == [] + + +# --------------------------------------------------------------------------- +# node MultipleOutputs +# --------------------------------------------------------------------------- + + +def test_multiple_outputs_expands_all_keys(): + node = _make_node( + "MultipleOutputs", + inputs=[ + { + "node_id": "a", + "target_name": "basic_7", + "value": 3028, + "source_name": "output", + }, + { + "node_id": "b", + "target_name": "basic_15", + "value": 7194, + "source_name": "output", + }, + { + "node_id": "c", + "target_name": "basic_30", + "value": 15720, + "source_name": "output", + }, + ], + data=[{"name": "message", "value": "cobertura"}], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [ + {"name": "basic_7", "value": 3028, "message": "cobertura"}, + {"name": "basic_15", "value": 7194, "message": "cobertura"}, + {"name": "basic_30", "value": 15720, "message": "cobertura"}, + ] + + +def test_multiple_outputs_replicates_message_on_all_entries(): + node = _make_node( + "MultipleOutputs", + inputs=[ + {"node_id": "a", "target_name": "k1", "value": 1, "source_name": "output"}, + {"node_id": "b", "target_name": "k2", "value": 2, "source_name": "output"}, + ], + data=[{"name": "message", "value": "msg"}], + ) + result = _normalize([node]) + messages = [e["message"] for e in result[0]["outputs"]] + assert messages == ["msg", "msg"] + + +def test_multiple_outputs_message_none(): + node = _make_node( + "MultipleOutputs", + inputs=[ + {"node_id": "a", "target_name": "k1", "value": 10, "source_name": "output"}, + ], + data=[{"name": "message", "value": None}], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [{"name": "k1", "value": 10, "message": None}] + + +def test_multiple_outputs_filters_none_value(): + node = _make_node( + "MultipleOutputs", + inputs=[ + { + "node_id": "a", + "target_name": "k1", + "value": 100, + "source_name": "output", + }, + { + "node_id": "b", + "target_name": "k2", + "value": None, + "source_name": "output", + }, + { + "node_id": "c", + "target_name": "k3", + "value": 200, + "source_name": "output", + }, + ], + data=[], + ) + result = _normalize([node]) + assert [e["name"] for e in result[0]["outputs"]] == ["k1", "k3"] + + +def test_multiple_outputs_filters_nan_value(): + node = _make_node( + "MultipleOutputs", + inputs=[ + { + "node_id": "a", + "target_name": "k1", + "value": float("nan"), + "source_name": "output", + }, + {"node_id": "b", "target_name": "k2", "value": 42, "source_name": "output"}, + ], + data=[], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [{"name": "k2", "value": 42, "message": None}] + + +def test_multiple_outputs_single_key(): + node = _make_node( + "MultipleOutputs", + inputs=[ + { + "node_id": "a", + "target_name": "only_key", + "value": 7, + "source_name": "output", + } + ], + data=[{"name": "message", "value": "solo"}], + ) + result = _normalize([node]) + assert result[0]["outputs"] == [{"name": "only_key", "value": 7, "message": "solo"}] + + +def test_multiple_outputs_preserves_order(): + keys = ["z", "a", "m"] + node = _make_node( + "MultipleOutputs", + inputs=[ + {"node_id": str(i), "target_name": k, "value": i, "source_name": "output"} + for i, k in enumerate(keys) + ], + data=[], + ) + result = _normalize([node]) + assert [e["name"] for e in result[0]["outputs"]] == keys + + +def test_multiple_outputs_empty_inputs(): + node = _make_node("MultipleOutputs", inputs=[], data=[]) + result = _normalize([node]) + assert result[0]["outputs"] == [] + + +# --------------------------------------------------------------------------- +# backward compatibility +# --------------------------------------------------------------------------- + + +def test_non_terminal_nodes_ignored(): + other = _make_node( + "Check", + inputs=[ + {"node_id": "x", "target_name": "v", "value": 99, "source_name": "output"} + ], + ) + result = _normalize([other]) + assert result[0]["outputs"] == [] + + +def test_output_and_non_terminal_together(): + output_node = _make_node( + "Output", + inputs=[ + { + "node_id": "a", + "target_name": "input_value", + "value": 5, + "source_name": "output", + } + ], + data=[{"name": "message", "value": "ok"}], + ) + other = _make_node("Math", inputs=[]) + result = _normalize([output_node, other]) + assert result[0]["outputs"] == [{"name": "output", "value": 5, "message": "ok"}] diff --git a/tests/test_validators/__init__.py b/tests/test_validators/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_validators/test_single_terminal_node.py b/tests/test_validators/test_single_terminal_node.py new file mode 100644 index 0000000..8d9fa05 --- /dev/null +++ b/tests/test_validators/test_single_terminal_node.py @@ -0,0 +1,109 @@ +import pytest + +from retrack.validators.single_terminal_node import SingleTerminalNodeValidator + + +@pytest.fixture +def validator(): + return SingleTerminalNodeValidator() + + +def _graph(node_names): + return { + "nodes": { + str(i): {"id": i, "name": name, "data": {}, "inputs": {}, "outputs": {}} + for i, name in enumerate(node_names) + } + } + + +# --- valid cases --- + + +def test_accepts_single_output_node(validator): + is_valid, msg = validator.validate(_graph(["Start", "Output"])) + assert is_valid is True + assert msg is None + + +def test_accepts_multiple_output_nodes_same_type(validator): + """Conditional branches produce multiple Output nodes — must remain valid.""" + is_valid, msg = validator.validate(_graph(["Start", "Output", "Output", "Output"])) + assert is_valid is True + assert msg is None + + +def test_accepts_single_multiple_outputs_node(validator): + is_valid, msg = validator.validate(_graph(["Start", "MultipleOutputs"])) + assert is_valid is True + assert msg is None + + +def test_accepts_multiple_multiple_outputs_nodes(validator): + is_valid, msg = validator.validate( + _graph(["Start", "MultipleOutputs", "MultipleOutputs"]) + ) + assert is_valid is True + assert msg is None + + +def test_accepts_graph_with_no_terminal_node(validator): + """Validator does not enforce presence — that is handled separately.""" + is_valid, msg = validator.validate(_graph(["Start", "Input"])) + assert is_valid is True + + +# --- invalid cases --- + + +def test_rejects_output_and_multiple_outputs_together(validator): + is_valid, msg = validator.validate(_graph(["Start", "Output", "MultipleOutputs"])) + assert is_valid is False + assert msg is not None + + +def test_rejects_multiple_outputs_and_output_reversed(validator): + is_valid, msg = validator.validate(_graph(["MultipleOutputs", "Output"])) + assert is_valid is False + + +def test_rejects_mix_with_many_output_nodes(validator): + """Conditional branches with Output + one MultipleOutputs must be rejected.""" + is_valid, msg = validator.validate( + _graph(["Start", "Output", "Output", "MultipleOutputs"]) + ) + assert is_valid is False + + +# --- case insensitivity --- + + +def test_case_insensitive_output(validator): + graph = { + "nodes": { + "0": {"id": 0, "name": "output", "data": {}, "inputs": {}, "outputs": {}} + } + } + is_valid, _ = validator.validate(graph) + assert is_valid is True + + +def test_case_insensitive_multiple_outputs(validator): + graph = { + "nodes": { + "0": { + "id": 0, + "name": "multipleoutputs", + "data": {}, + "inputs": {}, + "outputs": {}, + } + } + } + is_valid, _ = validator.validate(graph) + assert is_valid is True + + +def test_empty_graph(validator): + is_valid, _ = validator.validate({"nodes": {}}) + assert is_valid is True