mirror of
https://github.com/run-llama/workflows-py.git
synced 2026-07-21 04:05:25 -04:00
Draw nested workflows (#292)
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
---
|
||||
"llama-index-utils-workflow": patch
|
||||
---
|
||||
|
||||
Added in the utility function draw_all_possible_flows_nested_mermaid. This draws the possible flows (not execution) of nested workflows
|
||||
@@ -2,8 +2,12 @@
|
||||
# Copyright (c) 2026 LlamaIndex Inc.
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import inspect
|
||||
import json
|
||||
from typing import Any, Dict, List, Tuple, Union, cast
|
||||
import sys
|
||||
import textwrap
|
||||
from typing import Any, Callable, Dict, List, Tuple, Union, cast
|
||||
|
||||
from llama_index.core.agent.workflow import (
|
||||
AgentWorkflow,
|
||||
@@ -44,6 +48,8 @@ def _get_node_color(node: WorkflowGraphNode) -> str:
|
||||
return "#ADD8E6" # Light blue for steps
|
||||
elif node.node_type == "external":
|
||||
return "#BEDAE4" # Light blue-gray for external
|
||||
elif node.node_type == "child_connector":
|
||||
return "#E0E0E0" # Light gray for child workflow connectors
|
||||
elif node.node_type == "resource":
|
||||
return "#DDA0DD" # Plum/light purple for resources
|
||||
elif node.node_type == "resource_config":
|
||||
@@ -78,6 +84,8 @@ def _get_node_shape(node: WorkflowGraphNode) -> str:
|
||||
"""Determine shape for a node based on its type."""
|
||||
if node.node_type in ("step", "external"):
|
||||
return "box"
|
||||
elif node.node_type == "child_connector":
|
||||
return "ellipse"
|
||||
elif node.node_type == "event":
|
||||
return "ellipse"
|
||||
elif node.node_type == "resource":
|
||||
@@ -183,6 +191,8 @@ def _get_mermaid_css_class(node: WorkflowGraphNode) -> str:
|
||||
return "stepStyle"
|
||||
elif node.node_type == "external":
|
||||
return "externalStyle"
|
||||
elif node.node_type == "child_connector":
|
||||
return "childConnectorStyle"
|
||||
elif node.node_type == "resource":
|
||||
return "resourceStyle"
|
||||
elif node.node_type == "resource_config":
|
||||
@@ -302,6 +312,7 @@ def _render_mermaid(
|
||||
" classDef workflowAgentStyle fill:#66ccff,color:#000000",
|
||||
" classDef workflowToolStyle fill:#ff9966,color:#000000",
|
||||
" classDef workflowHandoffStyle fill:#E27AFF,color:#000000",
|
||||
" classDef childConnectorStyle fill:#E0E0E080,color:#555555,stroke-width:0px",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -512,11 +523,194 @@ def _extract_execution_graph(
|
||||
return nodes, edges
|
||||
|
||||
|
||||
def _get_workflow_classes_from_step(method_callable: Callable | Any) -> list[str]:
|
||||
"""
|
||||
Finds classes instantiated within a method that inherit from Workflow.
|
||||
Resolves names against the method's module namespace.
|
||||
"""
|
||||
if method_callable is None:
|
||||
return []
|
||||
|
||||
workflow_classes = []
|
||||
try:
|
||||
# Get source and module context
|
||||
source = inspect.getsource(method_callable)
|
||||
clean_source = textwrap.dedent(source)
|
||||
tree = ast.parse(clean_source)
|
||||
|
||||
# Get the module where the step is defined to look up names
|
||||
module_name = getattr(method_callable, "__module__", None)
|
||||
|
||||
# Add this guard:
|
||||
if not isinstance(module_name, str):
|
||||
return []
|
||||
|
||||
module = sys.modules.get(module_name)
|
||||
if not module:
|
||||
return []
|
||||
|
||||
for node in ast.walk(tree):
|
||||
# We are looking for instantiations: e.g., MySubWorkflow()
|
||||
if isinstance(node, ast.Call):
|
||||
# Handle direct calls: Name()
|
||||
if isinstance(node.func, ast.Name):
|
||||
class_name = node.func.id
|
||||
# Handle attribute calls: module.Name()
|
||||
elif isinstance(node.func, ast.Attribute):
|
||||
class_name = node.func.attr
|
||||
else:
|
||||
continue
|
||||
|
||||
# Look up the name in the module's namespace
|
||||
obj = getattr(module, class_name, None)
|
||||
|
||||
# Robust check: Is it a class, and is it a Workflow subclass?
|
||||
if (
|
||||
inspect.isclass(obj)
|
||||
and issubclass(obj, Workflow)
|
||||
and obj is not Workflow # Don't include the base class itself
|
||||
):
|
||||
if class_name not in workflow_classes:
|
||||
workflow_classes.append(class_name)
|
||||
|
||||
except Exception:
|
||||
# Fallback to empty list if source cannot be parsed or objects resolved
|
||||
return []
|
||||
|
||||
return workflow_classes
|
||||
|
||||
|
||||
def _get_nested_workflow_representation(
|
||||
workflow: Workflow, include_child_workflows: bool = False
|
||||
) -> WorkflowGraph:
|
||||
"""
|
||||
Introspects a workflow and builds a unified WorkflowGraph.
|
||||
|
||||
If include_child_workflows is True, it performs a 1-level deep scan of
|
||||
step source code to find instantiated sub-workflows and merges their
|
||||
graphs into the parent graph.
|
||||
"""
|
||||
parent_graph = _get_workflow_representation(workflow)
|
||||
if not include_child_workflows:
|
||||
return parent_graph
|
||||
|
||||
# Define the helper AFTER parent_graph is created so it can
|
||||
# modify parent_graph directly via closure (i.e. parent_graph is now in scope for this helper function.
|
||||
def _merge_subgraph_into_parent(
|
||||
child_graph: WorkflowGraph, parent_step_id: str, class_name: str
|
||||
) -> None:
|
||||
"""Internal helper to handle ID prefixing and edge stitching."""
|
||||
prefix = f"{parent_step_id}_{class_name}_"
|
||||
|
||||
# 1. Merge Nodes
|
||||
for c_node in list(child_graph.nodes):
|
||||
c_node_id = getattr(c_node, "id", str(c_node))
|
||||
new_node = WorkflowGenericNode(
|
||||
id=f"{prefix}{c_node_id}",
|
||||
label=getattr(c_node, "label", c_node_id),
|
||||
node_type=getattr(c_node, "node_type", "step"),
|
||||
event_type=getattr(c_node, "event_type", None),
|
||||
)
|
||||
parent_graph.nodes.append(new_node)
|
||||
|
||||
# 2. Merge Edges
|
||||
for edge in child_graph.edges:
|
||||
parent_graph.edges.append(
|
||||
WorkflowGraphEdge(
|
||||
source=f"{prefix}{edge.source}",
|
||||
target=f"{prefix}{edge.target}",
|
||||
label=edge.label,
|
||||
)
|
||||
)
|
||||
|
||||
# 3. Stitch: Parent Step -> "calls" node -> Child Start
|
||||
child_start_id = next(
|
||||
(
|
||||
n.id
|
||||
for n in child_graph.nodes
|
||||
if getattr(n, "event_type", None) == "StartEvent"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if child_start_id:
|
||||
calls_node_id = f"{prefix}calls"
|
||||
parent_graph.nodes.append(
|
||||
WorkflowGenericNode(
|
||||
id=calls_node_id,
|
||||
label=f"calls: {class_name}",
|
||||
node_type="child_connector",
|
||||
)
|
||||
)
|
||||
parent_graph.edges.append(
|
||||
WorkflowGraphEdge(source=parent_step_id, target=calls_node_id)
|
||||
)
|
||||
parent_graph.edges.append(
|
||||
WorkflowGraphEdge(
|
||||
source=calls_node_id, target=f"{prefix}{child_start_id}"
|
||||
)
|
||||
)
|
||||
|
||||
# 4. Stitch: Child Stop -> "returns" node -> Parent Step
|
||||
child_stop_id = next(
|
||||
(
|
||||
n.id
|
||||
for n in child_graph.nodes
|
||||
if getattr(n, "event_type", None) == "StopEvent"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if child_stop_id:
|
||||
returns_node_id = f"{prefix}returns"
|
||||
parent_graph.nodes.append(
|
||||
WorkflowGenericNode(
|
||||
id=returns_node_id,
|
||||
label=f"returns: {class_name}",
|
||||
node_type="child_connector",
|
||||
)
|
||||
)
|
||||
parent_graph.edges.append(
|
||||
WorkflowGraphEdge(
|
||||
source=f"{prefix}{child_stop_id}", target=returns_node_id
|
||||
)
|
||||
)
|
||||
parent_graph.edges.append(
|
||||
WorkflowGraphEdge(source=returns_node_id, target=parent_step_id)
|
||||
)
|
||||
|
||||
# --- Discovery and Execution Loop ---
|
||||
steps_lookup = workflow._get_steps()
|
||||
|
||||
for node in list(parent_graph.nodes):
|
||||
step_id = getattr(node, "id", str(node))
|
||||
if getattr(node, "node_type", None) == "step":
|
||||
step_method = steps_lookup.get(step_id)
|
||||
nested_wf_classnames = _get_workflow_classes_from_step(step_method)
|
||||
|
||||
for nested_wf_classname in nested_wf_classnames:
|
||||
try:
|
||||
module_name = step_method.__module__
|
||||
module = sys.modules.get(module_name) or sys.modules.get("__main__")
|
||||
wf_class = getattr(module, nested_wf_classname)
|
||||
|
||||
child_instance = wf_class()
|
||||
child_graph = _get_workflow_representation(child_instance)
|
||||
|
||||
# Executes the merge using the closure above
|
||||
_merge_subgraph_into_parent(
|
||||
child_graph, step_id, nested_wf_classname
|
||||
)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
return parent_graph
|
||||
|
||||
|
||||
def draw_all_possible_flows(
|
||||
workflow: Workflow,
|
||||
filename: str = "workflow_all_flows.html",
|
||||
notebook: bool = False,
|
||||
max_label_length: int | None = None,
|
||||
include_child_workflows: bool = True,
|
||||
) -> None:
|
||||
"""
|
||||
Draws all possible flows of the workflow using Pyvis.
|
||||
@@ -526,9 +720,12 @@ def draw_all_possible_flows(
|
||||
filename: Output HTML filename
|
||||
notebook: Whether running in notebook environment
|
||||
max_label_length: Maximum label length before truncation (None = no limit)
|
||||
include_child_workflows: Whether to include child workflow graphs
|
||||
|
||||
"""
|
||||
graph = _get_workflow_representation(workflow)
|
||||
graph = _get_nested_workflow_representation(
|
||||
workflow, include_child_workflows=include_child_workflows
|
||||
)
|
||||
_render_pyvis(graph, filename, notebook, max_label_length)
|
||||
|
||||
|
||||
@@ -536,21 +733,18 @@ def draw_all_possible_flows_mermaid(
|
||||
workflow: Workflow,
|
||||
filename: str = "workflow_all_flows.mermaid",
|
||||
max_label_length: int | None = None,
|
||||
include_child_workflows: bool = True,
|
||||
) -> str:
|
||||
"""
|
||||
Draws all possible flows of the workflow as a Mermaid diagram.
|
||||
|
||||
Args:
|
||||
workflow: The workflow to visualize
|
||||
filename: Output Mermaid filename
|
||||
max_label_length: Maximum label length before truncation (None = no limit)
|
||||
|
||||
Returns:
|
||||
The Mermaid diagram as a string
|
||||
|
||||
"""
|
||||
graph = _get_workflow_representation(workflow)
|
||||
return _render_mermaid(graph, filename, max_label_length)
|
||||
# Use the new helper to get the full graph structure
|
||||
full_graph = _get_nested_workflow_representation(
|
||||
workflow, include_child_workflows=include_child_workflows
|
||||
)
|
||||
|
||||
# Render to Mermaid format
|
||||
return _render_mermaid(full_graph, filename, max_label_length)
|
||||
|
||||
|
||||
def draw_agent_with_tools(
|
||||
|
||||
@@ -100,3 +100,31 @@ def workflow_with_resources() -> Workflow:
|
||||
@pytest.fixture()
|
||||
def events() -> list[type[Event]]:
|
||||
return [OneTestEvent, AnotherTestEvent]
|
||||
|
||||
|
||||
# --- Nested workflow for testing ---
|
||||
|
||||
|
||||
class ChildWorkflowA(Workflow):
|
||||
@step()
|
||||
async def child_start(self, ev: StartEvent) -> StopEvent:
|
||||
return StopEvent(result="Child processed")
|
||||
|
||||
|
||||
class ParentWorkflow(Workflow):
|
||||
@step()
|
||||
async def parent_start(self, ev: StartEvent) -> Event:
|
||||
# Instantiate the nested workflow. The drawing function will inspect this.
|
||||
child_wf = ChildWorkflowA()
|
||||
# The actual run logic is not important for drawing.
|
||||
_ = await child_wf.run(input="dummy")
|
||||
return Event(result="some result")
|
||||
|
||||
@step()
|
||||
async def parent_end(self, ev: Event) -> StopEvent:
|
||||
return StopEvent(result="Final Result")
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def nested_workflow() -> Workflow:
|
||||
return ParentWorkflow()
|
||||
|
||||
@@ -561,3 +561,96 @@ def test_resource_node_deduplication_in_rendering(
|
||||
assert len(resource_node_defs) == 2, (
|
||||
f"Expected 2 unique resource nodes, found {len(resource_node_defs)}: {resource_node_defs}"
|
||||
)
|
||||
|
||||
|
||||
def test_draw_all_possible_flows_with_child_workflow_mermaid(
|
||||
nested_workflow: Workflow,
|
||||
) -> None:
|
||||
"""Test Mermaid diagram generation for nested workflows."""
|
||||
result = draw_all_possible_flows_mermaid(
|
||||
nested_workflow, include_child_workflows=True
|
||||
)
|
||||
|
||||
# Basic checks
|
||||
assert isinstance(result, str)
|
||||
assert result.startswith("flowchart TD")
|
||||
|
||||
# Check for parent workflow nodes
|
||||
assert "step_parent_start" in result
|
||||
assert "step_parent_end" in result
|
||||
|
||||
# Check for child workflow nodes (prefixed)
|
||||
assert "step_parent_start_ChildWorkflowA_child_start" in result
|
||||
assert "event_parent_start_ChildWorkflowA_StartEvent" in result
|
||||
assert "event_parent_start_ChildWorkflowA_StopEvent" in result
|
||||
|
||||
# Check for connector nodes (calls/returns are now nodes, not edge labels)
|
||||
assert "child_connector_parent_start_ChildWorkflowA_calls" in result
|
||||
assert "calls: ChildWorkflowA" in result
|
||||
assert "child_connector_parent_start_ChildWorkflowA_returns" in result
|
||||
assert "returns: ChildWorkflowA" in result
|
||||
|
||||
# Parent step -> calls node -> child StartEvent
|
||||
assert (
|
||||
"step_parent_start --> child_connector_parent_start_ChildWorkflowA_calls"
|
||||
in result
|
||||
)
|
||||
assert (
|
||||
"child_connector_parent_start_ChildWorkflowA_calls --> event_parent_start_ChildWorkflowA_StartEvent"
|
||||
in result
|
||||
)
|
||||
# Child StopEvent -> returns node -> parent step
|
||||
assert (
|
||||
"event_parent_start_ChildWorkflowA_StopEvent --> child_connector_parent_start_ChildWorkflowA_returns"
|
||||
in result
|
||||
)
|
||||
assert (
|
||||
"child_connector_parent_start_ChildWorkflowA_returns --> step_parent_start"
|
||||
in result
|
||||
)
|
||||
|
||||
|
||||
def test_draw_all_possible_flows_with_child_workflow_pyvis(
|
||||
nested_workflow: Workflow,
|
||||
) -> None:
|
||||
"""Test Pyvis diagram generation includes nested child workflow nodes and edges."""
|
||||
with patch("llama_index.utils.workflow.Network") as mock_network:
|
||||
mock_net_instance = MagicMock()
|
||||
mock_network.return_value = mock_net_instance
|
||||
|
||||
draw_all_possible_flows(
|
||||
nested_workflow, filename="test.html", include_child_workflows=True
|
||||
)
|
||||
|
||||
node_ids = [call[0][0] for call in mock_net_instance.add_node.call_args_list]
|
||||
|
||||
# Parent nodes present
|
||||
assert "parent_start" in node_ids
|
||||
assert "parent_end" in node_ids
|
||||
|
||||
# Child workflow nodes present (prefixed)
|
||||
assert "parent_start_ChildWorkflowA_child_start" in node_ids
|
||||
assert "parent_start_ChildWorkflowA_StartEvent" in node_ids
|
||||
assert "parent_start_ChildWorkflowA_StopEvent" in node_ids
|
||||
|
||||
# Connector nodes present
|
||||
assert "parent_start_ChildWorkflowA_calls" in node_ids
|
||||
assert "parent_start_ChildWorkflowA_returns" in node_ids
|
||||
|
||||
# Check stitching edges through connector nodes
|
||||
edge_pairs = [
|
||||
(call[0][0], call[0][1])
|
||||
for call in mock_net_instance.add_edge.call_args_list
|
||||
]
|
||||
# parent_start -> calls node -> child StartEvent
|
||||
assert ("parent_start", "parent_start_ChildWorkflowA_calls") in edge_pairs
|
||||
assert (
|
||||
"parent_start_ChildWorkflowA_calls",
|
||||
"parent_start_ChildWorkflowA_StartEvent",
|
||||
) in edge_pairs
|
||||
# child StopEvent -> returns node -> parent_start
|
||||
assert (
|
||||
"parent_start_ChildWorkflowA_StopEvent",
|
||||
"parent_start_ChildWorkflowA_returns",
|
||||
) in edge_pairs
|
||||
assert ("parent_start_ChildWorkflowA_returns", "parent_start") in edge_pairs
|
||||
|
||||
Reference in New Issue
Block a user