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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[tool.poetry]
name = "retrack"
version = "3.5.0"
version = "3.6.0"
description = "A business rules engine"
authors = ["Gabriel Guarisa <gabriel.guarisa@pier.digital>"]
license = "MIT"
Expand Down
2 changes: 2 additions & 0 deletions retrack/nodes/dynamic/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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

Expand Down
65 changes: 65 additions & 0 deletions retrack/nodes/dynamic/multiple_outputs.py
Original file line number Diff line number Diff line change
@@ -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
2 changes: 1 addition & 1 deletion retrack/utils/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"}
54 changes: 38 additions & 16 deletions retrack/utils/transformers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
}
Expand Down
2 changes: 2 additions & 0 deletions retrack/validators/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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


Expand Down
30 changes: 30 additions & 0 deletions retrack/validators/single_terminal_node.py
Original file line number Diff line number Diff line change
@@ -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
82 changes: 82 additions & 0 deletions tests/resources/multiple-outputs.json
Original file line number Diff line number Diff line change
@@ -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"
}
}
}
93 changes: 93 additions & 0 deletions tests/test_engine/test_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading
Loading