chore: move representation utils to workflows core (#224)

* chore: move representation utils around

* chore: changesets

* chore: use fixture in tests
This commit is contained in:
Clelia (Astra) Bertelli
2025-11-20 17:22:31 +01:00
committed by GitHub
parent 300fd05416
commit 91159d7c56
5 changed files with 152 additions and 183 deletions
+6
View File
@@ -0,0 +1,6 @@
---
"llama-index-utils-workflow": patch
"llama-index-workflows": patch
---
Moving `_extract_workflow_structure` to its own module in workflow core
@@ -3,7 +3,6 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, List, Tuple, Union, cast
from llama_index.core.agent.workflow import (
@@ -15,194 +14,25 @@ from llama_index.core.agent.workflow import (
from llama_index.core.tools import AsyncBaseTool, BaseTool
from pyvis.network import Network
from workflows import Workflow
from workflows.decorators import StepConfig
from workflows.events import (
Event,
HumanResponseEvent,
InputRequiredEvent,
StartEvent,
StopEvent,
)
from workflows.handler import WorkflowHandler
from workflows.representation_utils import (
DrawWorkflowEdge,
DrawWorkflowGraph,
DrawWorkflowNode,
_truncate_label,
)
from workflows.representation_utils import (
extract_workflow_structure as _extract_workflow_structure,
)
from workflows.runtime.types.results import AddCollectedEvent, StepWorkerResult
from workflows.runtime.types.ticks import TickAddEvent, TickStepResult, WorkflowTick
@dataclass
class DrawWorkflowNode:
"""Represents a node in the workflow graph."""
id: str
label: str
node_type: str # 'step', 'event', 'external'
title: str | None = None
event_type: type | None = None # Store the actual event type for styling decisions
@dataclass
class DrawWorkflowEdge:
"""Represents an edge in the workflow graph."""
source: str
target: str
@dataclass
class DrawWorkflowGraph:
"""Intermediate representation of workflow structure."""
nodes: List[DrawWorkflowNode]
edges: List[DrawWorkflowEdge]
def _truncate_label(label: str, max_length: int) -> str:
"""Helper to truncate long labels."""
return label if len(label) <= max_length else f"{label[: max_length - 1]}*"
def _extract_workflow_structure(
workflow: Workflow, max_label_length: int | None = None
) -> DrawWorkflowGraph:
"""Extract workflow structure into an intermediate representation."""
# Get workflow steps
steps = workflow._get_steps()
nodes = []
edges = []
added_nodes = set() # Track added node IDs to avoid duplicates
step_config: StepConfig | None = None
# Only one kind of `StopEvent` is allowed in a `Workflow`.
# Assuming that `Workflow` is validated before drawing, it's enough to find the first one.
current_stop_event = None
for step_name, step_func in steps.items():
step_config = step_func._step_config
for return_type in step_config.return_types:
if issubclass(return_type, StopEvent):
current_stop_event = return_type
break
if current_stop_event:
break
# First pass: Add all nodes
for step_name, step_func in steps.items():
step_config = step_func._step_config
# Add step node
step_label = (
_truncate_label(step_name, max_label_length)
if max_label_length
else step_name
)
step_title = (
step_name
if max_label_length and len(step_name) > max_label_length
else None
)
if step_name not in added_nodes:
nodes.append(
DrawWorkflowNode(
id=step_name,
label=step_label,
node_type="step",
title=step_title,
)
)
added_nodes.add(step_name)
# Add event nodes for accepted events
for event_type in step_config.accepted_events:
if event_type == StopEvent and event_type != current_stop_event:
continue
event_label = (
_truncate_label(event_type.__name__, max_label_length)
if max_label_length
else event_type.__name__
)
event_title = (
event_type.__name__
if max_label_length and len(event_type.__name__) > max_label_length
else None
)
if event_type.__name__ not in added_nodes:
nodes.append(
DrawWorkflowNode(
id=event_type.__name__,
label=event_label,
node_type="event",
title=event_title,
event_type=event_type,
)
)
added_nodes.add(event_type.__name__)
# Add event nodes for return types
for return_type in step_config.return_types:
if return_type is type(None):
continue
return_label = (
_truncate_label(return_type.__name__, max_label_length)
if max_label_length
else return_type.__name__
)
return_title = (
return_type.__name__
if max_label_length and len(return_type.__name__) > max_label_length
else None
)
if return_type.__name__ not in added_nodes:
nodes.append(
DrawWorkflowNode(
id=return_type.__name__,
label=return_label,
node_type="event",
title=return_title,
event_type=return_type,
)
)
added_nodes.add(return_type.__name__)
# Add external_step node when InputRequiredEvent is found
if (
issubclass(return_type, InputRequiredEvent)
and "external_step" not in added_nodes
):
nodes.append(
DrawWorkflowNode(
id="external_step",
label="external_step",
node_type="external",
)
)
added_nodes.add("external_step")
# Second pass: Add edges
for step_name, step_func in steps.items():
step_config = step_func._step_config
# Edges from steps to return types
for return_type in step_config.return_types:
if return_type is not type(None):
edges.append(DrawWorkflowEdge(step_name, return_type.__name__))
if issubclass(return_type, InputRequiredEvent):
edges.append(DrawWorkflowEdge(return_type.__name__, "external_step"))
# Edges from events to steps
for event_type in step_config.accepted_events:
edges.append(DrawWorkflowEdge(event_type.__name__, step_name))
if issubclass(event_type, HumanResponseEvent):
edges.append(DrawWorkflowEdge("external_step", event_type.__name__))
return DrawWorkflowGraph(nodes=nodes, edges=edges)
def _get_node_color(node: DrawWorkflowNode) -> str:
"""Determine color for a node based on its type and event_type."""
if node.node_type == "step":
@@ -74,7 +74,7 @@ def _truncate_label(label: str, max_length: int) -> str:
return label if len(label) <= max_length else f"{label[: max_length - 1]}*"
def _extract_workflow_structure(
def extract_workflow_structure(
workflow: Workflow, max_label_length: Optional[int] = None
) -> DrawWorkflowGraph:
"""Extract workflow structure into an intermediate representation."""
@@ -50,6 +50,7 @@ from workflows.protocol.serializable_events import (
EventEnvelopeWithMetadata,
EventValidationError,
)
from workflows.representation_utils import extract_workflow_structure
from workflows.server.abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
@@ -62,8 +63,6 @@ from workflows.types import RunResultT
# Protocol models are used on the client side; server responds with plain dicts
from workflows.utils import _nanoid as nanoid
from .representation_utils import _extract_workflow_structure
logger = logging.getLogger()
@@ -623,7 +622,7 @@ class WorkflowServer:
"""
workflow = self._extract_workflow(request)
try:
workflow_graph = _extract_workflow_structure(workflow.workflow)
workflow_graph = extract_workflow_structure(workflow.workflow)
except Exception as e:
raise HTTPException(
detail=f"Error while getting JSON workflow representation: {e}",
@@ -0,0 +1,134 @@
import pytest
from workflows.events import StartEvent, StopEvent
from workflows.representation_utils import (
DrawWorkflowEdge,
DrawWorkflowGraph,
DrawWorkflowNode,
extract_workflow_structure,
)
from .conftest import DummyWorkflow, LastEvent, OneTestEvent # type: ignore[import]
@pytest.fixture()
def ground_truth_repr() -> DrawWorkflowGraph:
return DrawWorkflowGraph(
nodes=[
DrawWorkflowNode(
id="end_step",
label="end_step",
node_type="step",
title=None,
event_type=None,
),
DrawWorkflowNode(
id="LastEvent",
label="LastEvent",
node_type="event",
title=None,
event_type=LastEvent,
),
DrawWorkflowNode(
id="StopEvent",
label="StopEvent",
node_type="event",
title=None,
event_type=StopEvent,
),
DrawWorkflowNode(
id="middle_step",
label="middle_step",
node_type="step",
title=None,
event_type=None,
),
DrawWorkflowNode(
id="OneTestEvent",
label="OneTestEvent",
node_type="event",
title=None,
event_type=OneTestEvent,
),
DrawWorkflowNode(
id="start_step",
label="start_step",
node_type="step",
title=None,
event_type=None,
),
DrawWorkflowNode(
id="StartEvent",
label="StartEvent",
node_type="event",
title=None,
event_type=StartEvent,
),
],
edges=[
DrawWorkflowEdge(source="end_step", target="StopEvent"),
DrawWorkflowEdge(source="LastEvent", target="end_step"),
DrawWorkflowEdge(source="middle_step", target="LastEvent"),
DrawWorkflowEdge(source="OneTestEvent", target="middle_step"),
DrawWorkflowEdge(source="start_step", target="OneTestEvent"),
DrawWorkflowEdge(source="StartEvent", target="start_step"),
],
)
def test_extract_workflow_structure(ground_truth_repr: DrawWorkflowGraph) -> None:
wf = DummyWorkflow()
graph = extract_workflow_structure(workflow=wf)
assert isinstance(graph, DrawWorkflowGraph)
assert sorted(
[node.id for node in ground_truth_repr.nodes if node.node_type == "step"]
) == sorted([node.id for node in graph.nodes if node.node_type == "step"])
assert sorted(
[node.id for node in ground_truth_repr.nodes if node.node_type == "event"]
) == sorted([node.id for node in graph.nodes if node.node_type == "event"])
expected_edges = ground_truth_repr.edges
for edge in expected_edges:
assert edge in graph.edges
def test_extract_workflow_structure_trim_label() -> None:
wf = DummyWorkflow()
graph = extract_workflow_structure(workflow=wf, max_label_length=2)
assert sorted(["e*", "m*", "s*"]) == sorted(
[node.label for node in graph.nodes if node.node_type == "step"]
)
assert sorted(["S*", "S*", "O*", "L*"]) == sorted(
[node.label for node in graph.nodes if node.node_type == "event"]
)
def test_graph_to_response_model() -> None:
graph = DrawWorkflowGraph(
nodes=[
DrawWorkflowNode(
id="test", label="test", node_type="step", title=None, event_type=None
),
DrawWorkflowNode(
id="OneTestEvent",
label="OneTestEvent",
node_type="event",
title=None,
event_type=OneTestEvent,
),
],
edges=[DrawWorkflowEdge(source="test", target="OneTestEvent")],
)
res = graph.to_response_model()
assert len(res.nodes) == 2
assert res.nodes[0].event_type is None
assert res.nodes[0].title is None
assert res.nodes[0].node_type == "step"
assert res.nodes[0].label == "test"
assert res.nodes[0].id == "test"
assert res.nodes[1].event_type == OneTestEvent.__name__
assert res.nodes[1].title is None
assert res.nodes[1].node_type == "event"
assert res.nodes[1].label == "OneTestEvent"
assert res.nodes[1].id == "OneTestEvent"
assert len(res.edges) == 1
assert res.edges[0].source == "test"
assert res.edges[0].target == "OneTestEvent"