Add resource nodes to workflow graph export (#269)

* feat: add resource nodes to workflow graph export

Add support for visualizing resource dependencies in workflow graphs:

- Enhanced _Resource class with source location metadata (file, line,
  docstring, unique hash for deduplication)
- Added ResourceDefinition.type_annotation to capture the type from
  Annotated[T, Resource(...)]
- Added WorkflowGraphResourceNode protocol model with full metadata
- Added DrawWorkflowResourceNode class for intermediate representation
- Updated DrawWorkflowEdge to support edge labels (variable names)
- Modified extract_workflow_structure to extract and deduplicate
  resource nodes, connecting them to steps with labeled edges
- Updated Pyvis and Mermaid renderers to display resource nodes
  (hexagon shape, plum color) with metadata tooltips

* test: add tests and example for resource nodes visualization

- Add tests for resource node extraction in test_representation_utils.py
- Add tests for resource rendering in test_drawing.py (Mermaid + Pyvis)
- Add workflow_with_resources fixture with multiple resources
- Add runnable example (examples/visualization/resource_nodes_example.py)
  that generates both Mermaid and Pyvis output with resource nodes

* chore: add changeset for resource nodes feature

* fix: reverse resource edge direction to step -> resource

* refactor: unify DrawWorkflowNode and DrawWorkflowResourceNode

- Merge resource-specific fields into DrawWorkflowNode
- Add to_resource_response_model() for resource serialization
- Keep DrawWorkflowResourceNode as backwards compat alias
- Simplify type annotations in rendering code (remove unions)
- Remove isinstance checks since all nodes are now DrawWorkflowNode

* refactor: simplify _get_clean_node_id to use node_type prefix

* refactor: merge resource_nodes into nodes list

Unify resource nodes with the main nodes list instead of a separate
resource_nodes field. This simplifies the API and rendering code while
maintaining backwards compatibility through a property accessor.

- Remove resource_nodes field from DrawWorkflowGraph
- Add resource_nodes property for backwards compatibility
- Merge WorkflowGraphResourceNode fields into WorkflowGraphNode
- Simplify rendering loops to single iteration over nodes

* refactor: remove unnecessary backwards compatibility aliases

This is all new code in this PR - no need for backwards compat aliases.
Removes WorkflowGraphResourceNode, DrawWorkflowResourceNode,
to_resource_response_model, and resource_nodes property.

* perf: make resource metadata extraction lazy

Move expensive inspect operations (getfile, getsourcelines, getdoc)
and hash computation from Resource creation time to graph extraction
time. This avoids the overhead when resources are used at runtime
but visualization is not needed.

* test: remove unnecessary async from mermaid drawing tests

* consolidate types

* Split types

* feat: add description and schema fields to workflow graph nodes

- Add description field to WorkflowGraph (workflow class docstring)
- Add description field to WorkflowStepNode (step function docstring)
- Add event_schema field to WorkflowEventNode (Pydantic JSON schema)
- Rename WorkflowResourceNode.docstring to description for consistency
- Rename WorkflowGraphNodeEdges to WorkflowGraph

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* filter by node type

* Update add_resource_nodes.md

* remove resource hash

---------

Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
Adrian Lyjak
2026-01-10 13:01:12 -05:00
committed by GitHub
parent 3b043b8801
commit e53c654b42
13 changed files with 1797 additions and 388 deletions
+6
View File
@@ -0,0 +1,6 @@
---
"llama-index-workflows": minor
"llama-index-utils-workflow": minor
---
Add further detail to workflow graph, mainly adding `Resource` nodes to workflow graph and visualizations
@@ -0,0 +1,260 @@
#!/usr/bin/env python3
"""
Example demonstrating resource nodes in workflow graph visualization.
This example shows how resources (dependencies injected via Annotated types)
are rendered in both Mermaid and Pyvis diagrams.
Run this script to generate:
- workflow_with_resources.mermaid - A Mermaid diagram file
- workflow_with_resources.html - An interactive Pyvis HTML visualization
You can view the Mermaid diagram at https://mermaid.live/ by pasting the contents.
The HTML file can be opened directly in any web browser.
"""
import argparse
from typing import Annotated
from llama_index.utils.workflow import (
draw_all_possible_flows,
draw_all_possible_flows_mermaid,
)
from workflows import Workflow, step
from workflows.events import Event, StartEvent, StopEvent
from workflows.resource import Resource
# --- Mock resource types ---
class DatabaseClient:
"""A database client for persistent storage."""
def __init__(self, connection_string: str = "postgres://localhost/db"):
self.connection_string = connection_string
def query(self, sql: str) -> list:
"""Execute a SQL query."""
return []
class CacheClient:
"""A cache client for fast data retrieval."""
def __init__(self, host: str = "localhost", port: int = 6379):
self.host = host
self.port = port
def get(self, key: str) -> str | None:
"""Get a value from cache."""
return None
def set(self, key: str, value: str) -> None:
"""Set a value in cache."""
pass
class LLMClient:
"""A client for interacting with a large language model."""
def __init__(self, api_key: str = "sk-..."):
self.api_key = api_key
async def complete(self, prompt: str) -> str:
"""Generate a completion for the given prompt."""
return f"Response to: {prompt}"
# --- Resource factory functions ---
def get_database_client() -> DatabaseClient:
"""Factory function to create a database client.
This function creates a PostgreSQL database client configured
for the application's data storage needs.
"""
return DatabaseClient(connection_string="postgres://localhost/myapp")
def get_cache_client() -> CacheClient:
"""Factory function to create a Redis cache client.
Provides fast caching for frequently accessed data.
"""
return CacheClient(host="localhost", port=6379)
def get_llm_client() -> LLMClient:
"""Factory function to create an LLM client.
Creates a client for the language model API.
"""
return LLMClient(api_key="sk-example-key")
# --- Event types ---
class QueryProcessedEvent(Event):
"""Event emitted after processing a user query."""
query: str
cached: bool = False
class ContextRetrievedEvent(Event):
"""Event emitted after retrieving context from the database."""
context: str
class ResponseGeneratedEvent(Event):
"""Event emitted after generating an LLM response."""
response: str
# --- Workflow with resources ---
class RAGWorkflow(Workflow):
"""A RAG (Retrieval-Augmented Generation) workflow demonstrating resource usage.
This workflow shows how different steps can depend on shared resources
like database clients, cache clients, and LLM clients.
"""
@step
async def process_query(
self,
ev: StartEvent,
cache: Annotated[CacheClient, Resource(get_cache_client)],
) -> QueryProcessedEvent:
"""Process the incoming query, checking cache first."""
query = getattr(ev, "query", "default query")
# Check if query result is cached
cached_result = cache.get(f"query:{query}")
if cached_result:
return QueryProcessedEvent(query=query, cached=True)
return QueryProcessedEvent(query=query, cached=False)
@step
async def retrieve_context(
self,
ev: QueryProcessedEvent,
db: Annotated[DatabaseClient, Resource(get_database_client)],
cache: Annotated[CacheClient, Resource(get_cache_client)],
) -> ContextRetrievedEvent:
"""Retrieve relevant context from the database."""
if ev.cached:
context = "Cached context"
else:
# Query the database for relevant documents
results = db.query(
f"SELECT content FROM documents WHERE query = '{ev.query}'"
)
context = " ".join(str(r) for r in results) or "No context found"
# Cache the result
cache.set(f"context:{ev.query}", context)
return ContextRetrievedEvent(context=context)
@step
async def generate_response(
self,
ev: ContextRetrievedEvent,
llm: Annotated[LLMClient, Resource(get_llm_client)],
) -> ResponseGeneratedEvent:
"""Generate a response using the LLM with the retrieved context."""
prompt = f"Context: {ev.context}\n\nGenerate a response."
response = await llm.complete(prompt)
return ResponseGeneratedEvent(response=response)
@step
async def finalize_response(
self,
ev: ResponseGeneratedEvent,
cache: Annotated[CacheClient, Resource(get_cache_client)],
) -> StopEvent:
"""Finalize and cache the response."""
# Cache the final response
cache.set("last_response", ev.response)
return StopEvent(result=ev.response)
def main() -> None:
parser = argparse.ArgumentParser(
description="Generate workflow visualizations with resource nodes"
)
parser.add_argument(
"--output-dir",
default=".",
help="Directory to save output files (default: current directory)",
)
parser.add_argument(
"--mermaid-only",
action="store_true",
help="Only generate Mermaid output (print to stdout)",
)
args = parser.parse_args()
# Create the workflow
workflow = RAGWorkflow()
print("=" * 60)
print("Workflow Graph Visualization with Resource Nodes")
print("=" * 60)
# Generate Mermaid diagram
mermaid_file = f"{args.output_dir}/workflow_with_resources.mermaid"
mermaid_output = draw_all_possible_flows_mermaid(
workflow,
filename="" if args.mermaid_only else mermaid_file,
)
print("\n--- Mermaid Diagram ---")
print(mermaid_output)
print()
if not args.mermaid_only:
print(f"Mermaid diagram saved to: {mermaid_file}")
# Generate Pyvis HTML
html_file = f"{args.output_dir}/workflow_with_resources.html"
draw_all_possible_flows(workflow, filename=html_file)
print(f"Interactive Pyvis diagram saved to: {html_file}")
print("\n--- Resource Nodes Summary ---")
print(
"""
The diagram shows:
- HEXAGON nodes (plum color): Resource dependencies
- DatabaseClient: Database connection via get_database_client()
- CacheClient: Cache connection via get_cache_client()
- LLMClient: LLM API client via get_llm_client()
- Edge labels on resource connections show the variable name used in the step
e.g., "db", "cache", "llm"
- Resources are deduplicated: CacheClient appears once even though
it's used by multiple steps (process_query, retrieve_context, finalize_response)
To view the Mermaid diagram:
1. Go to https://mermaid.live/
2. Paste the diagram content above
3. See the interactive visualization
To view the Pyvis diagram:
1. Open workflow_with_resources.html in a web browser
2. Hover over nodes to see metadata (type, getter, source location, docstring)
3. Drag nodes to rearrange the layout
"""
)
if __name__ == "__main__":
main()
@@ -19,7 +19,7 @@ authors = [{name = "Adrian Lyjak", email = "adrianlyjak@gmail.com"}]
requires-python = ">=3.9"
dependencies = [
"llama-index-core>=0.14,<0.15.0",
"llama-index-workflows>=2.11.3,<3.0.0",
"llama-index-workflows>=2.12.0,<3.0.0",
"pyvis>=0.3.2"
]
@@ -8,8 +8,6 @@ from typing import Any, Dict, List, Tuple, Union, cast
from llama_index.core.agent.workflow import (
AgentWorkflow,
BaseWorkflowAgent,
CodeActAgent,
ReActAgent,
)
from llama_index.core.tools import AsyncBaseTool, BaseTool
from pyvis.network import Network
@@ -20,11 +18,12 @@ from workflows.events import (
StopEvent,
)
from workflows.handler import WorkflowHandler
from workflows.representation_utils import (
DrawWorkflowEdge,
DrawWorkflowGraph,
DrawWorkflowNode,
_truncate_label,
from workflows.protocol import (
WorkflowGenericNode,
WorkflowGraph,
WorkflowGraphEdge,
WorkflowGraphNode,
WorkflowResourceNode,
)
from workflows.representation_utils import (
extract_workflow_structure as _extract_workflow_structure,
@@ -33,22 +32,31 @@ from workflows.runtime.types.results import AddCollectedEvent, StepWorkerResult
from workflows.runtime.types.ticks import TickAddEvent, TickStepResult, WorkflowTick
def _get_node_color(node: DrawWorkflowNode) -> str:
def _truncate_label(label: str, max_length: int) -> str:
"""Truncate long labels for visualization."""
return label if len(label) <= max_length else f"{label[: max_length - 1]}*"
def _get_node_color(node: WorkflowGraphNode) -> str:
"""Determine color for a node based on its type and event_type."""
if node.node_type == "step":
return "#ADD8E6" # Light blue for steps
elif node.node_type == "external":
return "#BEDAE4" # Light blue-gray for external
elif node.node_type == "event" and node.event_type:
return _determine_event_color(node.event_type) # Uses original function
elif node.node_type == "resource":
return "#DDA0DD" # Plum/light purple for resources
elif node.node_type == "event":
if node.is_subclass_of("StartEvent"):
return "#E27AFF" # Pink for start events
elif node.is_subclass_of("StopEvent"):
return "#FFA07A" # Orange for stop events
return "#90EE90" # Light green for other events
elif node.node_type == "agent":
# Determine color based on agent type
if node.event_type and issubclass(node.event_type, ReActAgent):
if node.is_subclass_of("ReActAgent"):
return "#E27AFF"
elif node.event_type and issubclass(node.event_type, CodeActAgent):
elif node.is_subclass_of("CodeActAgent"):
return "#66ccff"
else:
return "#90EE90"
return "#90EE90"
elif node.node_type == "tool":
return "#ff9966" # Orange for tools
elif node.node_type == "workflow_base":
@@ -63,30 +71,29 @@ def _get_node_color(node: DrawWorkflowNode) -> str:
return "#90EE90" # Default light green
def _get_node_shape(node: DrawWorkflowNode) -> str:
def _get_node_shape(node: WorkflowGraphNode) -> str:
"""Determine shape for a node based on its type."""
if node.node_type == "step" or node.node_type == "external":
return "box" # Steps and external_step use box
if node.node_type in ("step", "external"):
return "box"
elif node.node_type == "event":
return "ellipse" # Events use ellipse
elif node.node_type == "agent":
return "ellipse" # Agents use ellipse
elif node.node_type == "tool":
return "ellipse" # Tools use ellipse (matching original ellipsis behavior)
return "ellipse"
elif node.node_type == "resource":
return "hexagon"
elif node.node_type in ("agent", "tool", "workflow_agent", "workflow_handoff"):
return "ellipse"
elif node.node_type == "workflow_base":
return "diamond" # Workflow base uses diamond
elif node.node_type == "workflow_agent":
return "ellipse" # Workflow agents use ellipse
return "diamond"
elif node.node_type == "workflow_tool":
return "box" # Workflow tools use box (matching original square behavior)
elif node.node_type == "workflow_handoff":
return "ellipse" # Handoff nodes use ellipse
return "box"
else:
return "box" # Default shape
return "box"
def _render_pyvis(
graph: DrawWorkflowGraph, filename: str, notebook: bool = False
graph: WorkflowGraph,
filename: str,
notebook: bool = False,
max_label_length: int | None = None,
) -> None:
"""Render workflow graph using Pyvis."""
@@ -96,17 +103,43 @@ def _render_pyvis(
for node in graph.nodes:
color = _get_node_color(node)
shape = _get_node_shape(node)
# Compute display label (with optional truncation)
display_label = node.label
if max_label_length:
display_label = _truncate_label(node.label, max_label_length)
# Build title - show full label if truncated, plus resource metadata
title: str | None = None
if max_label_length and len(node.label) > max_label_length:
title = node.label # Show full label on hover
if isinstance(node, WorkflowResourceNode):
title_parts = [f"Type: {node.type_name or 'Unknown'}"]
if node.getter_name:
title_parts.append(f"Getter: {node.getter_name}")
if node.source_file:
location = node.source_file
if node.source_line:
location += f":{node.source_line}"
title_parts.append(f"Source: {location}")
if node.description:
title_parts.append(f"Doc: {node.description[:100]}...")
title = "\n".join(title_parts)
net.add_node(
node.id,
label=node.label,
title=node.title,
label=display_label,
title=title,
color=color,
shape=shape,
)
# Add edges
for edge in graph.edges:
net.add_edge(edge.source, edge.target)
if edge.label:
net.add_edge(edge.source, edge.target, label=edge.label)
else:
net.add_edge(edge.source, edge.target)
net.show(filename, notebook=notebook)
@@ -130,27 +163,26 @@ def _clean_id_for_mermaid(name: str) -> str:
return name.replace(" ", "_").replace("-", "_").replace(".", "_")
def _get_mermaid_css_class(node: DrawWorkflowNode) -> str:
def _get_mermaid_css_class(node: WorkflowGraphNode) -> str:
"""Determine CSS class for a node in Mermaid based on its type and event_type."""
if node.node_type == "step":
return "stepStyle"
elif node.node_type == "external":
return "externalStyle"
elif node.node_type == "event" and node.event_type:
if issubclass(node.event_type, StartEvent):
elif node.node_type == "resource":
return "resourceStyle"
elif node.node_type == "event":
if node.is_subclass_of("StartEvent"):
return "startEventStyle"
elif issubclass(node.event_type, StopEvent):
elif node.is_subclass_of("StopEvent"):
return "stopEventStyle"
else:
return "defaultEventStyle"
return "defaultEventStyle"
elif node.node_type == "agent":
# Determine class based on agent type
if node.event_type and issubclass(node.event_type, ReActAgent):
if node.is_subclass_of("ReActAgent"):
return "reactAgentStyle"
elif node.event_type and issubclass(node.event_type, CodeActAgent):
elif node.is_subclass_of("CodeActAgent"):
return "codeActAgentStyle"
else:
return "defaultAgentStyle"
return "defaultAgentStyle"
elif node.node_type == "tool":
return "toolStyle"
elif node.node_type == "workflow_base":
@@ -165,88 +197,73 @@ def _get_mermaid_css_class(node: DrawWorkflowNode) -> str:
return "defaultEventStyle"
def _render_mermaid(graph: DrawWorkflowGraph, filename: str) -> str:
def _get_clean_node_id(node: WorkflowGraphNode) -> str:
"""Get a clean Mermaid-compatible ID for a node."""
return f"{node.node_type}_{_clean_id_for_mermaid(node.id)}"
def _get_mermaid_shape(shape: str) -> tuple[str, str]:
"""Get Mermaid shape delimiters for a given shape."""
if shape == "box":
return "[", "]"
elif shape == "ellipse":
return "([", "])"
elif shape == "diamond":
return "{", "}"
elif shape == "hexagon":
return "{{", "}}"
else:
return "[", "]"
def _render_mermaid(
graph: WorkflowGraph, filename: str, max_label_length: int | None = None
) -> str:
"""Render workflow graph using Mermaid."""
mermaid_lines = ["flowchart TD"]
added_nodes = set()
added_edges = set()
added_nodes: set[str] = set()
added_edges: set[str] = set()
# Build lookup dictionary for all nodes
node_by_id: dict[str, WorkflowGraphNode] = {node.id: node for node in graph.nodes}
# Add nodes
for node in graph.nodes:
# Clean ID for Mermaid
if node.node_type == "step":
clean_id = f"step_{_clean_id_for_mermaid(node.id)}"
elif node.node_type == "external":
clean_id = node.id # external_step is already clean
elif node.node_type in [
"agent",
"tool",
"workflow_base",
"workflow_agent",
"workflow_tool",
"workflow_handoff",
]:
clean_id = _clean_id_for_mermaid(node.id)
else: # event
clean_id = f"event_{_clean_id_for_mermaid(node.id)}"
clean_id = _get_clean_node_id(node)
if clean_id not in added_nodes:
added_nodes.add(clean_id)
# Format node based on shape
# Compute display label (with optional truncation)
display_label = node.label
if max_label_length:
display_label = _truncate_label(node.label, max_label_length)
shape = _get_node_shape(node)
if shape == "box":
shape_start, shape_end = "[", "]"
elif shape == "ellipse":
shape_start, shape_end = "([", "])"
elif shape == "diamond":
shape_start, shape_end = "{", "}"
else:
shape_start, shape_end = "[", "]"
shape_start, shape_end = _get_mermaid_shape(shape)
css_class = _get_mermaid_css_class(node)
mermaid_lines.append(
f' {clean_id}{shape_start}"{node.label}"{shape_end}:::{css_class}'
f' {clean_id}{shape_start}"{display_label}"{shape_end}:::{css_class}'
)
# Add edges
for edge in graph.edges:
source_node = next(n for n in graph.nodes if n.id == edge.source)
target_node = next(n for n in graph.nodes if n.id == edge.target)
source_node = node_by_id.get(edge.source)
target_node = node_by_id.get(edge.target)
if source_node.node_type == "step":
source_id = f"step_{_clean_id_for_mermaid(edge.source)}"
elif source_node.node_type == "external":
source_id = edge.source
elif source_node.node_type in [
"agent",
"tool",
"workflow_base",
"workflow_agent",
"workflow_tool",
"workflow_handoff",
]:
source_id = _clean_id_for_mermaid(edge.source)
else: # event
source_id = f"event_{_clean_id_for_mermaid(edge.source)}"
if source_node is None or target_node is None:
continue
if target_node.node_type == "step":
target_id = f"step_{_clean_id_for_mermaid(edge.target)}"
elif target_node.node_type == "external":
target_id = edge.target
elif target_node.node_type in [
"agent",
"tool",
"workflow_base",
"workflow_agent",
"workflow_tool",
"workflow_handoff",
]:
target_id = _clean_id_for_mermaid(edge.target)
else: # event
target_id = f"event_{_clean_id_for_mermaid(edge.target)}"
source_id = _get_clean_node_id(source_node)
target_id = _get_clean_node_id(target_node)
# Handle edge labels (e.g., variable names for resources)
if edge.label:
edge_str = f'{source_id} -->|"{edge.label}"| {target_id}'
else:
edge_str = f"{source_id} --> {target_id}"
edge_str = f"{source_id} --> {target_id}"
if edge_str not in added_edges:
added_edges.add(edge_str)
mermaid_lines.append(f" {edge_str}")
@@ -256,6 +273,7 @@ def _render_mermaid(graph: DrawWorkflowGraph, filename: str) -> str:
[
" classDef stepStyle fill:#ADD8E6,color:#000000,line-height:1.2",
" classDef externalStyle fill:#BEDAE4,color:#000000,line-height:1.2",
" classDef resourceStyle fill:#DDA0DD,color:#000000,line-height:1.2",
" classDef startEventStyle fill:#E27AFF,color:#000000",
" classDef stopEventStyle fill:#FFA07A,color:#000000",
" classDef defaultEventStyle fill:#90EE90,color:#000000",
@@ -279,17 +297,30 @@ def _render_mermaid(graph: DrawWorkflowGraph, filename: str) -> str:
return diagram_string
def _extract_single_agent_structure(agent: BaseWorkflowAgent) -> DrawWorkflowGraph:
def _get_type_chain(cls: type, base: type) -> list[str]:
"""Get type inheritance chain up to (but not including) base class."""
names: list[str] = [cls.__name__]
for parent in cls.mro()[1:]:
if parent is base:
break
if isinstance(parent, type) and issubclass(parent, base):
names.append(parent.__name__)
return names
def _extract_single_agent_structure(agent: BaseWorkflowAgent) -> WorkflowGraph:
"""Extract the structure of a single agent."""
nodes = []
edges = []
nodes: List[WorkflowGraphNode] = []
edges: List[WorkflowGraphEdge] = []
# Add agent node
agent_node = DrawWorkflowNode(
agent_type = type(agent)
agent_node = WorkflowGenericNode(
id="agent",
label=agent.name,
node_type="agent",
event_type=type(agent), # Store agent type for color determination
event_type=agent_type.__name__,
event_types=_get_type_chain(agent_type, BaseWorkflowAgent),
)
nodes.append(agent_node)
@@ -298,7 +329,7 @@ def _extract_single_agent_structure(agent: BaseWorkflowAgent) -> DrawWorkflowGra
if tools is not None and len(tools) > 0:
for i, tool in enumerate(tools):
tool_id = f"tool_{i}"
tool_node = DrawWorkflowNode(
tool_node = WorkflowGenericNode(
id=tool_id,
label=f"Tool {i + 1}: {tool.metadata.get_name()}",
node_type="tool",
@@ -306,27 +337,27 @@ def _extract_single_agent_structure(agent: BaseWorkflowAgent) -> DrawWorkflowGra
nodes.append(tool_node)
# Add edge from agent to tool
edges.append(DrawWorkflowEdge("agent", tool_id))
edges.append(WorkflowGraphEdge(source="agent", target=tool_id))
return DrawWorkflowGraph(nodes=nodes, edges=edges)
return WorkflowGraph(nodes=nodes, edges=edges)
def _process_tools_and_handoffs(
agent: BaseWorkflowAgent,
processed_agents: List[str],
all_agents: Dict[str, BaseWorkflowAgent],
nodes: List[DrawWorkflowNode],
edges: List[DrawWorkflowEdge],
nodes: List[WorkflowGraphNode],
edges: List[WorkflowGraphEdge],
root_agent: str,
) -> Tuple[List[DrawWorkflowNode], List[DrawWorkflowEdge], List[str]]:
) -> Tuple[List[WorkflowGraphNode], List[WorkflowGraphEdge], List[str]]:
if agent.name not in processed_agents:
nodes.append(
DrawWorkflowNode(
WorkflowGenericNode(
id=agent.name, label=agent.name, node_type="workflow_agent"
)
)
if agent.name == root_agent:
edges.append(DrawWorkflowEdge("user", root_agent))
edges.append(WorkflowGraphEdge(source="user", target=root_agent))
for t in agent.tools or []:
if isinstance(t, BaseTool):
fn_name = t.metadata.get_name()
@@ -335,28 +366,18 @@ def _process_tools_and_handoffs(
fn_name = getattr(t, "__name__", type(t).__name__)
node_id = f"{agent.name}_{fn_name}"
nodes.append(
DrawWorkflowNode(
WorkflowGenericNode(
id=node_id,
label=fn_name,
node_type="workflow_tool",
)
)
edges.append(DrawWorkflowEdge(agent.name, node_id))
edges.append(WorkflowGraphEdge(source=agent.name, target=node_id))
if agent.can_handoff_to:
for a in agent.can_handoff_to:
edges.append(
DrawWorkflowEdge(
agent.name,
a,
)
)
edges.append(WorkflowGraphEdge(source=agent.name, target=a))
else:
edges.append(
DrawWorkflowEdge(
agent.name,
"output",
)
)
edges.append(WorkflowGraphEdge(source=agent.name, target="output"))
processed_agents.append(agent.name)
if agent.can_handoff_to:
@@ -376,18 +397,18 @@ def _process_tools_and_handoffs(
def _extract_agent_workflow_structure(
agent_workflow: AgentWorkflow,
) -> DrawWorkflowGraph:
) -> WorkflowGraph:
"""Extract the structure of an agent workflow."""
nodes: List[DrawWorkflowNode] = []
edges: List[DrawWorkflowEdge] = []
nodes: List[WorkflowGraphNode] = []
edges: List[WorkflowGraphEdge] = []
# Add base workflow node
user_node = DrawWorkflowNode(
user_node = WorkflowGenericNode(
id="user",
label="User",
node_type="workflow_base",
)
output_node = DrawWorkflowNode(
output_node = WorkflowGenericNode(
id="output", label="Output", node_type="workflow_base"
)
nodes.extend([user_node, output_node])
@@ -405,9 +426,9 @@ def _extract_agent_workflow_structure(
)
if all(edge.target != "output" for edge in edges):
agent_nodes = [n for n in nodes if n.node_type == "workflow_agent"]
edges.append(DrawWorkflowEdge(agent_nodes[-1].id, "output"))
edges.append(WorkflowGraphEdge(source=agent_nodes[-1].id, target="output"))
return DrawWorkflowGraph(nodes=nodes, edges=edges)
return WorkflowGraph(nodes=nodes, edges=edges)
def _extract_execution_graph(
@@ -490,8 +511,8 @@ def draw_all_possible_flows(
max_label_length: Maximum label length before truncation (None = no limit)
"""
graph = _extract_workflow_structure(workflow, max_label_length)
_render_pyvis(graph, filename, notebook)
graph = _extract_workflow_structure(workflow)
_render_pyvis(graph, filename, notebook, max_label_length)
def draw_all_possible_flows_mermaid(
@@ -511,8 +532,8 @@ def draw_all_possible_flows_mermaid(
The Mermaid diagram as a string
"""
graph = _extract_workflow_structure(workflow, max_label_length)
return _render_mermaid(graph, filename)
graph = _extract_workflow_structure(workflow)
return _render_mermaid(graph, filename, max_label_length)
def draw_agent_with_tools(
@@ -680,6 +701,7 @@ def draw_most_recent_execution_mermaid(
styles = [
"classDef stepStyle fill:#ADD8E6,color:#000000,line-height:1.2",
"classDef externalStyle fill:#BEDAE4,color:#000000,line-height:1.2",
"classDef resourceStyle fill:#DDA0DD,color:#000000,line-height:1.2",
"classDef startEventStyle fill:#E27AFF,color:#000000",
"classDef stopEventStyle fill:#FFA07A,color:#000000",
"classDef defaultEventStyle fill:#90EE90,color:#000000",
@@ -1,7 +1,10 @@
from typing import Annotated
import pytest
from pydantic import Field
from workflows.decorators import step
from workflows.events import Event, StartEvent, StopEvent
from workflows.resource import Resource
from workflows.workflow import Workflow
@@ -31,11 +34,69 @@ class DummyWorkflow(Workflow):
return StopEvent(result="Workflow completed")
# --- Resource-based workflow for testing ---
class DatabaseClient:
"""A mock database client for testing resources."""
pass
def get_database_client() -> DatabaseClient:
"""Factory function to create a database client.
This is a test docstring that should appear in the resource metadata.
"""
return DatabaseClient()
class CacheClient:
"""A mock cache client for testing resources."""
pass
def get_cache_client() -> CacheClient:
"""Factory function to create a cache client."""
return CacheClient()
class ResourceWorkflow(Workflow):
"""A workflow with resource dependencies for testing visualization."""
@step()
async def start_step(self, ev: StartEvent) -> OneTestEvent:
return OneTestEvent()
@step()
async def step_with_db(
self,
ev: OneTestEvent,
db_client: Annotated[DatabaseClient, Resource(get_database_client)],
) -> LastEvent:
return LastEvent()
@step()
async def step_with_both_resources(
self,
ev: LastEvent,
db: Annotated[DatabaseClient, Resource(get_database_client)],
cache: Annotated[CacheClient, Resource(get_cache_client)],
) -> StopEvent:
return StopEvent(result="Workflow completed")
@pytest.fixture()
def workflow() -> Workflow:
return DummyWorkflow()
@pytest.fixture()
def workflow_with_resources() -> Workflow:
return ResourceWorkflow()
@pytest.fixture()
def events() -> list[type[Event]]:
return [OneTestEvent, AnotherTestEvent]
@@ -29,8 +29,7 @@ async def test_workflow_draw_methods(workflow: Workflow) -> None:
)
@pytest.mark.asyncio
async def test_draw_all_possible_flows_with_max_label_length(
def test_draw_all_possible_flows_with_max_label_length(
workflow: Workflow,
) -> None:
"""Test the max_label_length parameter."""
@@ -80,8 +79,7 @@ async def test_draw_all_possible_flows_with_max_label_length(
)
@pytest.mark.asyncio
async def test_draw_all_possible_flows_mermaid_basic(workflow: Workflow) -> None:
def test_draw_all_possible_flows_mermaid_basic(workflow: Workflow) -> None:
"""Test basic Mermaid diagram generation."""
with patch("builtins.open", mock_open()) as mock_file:
result = draw_all_possible_flows_mermaid(
@@ -103,8 +101,7 @@ async def test_draw_all_possible_flows_mermaid_basic(workflow: Workflow) -> None
assert "classDef externalStyle fill:#BEDAE4" in result
@pytest.mark.asyncio
async def test_draw_all_possible_flows_mermaid_no_file(workflow: Workflow) -> None:
def test_draw_all_possible_flows_mermaid_no_file(workflow: Workflow) -> None:
"""Test Mermaid diagram generation without file output."""
result = draw_all_possible_flows_mermaid(workflow)
@@ -113,8 +110,7 @@ async def test_draw_all_possible_flows_mermaid_no_file(workflow: Workflow) -> No
assert result.startswith("flowchart TD")
@pytest.mark.asyncio
async def test_mermaid_node_shapes_and_styles(workflow: Workflow) -> None:
def test_mermaid_node_shapes_and_styles(workflow: Workflow) -> None:
"""Test that Mermaid nodes have correct shapes and styles."""
result = draw_all_possible_flows_mermaid(workflow)
@@ -147,8 +143,7 @@ async def test_mermaid_node_shapes_and_styles(workflow: Workflow) -> None:
)
@pytest.mark.asyncio
async def test_mermaid_edges_generation(workflow: Workflow) -> None:
def test_mermaid_edges_generation(workflow: Workflow) -> None:
"""Test that Mermaid edges are properly generated."""
result = draw_all_possible_flows_mermaid(workflow)
@@ -168,8 +163,7 @@ async def test_mermaid_edges_generation(workflow: Workflow) -> None:
assert target.strip(), f"Edge target should not be empty: {edge_line}"
@pytest.mark.asyncio
async def test_mermaid_id_cleaning(workflow: Workflow) -> None:
def test_mermaid_id_cleaning(workflow: Workflow) -> None:
"""Test that Mermaid IDs are properly cleaned for validity."""
result = draw_all_possible_flows_mermaid(workflow)
@@ -190,8 +184,7 @@ async def test_mermaid_id_cleaning(workflow: Workflow) -> None:
# Note: We allow underscores as they're valid in Mermaid
@pytest.mark.asyncio
async def test_mermaid_vs_pyvis_consistency(workflow: Workflow) -> None:
def test_mermaid_vs_pyvis_consistency(workflow: Workflow) -> None:
"""Test that Mermaid and Pyvis generate consistent node/edge counts."""
# Generate Pyvis version
with patch("llama_index.utils.workflow.Network") as mock_network:
@@ -237,8 +230,7 @@ async def test_mermaid_vs_pyvis_consistency(workflow: Workflow) -> None:
)
@pytest.mark.asyncio
async def test_mermaid_file_writing(workflow: Workflow) -> None:
def test_mermaid_file_writing(workflow: Workflow) -> None:
"""Test that Mermaid diagram is correctly written to file."""
mock_file_handle = mock_open()
@@ -261,8 +253,7 @@ async def test_mermaid_file_writing(workflow: Workflow) -> None:
)
@pytest.mark.asyncio
async def test_mermaid_empty_filename(workflow: Workflow) -> None:
def test_mermaid_empty_filename(workflow: Workflow) -> None:
"""Test that Mermaid works with empty/None filename."""
# Test without filename (defaults internally)
result1 = draw_all_possible_flows_mermaid(workflow)
@@ -309,3 +300,156 @@ async def test_draw_most_recent_execution_mermaid(workflow: Workflow) -> None:
edge_lines = [line for line in lines if " --> " in line]
assert len(node_lines) > 0
assert len(edge_lines) > 0
# --- Resource node rendering tests ---
def test_mermaid_resource_nodes_rendered(
workflow_with_resources: Workflow,
) -> None:
"""Test that resource nodes are rendered in Mermaid output."""
result = draw_all_possible_flows_mermaid(workflow_with_resources)
# Verify resource style is defined
assert "classDef resourceStyle fill:#DDA0DD" in result
# Verify resource nodes are present (hexagon shape with {{}})
lines = result.split("\n")
resource_lines = [line for line in lines if "resource_" in line and ":::" in line]
assert len(resource_lines) > 0
# Check resource nodes use hexagon shape
for line in resource_lines:
assert "{{" in line and "}}" in line, (
f"Resource node should use hexagon shape: {line}"
)
assert ":::resourceStyle" in line, (
f"Resource node should use resourceStyle: {line}"
)
def test_mermaid_resource_edges_have_labels(
workflow_with_resources: Workflow,
) -> None:
"""Test that edges from resources to steps have labels (variable names)."""
result = draw_all_possible_flows_mermaid(workflow_with_resources)
lines = result.split("\n")
# Look for edges with labels: resource_xxx -->|"var_name"| step_yyy
labeled_edge_lines = [line for line in lines if '-->|"' in line]
# Should have labeled edges for resource connections
assert len(labeled_edge_lines) > 0
# Check that the labels are variable names
expected_labels = {"db_client", "db", "cache"}
found_labels = set()
for line in labeled_edge_lines:
# Extract label from -->|"label"|
if '-->|"' in line:
start = line.index('-->|"') + 5
end = line.index('"|', start)
label = line[start:end]
found_labels.add(label)
assert found_labels.intersection(expected_labels), (
f"Expected some of {expected_labels}, found {found_labels}"
)
def test_pyvis_resource_nodes_rendered(workflow_with_resources: Workflow) -> None:
"""Test that resource nodes are rendered in Pyvis output."""
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(workflow_with_resources, filename="test.html")
# Extract all add_node calls
node_calls = mock_net_instance.add_node.call_args_list
# Find resource nodes (should have hexagon shape and plum color)
resource_nodes = []
for call in node_calls:
args, kwargs = call
node_id = args[0]
if "resource_" in node_id:
resource_nodes.append((node_id, kwargs))
assert len(resource_nodes) > 0, "Should have resource nodes"
for node_id, kwargs in resource_nodes:
assert kwargs.get("shape") == "hexagon", (
f"Resource node {node_id} should be hexagon"
)
assert kwargs.get("color") == "#DDA0DD", (
f"Resource node {node_id} should be plum color"
)
# Should have a title with metadata
assert kwargs.get("title") is not None, (
f"Resource node {node_id} should have title"
)
def test_pyvis_resource_edges_have_labels(
workflow_with_resources: Workflow,
) -> None:
"""Test that Pyvis edges from resources have labels."""
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(workflow_with_resources, filename="test.html")
# Extract all add_edge calls
edge_calls = mock_net_instance.add_edge.call_args_list
# Find edges with labels
labeled_edges = []
for call in edge_calls:
args, kwargs = call
if "label" in kwargs:
labeled_edges.append((args, kwargs["label"]))
assert len(labeled_edges) > 0, "Should have labeled edges"
# Check that labels are variable names
labels = {label for _, label in labeled_edges}
expected_labels = {"db_client", "db", "cache"}
assert labels.intersection(expected_labels), (
f"Expected some of {expected_labels}, found {labels}"
)
def test_mermaid_resource_style_always_defined(workflow: Workflow) -> None:
"""Test that resourceStyle is always defined even for workflows without resources."""
result = draw_all_possible_flows_mermaid(workflow)
# resourceStyle should be defined even if not used
assert "classDef resourceStyle fill:#DDA0DD" in result
def test_resource_node_deduplication_in_rendering(
workflow_with_resources: Workflow,
) -> None:
"""Test that deduplicated resource nodes render correctly."""
result = draw_all_possible_flows_mermaid(workflow_with_resources)
lines = result.split("\n")
# Count unique resource node definitions (not edges)
resource_node_defs = [
line
for line in lines
if "resource_" in line
and ":::" in line
and " --> " not in line
and "-->|" not in line
]
# The workflow has 2 unique resources (DatabaseClient used twice, CacheClient once)
# So we should see exactly 2 resource node definitions
assert len(resource_node_defs) == 2, (
f"Expected 2 unique resource nodes, found {len(resource_node_defs)}: {resource_node_defs}"
)
@@ -1,8 +1,8 @@
from __future__ import annotations
from typing import Any, Literal
from typing import Any, Literal, Union
from pydantic import BaseModel
from pydantic import BaseModel, Field
from workflows.protocol.serializable_events import EventEnvelopeWithMetadata
@@ -58,25 +58,241 @@ class WorkflowEventsListResponse(BaseModel):
class WorkflowGraphResponse(BaseModel):
graph: WorkflowGraphNodeEdges
graph: WorkflowGraph
class WorkflowGraphNode(BaseModel):
id: str
label: str
node_type: str
title: str | None
event_type: str | None
class WorkflowNodeBase(BaseModel):
"""Base class for all workflow graph nodes."""
id: str = Field(description="Unique identifier for the node")
label: str = Field(description="Display text for the node")
def truncated_label(self, max_length: int) -> str:
"""Get truncated label for visualization (adds * suffix if truncated)."""
if len(self.label) <= max_length:
return self.label
return f"{self.label[: max_length - 1]}*"
class WorkflowStepNode(WorkflowNodeBase):
"""A workflow step node representing a function decorated with @step."""
node_type: Literal["step"] = Field(
default="step", description="Discriminator field for node type"
)
description: str | None = Field(
default=None,
description="Documentation string extracted from the step function",
)
class WorkflowEventNode(WorkflowNodeBase):
"""An event node representing an Event class that flows between steps."""
node_type: Literal["event"] = Field(
default="event", description="Discriminator field for node type"
)
event_type: str = Field(
description="The event class name (e.g., 'StartEvent', 'MyCustomEvent')"
)
event_types: list[str] = Field(
description="Event class inheritance chain for subclass checking. "
"First element is the class itself, followed by parent Event subclasses."
)
event_schema: dict[str, Any] | None = Field(
default=None,
description="Pydantic JSON schema for the event type",
)
def is_subclass_of(self, *type_names: str) -> bool:
"""Check if this node's event_type is a subclass of any of the given types."""
return any(name in self.event_types for name in type_names)
class WorkflowExternalNode(WorkflowNodeBase):
"""An external node representing human-in-the-loop or external system interaction."""
node_type: Literal["external"] = Field(
default="external", description="Discriminator field for node type"
)
class WorkflowResourceNode(WorkflowNodeBase):
"""A resource node representing an injected dependency (e.g., database client, API client)."""
node_type: Literal["resource"] = Field(
default="resource", description="Discriminator field for node type"
)
type_name: str | None = Field(
default=None,
description="The type annotation of the resource (e.g., 'DatabaseClient', 'AsyncLlamaCloud')",
)
getter_name: str | None = Field(
default=None,
description="Name of the factory function that creates the resource",
)
source_file: str | None = Field(
default=None,
description="Absolute path to the source file containing the getter function",
)
source_line: int | None = Field(
default=None, description="Line number where the getter function is defined"
)
description: str | None = Field(
default=None,
description="Documentation string extracted from the getter function",
)
class WorkflowGenericNode(WorkflowNodeBase):
"""A generic node for custom visualization types not covered by standard node types.
Used for agent visualization (node_type='agent', 'tool', 'workflow_agent', etc.)
and other custom extensions. Supports optional event_type fields for type checking.
"""
node_type: str = Field(
description="Custom node type string (e.g., 'agent', 'tool', 'workflow_base')"
)
event_type: str | None = Field(
default=None,
description="Optional type name for nodes that support inheritance checking (e.g., agent types)",
)
event_types: list[str] | None = Field(
default=None,
description="Optional inheritance chain for subclass checking, similar to WorkflowEventNode",
)
def is_subclass_of(self, *type_names: str) -> bool:
"""Check if this node's event_type is a subclass of any of the given types."""
if not self.event_types:
return False
return any(name in self.event_types for name in type_names)
# Union type for workflow graph nodes
# Pydantic will try to match against types in order; WorkflowGenericNode is last as catch-all
WorkflowGraphNode = Union[
WorkflowStepNode,
WorkflowEventNode,
WorkflowExternalNode,
WorkflowResourceNode,
WorkflowGenericNode,
]
class WorkflowGraphEdge(BaseModel):
source: str
target: str
"""A directed edge connecting two nodes in the workflow graph."""
source: str = Field(description="ID of the source node (where the edge originates)")
target: str = Field(description="ID of the target node (where the edge points to)")
label: str | None = Field(
default=None,
description="Optional edge label, used for resource edges to show the variable name",
)
class WorkflowGraphNodeEdges(BaseModel):
nodes: list[WorkflowGraphNode]
edges: list[WorkflowGraphEdge]
class WorkflowGraph(BaseModel):
"""Complete workflow graph structure containing all nodes and edges."""
nodes: list[WorkflowGraphNode] = Field(
description="All nodes in the workflow graph"
)
edges: list[WorkflowGraphEdge] = Field(
description="All directed edges connecting the nodes"
)
description: str | None = Field(
default=None,
description="Documentation string extracted from the workflow class",
)
def filter_by_node_type(self, *node_types: str) -> WorkflowGraph:
"""Create a simplified graph by removing nodes of specified types.
Edges passing through filtered nodes are resolved:
Node1 -> FilteredNode -> Node2 becomes Node1 -> Node2
Args:
*node_types: One or more node type strings to filter out
(e.g., "event", "resource", "step", "external")
Returns:
A new WorkflowGraph with the specified node types removed
and edges resolved through them.
"""
filter_types = set(node_types)
# Identify nodes to filter out
filtered_node_ids: set[str] = set()
for node in self.nodes:
if node.node_type in filter_types:
filtered_node_ids.add(node.id)
# Keep remaining nodes
remaining_nodes = [n for n in self.nodes if n.id not in filtered_node_ids]
remaining_node_ids = {n.id for n in remaining_nodes}
# Build outgoing edge map and node lookup
outgoing_map: dict[str, list[WorkflowGraphEdge]] = {}
for edge in self.edges:
outgoing_map.setdefault(edge.source, []).append(edge)
node_by_id: dict[str, WorkflowGraphNode] = {n.id: n for n in self.nodes}
def resolve_targets(
from_id: str,
first_filtered_label: str | None,
visited: set[str],
) -> list[tuple[str, str | None]]:
"""Find remaining nodes reachable from from_id, through filtered nodes."""
results: list[tuple[str, str | None]] = []
for edge in outgoing_map.get(from_id, []):
target = edge.target
if target in visited:
continue
if target in remaining_node_ids:
# Use the first filtered node's label, or the edge label if direct
label = (
first_filtered_label
if first_filtered_label is not None
else edge.label
)
results.append((target, label))
elif target in filtered_node_ids:
# Follow through filtered node, capturing its label if first
visited.add(target)
filtered_node = node_by_id[target]
label = (
first_filtered_label
if first_filtered_label is not None
else filtered_node.label
)
results.extend(resolve_targets(target, label, visited))
return results
# Build new edges
new_edges: list[WorkflowGraphEdge] = []
seen_edges: set[tuple[str, str]] = set()
for source_id in remaining_node_ids:
for target_id, label in resolve_targets(source_id, None, set()):
edge_key = (source_id, target_id)
if edge_key not in seen_edges:
seen_edges.add(edge_key)
new_edges.append(
WorkflowGraphEdge(
source=source_id,
target=target_id,
label=label,
)
)
return WorkflowGraph(
nodes=remaining_nodes,
edges=new_edges,
description=self.description,
)
__all__ = [
@@ -90,7 +306,13 @@ __all__ = [
"WorkflowSchemaResponse",
"WorkflowEventsListResponse",
"WorkflowGraphResponse",
"WorkflowNodeBase",
"WorkflowStepNode",
"WorkflowEventNode",
"WorkflowExternalNode",
"WorkflowResourceNode",
"WorkflowGenericNode",
"WorkflowGraphNode",
"WorkflowGraphEdge",
"WorkflowGraphNodeEdges",
"WorkflowGraph",
]
@@ -66,8 +66,7 @@ class EventEnvelopeWithMetadata(BaseModel):
class EventEnvelope(BaseModel):
"""
Client write representation of an Event. Includes class metadata in order to support
matching event types semantically in an extendable manner (e.g. "StartEvent", "StopEvent", etc.).
Client write representation of an Event. Simpler than the server provided EventEnvelopeWithMetadata, as the metadata can be inferred based on looking up the runtime type
"""
value: Any | None
@@ -1,93 +1,110 @@
from dataclasses import dataclass
from typing import List, Optional
from __future__ import annotations
import hashlib
import inspect
from workflows import Workflow
from workflows.decorators import StepConfig, StepFunction
from workflows.events import (
Event,
HumanResponseEvent,
InputRequiredEvent,
StopEvent,
)
from workflows.protocol import (
WorkflowEventNode,
WorkflowExternalNode,
WorkflowGraph,
WorkflowGraphEdge,
WorkflowGraphNode,
WorkflowGraphNodeEdges,
WorkflowResourceNode,
WorkflowStepNode,
)
from workflows.resource import ResourceDefinition
from workflows.utils import (
get_steps_from_class,
get_steps_from_instance,
)
@dataclass
class DrawWorkflowNode:
"""Represents a node in the workflow graph."""
def _get_event_type_chain(cls: type) -> list[str]:
"""Get the event type inheritance chain including the class itself.
id: str
label: str
node_type: str # 'step', 'event', 'external'
title: Optional[str] = None
event_type: Optional[type] = (
None # Store the actual event type for styling decisions
Returns a list starting with the class name, followed by parent Event
subclasses up to (but not including) Event itself.
"""
names: list[str] = [cls.__name__]
for parent in cls.mro()[1:]:
if parent is Event:
break
if isinstance(parent, type) and issubclass(parent, Event):
names.append(parent.__name__)
return names
def _create_resource_node(resource_def: ResourceDefinition) -> WorkflowResourceNode:
"""Create a WorkflowResourceNode from a ResourceDefinition.
Extracts metadata (source file, line number, docstring) lazily here
rather than at Resource creation time for performance.
"""
resource = resource_def.resource
factory = resource._factory
# Get type name from annotation
type_name: str | None = None
if resource_def.type_annotation is not None:
type_annotation = resource_def.type_annotation
if hasattr(type_annotation, "__name__"):
type_name = type_annotation.__name__
else:
type_name = str(type_annotation)
# Extract source metadata lazily
source_file: str | None = None
source_line: int | None = None
try:
source_file = inspect.getfile(factory)
except (TypeError, OSError):
pass
try:
_, source_line = inspect.getsourcelines(factory)
except (TypeError, OSError):
pass
resource_description = inspect.getdoc(factory)
# Compute unique hash for deduplication
hash_input = f"{resource.name}:{source_file or 'unknown'}"
unique_hash = hashlib.sha256(hash_input.encode()).hexdigest()[:12]
# Label: prefer type_name, then getter_name, then id
node_id = f"resource_{unique_hash}"
label = type_name or resource.name or node_id
return WorkflowResourceNode(
id=node_id,
label=label,
type_name=type_name,
getter_name=resource.name,
source_file=source_file,
source_line=source_line,
description=resource_description,
)
def to_response_model(self) -> WorkflowGraphNode:
return WorkflowGraphNode(
id=self.id,
label=self.label,
node_type=self.node_type,
title=self.title,
event_type=self.event_type.__name__ if self.event_type else None,
)
@dataclass
class DrawWorkflowEdge:
"""Represents an edge in the workflow graph."""
source: str
target: str
def to_response_model(self) -> WorkflowGraphEdge:
return WorkflowGraphEdge(
source=self.source,
target=self.target,
)
@dataclass
class DrawWorkflowGraph:
"""Intermediate representation of workflow structure."""
nodes: List[DrawWorkflowNode]
edges: List[DrawWorkflowEdge]
def to_response_model(self) -> WorkflowGraphNodeEdges:
return WorkflowGraphNodeEdges(
nodes=[node.to_response_model() for node in self.nodes],
edges=[edge.to_response_model() for edge in self.edges],
)
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: Optional[int] = None
) -> DrawWorkflowGraph:
"""Extract workflow structure into an intermediate representation."""
def extract_workflow_structure(workflow: Workflow) -> WorkflowGraph:
"""Extract workflow structure into a graph representation."""
# Get workflow steps
steps: dict[str, StepFunction] = get_steps_from_class(workflow)
if not steps:
steps = get_steps_from_instance(workflow)
nodes = []
edges = []
added_nodes = set() # Track added node IDs to avoid duplicates
nodes: list[WorkflowGraphNode] = []
edges: list[WorkflowGraphEdge] = []
added_nodes: set[str] = set() # Track added node IDs to avoid duplicates
added_resource_nodes: dict[int, WorkflowResourceNode] = {} # Track by factory id
step_config: Optional[StepConfig] = None
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.
@@ -108,24 +125,11 @@ def extract_workflow_structure(
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:
step_description = inspect.getdoc(step_func)
nodes.append(
DrawWorkflowNode(
id=step_name,
label=step_label,
node_type="step",
title=step_title,
WorkflowStepNode(
id=step_name, label=step_name, description=step_description
)
)
added_nodes.add(step_name)
@@ -135,25 +139,14 @@ def extract_workflow_structure(
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(
WorkflowEventNode(
id=event_type.__name__,
label=event_label,
node_type="event",
title=event_title,
event_type=event_type,
label=event_type.__name__,
event_type=event_type.__name__,
event_types=_get_event_type_chain(event_type),
event_schema=event_type.model_json_schema(),
)
)
added_nodes.add(event_type.__name__)
@@ -163,25 +156,14 @@ def extract_workflow_structure(
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(
WorkflowEventNode(
id=return_type.__name__,
label=return_label,
node_type="event",
title=return_title,
event_type=return_type,
label=return_type.__name__,
event_type=return_type.__name__,
event_types=_get_event_type_chain(return_type),
event_schema=return_type.model_json_schema(),
)
)
added_nodes.add(return_type.__name__)
@@ -192,14 +174,18 @@ def extract_workflow_structure(
and "external_step" not in added_nodes
):
nodes.append(
DrawWorkflowNode(
id="external_step",
label="external_step",
node_type="external",
)
WorkflowExternalNode(id="external_step", label="external_step")
)
added_nodes.add("external_step")
# Add resource nodes (deduplicated by factory identity)
for resource_def in step_config.resources:
factory_id = id(resource_def.resource._factory)
if factory_id not in added_resource_nodes:
resource_node = _create_resource_node(resource_def)
nodes.append(resource_node)
added_resource_nodes[factory_id] = resource_node
# Second pass: Add edges
for step_name, step_func in steps.items():
step_config = step_func._step_config
@@ -207,22 +193,49 @@ def extract_workflow_structure(
# 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__))
edges.append(
WorkflowGraphEdge(source=step_name, target=return_type.__name__)
)
if issubclass(return_type, InputRequiredEvent):
edges.append(DrawWorkflowEdge(return_type.__name__, "external_step"))
edges.append(
WorkflowGraphEdge(
source=return_type.__name__, target="external_step"
)
)
# Edges from events to steps
for event_type in step_config.accepted_events:
if step_name == "_done" and issubclass(event_type, StopEvent):
if current_stop_event:
edges.append(
DrawWorkflowEdge(current_stop_event.__name__, step_name)
WorkflowGraphEdge(
source=current_stop_event.__name__, target=step_name
)
)
else:
edges.append(DrawWorkflowEdge(event_type.__name__, step_name))
edges.append(
WorkflowGraphEdge(source=event_type.__name__, target=step_name)
)
if issubclass(event_type, HumanResponseEvent):
edges.append(DrawWorkflowEdge("external_step", event_type.__name__))
edges.append(
WorkflowGraphEdge(
source="external_step", target=event_type.__name__
)
)
return DrawWorkflowGraph(nodes=nodes, edges=edges)
# Edges from steps to resources (with variable name as label)
for resource_def in step_config.resources:
factory_id = id(resource_def.resource._factory)
resource_node = added_resource_nodes[factory_id]
edges.append(
WorkflowGraphEdge(
source=step_name,
target=resource_node.id,
label=resource_def.name, # The variable name
)
)
workflow_description = inspect.getdoc(workflow)
return WorkflowGraph(nodes=nodes, edges=edges, description=workflow_description)
@@ -49,11 +49,13 @@ class ResourceDefinition(BaseModel):
Attributes:
name (str): Parameter name in the step function.
resource (_Resource): Factory wrapper used by the manager to produce the dependency.
type_annotation (type | None): The type annotation from Annotated[T, Resource(...)].
"""
model_config = ConfigDict(arbitrary_types_allowed=True)
name: str
resource: _Resource
type_annotation: Any = None
def Resource(factory: Callable[..., T], cache: bool = True) -> _Resource[T]:
@@ -629,9 +629,7 @@ class WorkflowServer:
detail=f"Error while getting JSON workflow representation: {e}",
status_code=500,
)
return JSONResponse(
WorkflowGraphResponse(graph=workflow_graph.to_response_model()).model_dump()
)
return JSONResponse(WorkflowGraphResponse(graph=workflow_graph).model_dump())
async def _run_workflow_nowait(self, request: Request) -> JSONResponse:
"""
@@ -102,8 +102,15 @@ def inspect_signature(fn: Callable) -> StepSignatureSpec:
# Handle Annotated types for resources
if get_origin(annotation) is Annotated:
_, resource = get_args(annotation)
resources.append(ResourceDefinition(name=name, resource=resource))
args = get_args(annotation)
type_annotation = args[0] if args else None
resource = args[1] if len(args) > 1 else None
if resource is not None:
resources.append(
ResourceDefinition(
name=name, resource=resource, type_annotation=type_annotation
)
)
continue
# Get name and type of the Context param (without state type)
@@ -1,84 +1,79 @@
import pytest
from workflows.events import StartEvent, StopEvent
from workflows.representation_utils import (
DrawWorkflowEdge,
DrawWorkflowGraph,
DrawWorkflowNode,
extract_workflow_structure,
)
from typing import Annotated
from .conftest import DummyWorkflow, LastEvent, OneTestEvent # type: ignore[import]
import pytest
from workflows.decorators import step
from workflows.events import Event, StartEvent, StopEvent
from workflows.protocol import (
WorkflowEventNode,
WorkflowExternalNode,
WorkflowGraph,
WorkflowGraphEdge,
WorkflowResourceNode,
WorkflowStepNode,
)
from workflows.representation_utils import extract_workflow_structure
from workflows.resource import Resource
from workflows.workflow import Workflow
from .conftest import DummyWorkflow # type: ignore[import]
@pytest.fixture()
def ground_truth_repr() -> DrawWorkflowGraph:
return DrawWorkflowGraph(
def ground_truth_repr() -> WorkflowGraph:
return WorkflowGraph(
nodes=[
DrawWorkflowNode(
WorkflowStepNode(
id="end_step",
label="end_step",
node_type="step",
title=None,
event_type=None,
),
DrawWorkflowNode(
WorkflowEventNode(
id="LastEvent",
label="LastEvent",
node_type="event",
title=None,
event_type=LastEvent,
event_type="LastEvent",
event_types=["LastEvent"],
),
DrawWorkflowNode(
WorkflowEventNode(
id="StopEvent",
label="StopEvent",
node_type="event",
title=None,
event_type=StopEvent,
event_type="StopEvent",
event_types=["StopEvent"],
),
DrawWorkflowNode(
WorkflowStepNode(
id="middle_step",
label="middle_step",
node_type="step",
title=None,
event_type=None,
),
DrawWorkflowNode(
WorkflowEventNode(
id="OneTestEvent",
label="OneTestEvent",
node_type="event",
title=None,
event_type=OneTestEvent,
event_type="OneTestEvent",
event_types=["OneTestEvent"],
),
DrawWorkflowNode(
WorkflowStepNode(
id="start_step",
label="start_step",
node_type="step",
title=None,
event_type=None,
),
DrawWorkflowNode(
WorkflowEventNode(
id="StartEvent",
label="StartEvent",
node_type="event",
title=None,
event_type=StartEvent,
event_type="StartEvent",
event_types=["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"),
WorkflowGraphEdge(source="end_step", target="StopEvent"),
WorkflowGraphEdge(source="LastEvent", target="end_step"),
WorkflowGraphEdge(source="middle_step", target="LastEvent"),
WorkflowGraphEdge(source="OneTestEvent", target="middle_step"),
WorkflowGraphEdge(source="start_step", target="OneTestEvent"),
WorkflowGraphEdge(source="StartEvent", target="start_step"),
],
)
def test_extract_workflow_structure(ground_truth_repr: DrawWorkflowGraph) -> None:
def test_extract_workflow_structure(ground_truth_repr: WorkflowGraph) -> None:
wf = DummyWorkflow()
graph = extract_workflow_structure(workflow=wf)
assert isinstance(graph, DrawWorkflowGraph)
assert isinstance(graph, WorkflowGraph)
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"])
@@ -90,45 +85,725 @@ def test_extract_workflow_structure(ground_truth_repr: DrawWorkflowGraph) -> Non
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_truncated_label() -> None:
"""Test that truncated_label method works correctly."""
node = WorkflowStepNode(id="my_step", label="my_long_step_name")
assert node.truncated_label(5) == "my_l*"
assert node.truncated_label(20) == "my_long_step_name"
assert node.truncated_label(17) == "my_long_step_name"
def test_graph_to_response_model() -> None:
graph = DrawWorkflowGraph(
def test_graph_serialization() -> None:
"""Test that WorkflowGraphNodeEdges serializes correctly to JSON."""
graph = WorkflowGraph(
nodes=[
DrawWorkflowNode(
id="test", label="test", node_type="step", title=None, event_type=None
),
DrawWorkflowNode(
WorkflowStepNode(id="test", label="test"),
WorkflowEventNode(
id="OneTestEvent",
label="OneTestEvent",
node_type="event",
title=None,
event_type=OneTestEvent,
event_type="OneTestEvent",
event_types=["OneTestEvent"],
),
],
edges=[DrawWorkflowEdge(source="test", target="OneTestEvent")],
edges=[WorkflowGraphEdge(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"
# Test direct access
assert len(graph.nodes) == 2
step_node = graph.nodes[0]
assert isinstance(step_node, WorkflowStepNode)
assert step_node.node_type == "step"
assert step_node.label == "test"
assert step_node.id == "test"
event_node = graph.nodes[1]
assert isinstance(event_node, WorkflowEventNode)
assert event_node.event_type == "OneTestEvent"
assert event_node.event_types == ["OneTestEvent"]
assert event_node.node_type == "event"
assert event_node.label == "OneTestEvent"
assert event_node.id == "OneTestEvent"
assert len(graph.edges) == 1
assert graph.edges[0].source == "test"
assert graph.edges[0].target == "OneTestEvent"
# Test JSON serialization (round-trip works)
data = graph.model_dump()
assert "event_type" not in data["nodes"][0] # Step nodes don't have event_type
assert data["nodes"][1]["event_type"] == "OneTestEvent"
assert data["nodes"][1]["event_types"] == ["OneTestEvent"]
# Test deserialization
restored = WorkflowGraph.model_validate(data)
restored_event = restored.nodes[1]
assert isinstance(restored_event, WorkflowEventNode)
assert restored_event.event_type == "OneTestEvent"
assert restored_event.is_subclass_of("OneTestEvent")
# --- Resource node tests ---
class DatabaseClient:
"""A mock database client for testing resources."""
pass
def get_database_client() -> DatabaseClient:
"""Factory function to create a database client.
This docstring should appear in the resource metadata.
"""
return DatabaseClient()
class MiddleEvent(Event):
pass
class WorkflowWithResources(Workflow):
@step
async def start_step(self, ev: StartEvent) -> MiddleEvent:
return MiddleEvent()
@step
async def step_with_resource(
self,
ev: MiddleEvent,
db_client: Annotated[DatabaseClient, Resource(get_database_client)],
) -> StopEvent:
return StopEvent(result="done")
def test_extract_workflow_structure_with_resources() -> None:
"""Test that resource nodes are extracted from workflow with resources."""
wf = WorkflowWithResources()
graph = extract_workflow_structure(workflow=wf)
# Should have resource nodes
resource_nodes = [n for n in graph.nodes if isinstance(n, WorkflowResourceNode)]
assert len(resource_nodes) == 1
resource_node = resource_nodes[0]
assert resource_node.node_type == "resource"
assert resource_node.type_name == "DatabaseClient"
assert resource_node.getter_name == "get_database_client"
assert resource_node.description is not None
assert "Factory function" in resource_node.description
assert resource_node.source_file is not None
assert resource_node.source_line is not None
def test_resource_node_edges_have_variable_names() -> None:
"""Test that edges from steps to resources have the variable name as label."""
wf = WorkflowWithResources()
graph = extract_workflow_structure(workflow=wf)
# Find edges to resource nodes
resource_edges = [e for e in graph.edges if e.target.startswith("resource_")]
assert len(resource_edges) == 1
edge = resource_edges[0]
assert edge.label == "db_client" # The variable name
assert edge.source == "step_with_resource"
def test_resource_nodes_are_deduplicated() -> None:
"""Test that the same resource used in multiple steps appears only once."""
class StepEvent(Event):
pass
class WorkflowWithSharedResource(Workflow):
@step
async def start_step(self, ev: StartEvent) -> StepEvent:
return StepEvent()
@step
async def step_one(
self,
ev: StepEvent,
db: Annotated[DatabaseClient, Resource(get_database_client)],
) -> MiddleEvent:
return MiddleEvent()
@step
async def step_two(
self,
ev: MiddleEvent,
db: Annotated[DatabaseClient, Resource(get_database_client)],
) -> StopEvent:
return StopEvent(result="done")
wf = WorkflowWithSharedResource()
graph = extract_workflow_structure(workflow=wf)
# Should have only one resource node (deduplicated)
resource_nodes = [n for n in graph.nodes if isinstance(n, WorkflowResourceNode)]
assert len(resource_nodes) == 1
# But should have two edges (one from each step)
resource_edges = [e for e in graph.edges if e.target.startswith("resource_")]
assert len(resource_edges) == 2
# Both edges should have the variable name "db"
for edge in resource_edges:
assert edge.label == "db"
def test_multiple_different_resources() -> None:
"""Test workflow with multiple different resources."""
class CacheClient:
pass
def get_cache_client() -> CacheClient:
return CacheClient()
class WorkflowWithMultipleResources(Workflow):
@step
async def start_step(
self,
ev: StartEvent,
db: Annotated[DatabaseClient, Resource(get_database_client)],
cache: Annotated[CacheClient, Resource(get_cache_client)],
) -> StopEvent:
return StopEvent(result="done")
wf = WorkflowWithMultipleResources()
graph = extract_workflow_structure(workflow=wf)
# Should have two different resource nodes
resource_nodes = [n for n in graph.nodes if isinstance(n, WorkflowResourceNode)]
assert len(resource_nodes) == 2
type_names = {rn.type_name for rn in resource_nodes}
assert type_names == {"DatabaseClient", "CacheClient"}
# Should have two edges with different labels
resource_edges = [e for e in graph.edges if e.target.startswith("resource_")]
assert len(resource_edges) == 2
labels = {e.label for e in resource_edges}
assert labels == {"db", "cache"}
def test_resource_node_serialization() -> None:
"""Test that WorkflowResourceNode serializes correctly."""
resource_node = WorkflowResourceNode(
id="resource_abc123",
label="TestType",
type_name="TestType",
getter_name="get_test_type",
source_file="/path/to/file.py",
source_line=42,
description="Test docstring",
)
assert resource_node.id == "resource_abc123"
assert resource_node.label == "TestType"
assert resource_node.node_type == "resource"
assert resource_node.type_name == "TestType"
assert resource_node.getter_name == "get_test_type"
assert resource_node.source_file == "/path/to/file.py"
assert resource_node.source_line == 42
assert resource_node.description == "Test docstring"
# Test serialization
data = resource_node.model_dump()
assert data["id"] == "resource_abc123"
assert data["label"] == "TestType"
assert data["type_name"] == "TestType"
assert data["node_type"] == "resource"
# Test deserialization
restored = WorkflowResourceNode.model_validate(data)
assert restored.id == "resource_abc123"
assert restored.label == "TestType"
assert restored.type_name == "TestType"
assert restored.node_type == "resource"
def test_graph_with_resources() -> None:
"""Test that workflow graph with resources is correct."""
wf = WorkflowWithResources()
graph = extract_workflow_structure(workflow=wf)
# Check resource nodes are in the nodes list
resource_nodes = [n for n in graph.nodes if isinstance(n, WorkflowResourceNode)]
assert len(resource_nodes) == 1
rn = resource_nodes[0]
assert rn.type_name == "DatabaseClient"
assert rn.getter_name == "get_database_client"
# Check edges with labels
resource_edges = [e for e in graph.edges if e.label is not None]
assert len(resource_edges) == 1
assert resource_edges[0].label == "db_client"
def test_edge_with_label() -> None:
"""Test that WorkflowGraphEdge with label works correctly."""
edge = WorkflowGraphEdge(source="resource_123", target="my_step", label="my_var")
assert edge.source == "resource_123"
assert edge.target == "my_step"
assert edge.label == "my_var"
def test_edge_without_label() -> None:
"""Test that WorkflowGraphEdge without label works correctly."""
edge = WorkflowGraphEdge(source="event_A", target="step_B")
assert edge.source == "event_A"
assert edge.target == "step_B"
assert edge.label is None
# --- Serialization/Deserialization tests for all node types ---
def test_step_node_serialization_roundtrip() -> None:
"""Test WorkflowStepNode serialization and deserialization."""
node = WorkflowStepNode(id="my_step", label="My Step")
data = node.model_dump()
assert data["id"] == "my_step"
assert data["label"] == "My Step"
assert data["node_type"] == "step"
restored = WorkflowStepNode.model_validate(data)
assert restored.id == "my_step"
assert restored.label == "My Step"
assert restored.node_type == "step"
def test_event_node_serialization_roundtrip() -> None:
"""Test WorkflowEventNode serialization and deserialization."""
node = WorkflowEventNode(
id="MyEvent",
label="My Event",
event_type="MyEvent",
event_types=["MyEvent", "ParentEvent"],
)
data = node.model_dump()
assert data["id"] == "MyEvent"
assert data["label"] == "My Event"
assert data["node_type"] == "event"
assert data["event_type"] == "MyEvent"
assert data["event_types"] == ["MyEvent", "ParentEvent"]
restored = WorkflowEventNode.model_validate(data)
assert restored.id == "MyEvent"
assert restored.label == "My Event"
assert restored.node_type == "event"
assert restored.event_type == "MyEvent"
assert restored.event_types == ["MyEvent", "ParentEvent"]
assert restored.is_subclass_of("ParentEvent")
assert not restored.is_subclass_of("UnrelatedEvent")
def test_external_node_serialization_roundtrip() -> None:
"""Test WorkflowExternalNode serialization and deserialization."""
node = WorkflowExternalNode(id="external_step", label="External Step")
data = node.model_dump()
assert data["id"] == "external_step"
assert data["label"] == "External Step"
assert data["node_type"] == "external"
restored = WorkflowExternalNode.model_validate(data)
assert restored.id == "external_step"
assert restored.label == "External Step"
assert restored.node_type == "external"
def test_resource_node_serialization_roundtrip() -> None:
"""Test WorkflowResourceNode serialization and deserialization."""
node = WorkflowResourceNode(
id="resource_abc123",
label="MyResourceType",
type_name="MyResourceType",
getter_name="get_my_resource",
source_file="/path/to/source.py",
source_line=100,
description="Resource docstring",
)
data = node.model_dump()
assert data["id"] == "resource_abc123"
assert data["label"] == "MyResourceType"
assert data["node_type"] == "resource"
assert data["type_name"] == "MyResourceType"
assert data["getter_name"] == "get_my_resource"
assert data["source_file"] == "/path/to/source.py"
assert data["source_line"] == 100
assert data["description"] == "Resource docstring"
restored = WorkflowResourceNode.model_validate(data)
assert restored.id == "resource_abc123"
assert restored.label == "MyResourceType"
assert restored.type_name == "MyResourceType"
assert restored.getter_name == "get_my_resource"
assert restored.node_type == "resource"
def test_graph_with_all_node_types_serialization() -> None:
"""Test full graph serialization/deserialization with all node types."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="StartEvent",
label="StartEvent",
event_type="StartEvent",
event_types=["StartEvent"],
),
WorkflowExternalNode(id="external", label="External"),
WorkflowResourceNode(
id="resource_123",
label="DB",
type_name="DatabaseClient",
getter_name="get_db",
),
],
edges=[
WorkflowGraphEdge(source="StartEvent", target="step1"),
WorkflowGraphEdge(source="step1", target="resource_123", label="db"),
],
)
# Serialize
data = graph.model_dump()
assert len(data["nodes"]) == 4
assert len(data["edges"]) == 2
# Check discriminator values are present
node_types = {n["node_type"] for n in data["nodes"]}
assert node_types == {"step", "event", "external", "resource"}
# Deserialize
restored = WorkflowGraph.model_validate(data)
assert len(restored.nodes) == 4
assert len(restored.edges) == 2
# Check correct types restored
step_nodes = [n for n in restored.nodes if isinstance(n, WorkflowStepNode)]
event_nodes = [n for n in restored.nodes if isinstance(n, WorkflowEventNode)]
external_nodes = [n for n in restored.nodes if isinstance(n, WorkflowExternalNode)]
resource_nodes = [n for n in restored.nodes if isinstance(n, WorkflowResourceNode)]
assert len(step_nodes) == 1
assert len(event_nodes) == 1
assert len(external_nodes) == 1
assert len(resource_nodes) == 1
# Verify event node has its method
assert event_nodes[0].is_subclass_of("StartEvent")
# Verify resource node has its fields
assert resource_nodes[0].type_name == "DatabaseClient"
assert resource_nodes[0].getter_name == "get_db"
def test_graph_deserialization_from_raw_json() -> None:
"""Test that graph can be deserialized from raw JSON dict."""
raw_data = {
"nodes": [
{"id": "step1", "label": "Step 1", "node_type": "step"},
{
"id": "MyEvent",
"label": "MyEvent",
"node_type": "event",
"event_type": "MyEvent",
"event_types": ["MyEvent"],
},
{"id": "external", "label": "External", "node_type": "external"},
{
"id": "resource_xyz",
"label": "Resource",
"node_type": "resource",
"type_name": "SomeType",
},
],
"edges": [{"source": "MyEvent", "target": "step1"}],
}
graph = WorkflowGraph.model_validate(raw_data)
assert len(graph.nodes) == 4
assert isinstance(graph.nodes[0], WorkflowStepNode)
assert isinstance(graph.nodes[1], WorkflowEventNode)
assert isinstance(graph.nodes[2], WorkflowExternalNode)
assert isinstance(graph.nodes[3], WorkflowResourceNode)
# --- filter_by_node_type tests ---
def test_filter_by_node_type_removes_nodes() -> None:
"""Test that filter_by_node_type removes specified node types."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="EventA",
label="EventA",
event_type="EventA",
event_types=["EventA"],
),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[
WorkflowGraphEdge(source="step1", target="EventA"),
WorkflowGraphEdge(source="EventA", target="step2"),
],
)
filtered = graph.filter_by_node_type("event")
# Event nodes should be removed
assert len(filtered.nodes) == 2
assert all(n.node_type == "step" for n in filtered.nodes)
node_ids = {n.id for n in filtered.nodes}
assert node_ids == {"step1", "step2"}
def test_filter_by_node_type_resolves_edges() -> None:
"""Test that edges through filtered nodes are resolved."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="EventA",
label="EventA",
event_type="EventA",
event_types=["EventA"],
),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[
WorkflowGraphEdge(source="step1", target="EventA"),
WorkflowGraphEdge(source="EventA", target="step2"),
],
)
filtered = graph.filter_by_node_type("event")
# Edge should be resolved: step1 -> step2
assert len(filtered.edges) == 1
assert filtered.edges[0].source == "step1"
assert filtered.edges[0].target == "step2"
def test_filter_by_node_type_chain_of_filtered_nodes() -> None:
"""Test filtering handles chains of filtered nodes."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="EventA",
label="First Filtered Node",
event_type="EventA",
event_types=["EventA"],
),
WorkflowEventNode(
id="EventB",
label="Second Filtered Node",
event_type="EventB",
event_types=["EventB"],
),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[
WorkflowGraphEdge(source="step1", target="EventA"),
WorkflowGraphEdge(source="EventA", target="EventB"),
WorkflowGraphEdge(source="EventB", target="step2"),
],
)
filtered = graph.filter_by_node_type("event")
# Chain resolved: step1 -> step2, with first filtered node's label
assert len(filtered.nodes) == 2
assert len(filtered.edges) == 1
assert filtered.edges[0].source == "step1"
assert filtered.edges[0].target == "step2"
assert filtered.edges[0].label == "First Filtered Node"
def test_filter_by_node_type_multiple_types() -> None:
"""Test filtering multiple node types at once."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="EventA",
label="EventA",
event_type="EventA",
event_types=["EventA"],
),
WorkflowResourceNode(id="resource1", label="Resource"),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[
WorkflowGraphEdge(source="step1", target="EventA"),
WorkflowGraphEdge(source="step1", target="resource1", label="db"),
WorkflowGraphEdge(source="EventA", target="step2"),
],
)
filtered = graph.filter_by_node_type("event", "resource")
# Only step nodes remain
assert len(filtered.nodes) == 2
assert all(n.node_type == "step" for n in filtered.nodes)
# step1 -> step2 edge remains (resolved through EventA)
assert len(filtered.edges) == 1
assert filtered.edges[0].source == "step1"
assert filtered.edges[0].target == "step2"
def test_filter_by_node_type_preserves_direct_edges() -> None:
"""Test that direct edges between remaining nodes are preserved."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowStepNode(id="step2", label="Step 2"),
WorkflowEventNode(
id="EventA",
label="EventA",
event_type="EventA",
event_types=["EventA"],
),
],
edges=[
WorkflowGraphEdge(source="step1", target="step2"), # Direct edge
WorkflowGraphEdge(source="step2", target="EventA"),
],
)
filtered = graph.filter_by_node_type("event")
# Direct edge should be preserved
assert len(filtered.edges) == 1
assert filtered.edges[0].source == "step1"
assert filtered.edges[0].target == "step2"
def test_filter_by_node_type_uses_filtered_node_label() -> None:
"""Test that the first filtered node's label becomes the new edge label."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="EventA",
label="My Event Label",
event_type="EventA",
event_types=["EventA"],
),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[
WorkflowGraphEdge(source="step1", target="EventA"),
WorkflowGraphEdge(source="EventA", target="step2"),
],
)
filtered = graph.filter_by_node_type("event")
# Label from filtered node should be on the new edge
assert len(filtered.edges) == 1
assert filtered.edges[0].label == "My Event Label"
def test_filter_by_node_type_preserves_direct_edge_labels() -> None:
"""Test that labels on direct edges are preserved."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowResourceNode(id="resource1", label="Resource"),
WorkflowEventNode(
id="EventA",
label="EventA",
event_type="EventA",
event_types=["EventA"],
),
],
edges=[
WorkflowGraphEdge(source="step1", target="resource1", label="db"),
WorkflowGraphEdge(source="step1", target="EventA"),
],
)
filtered = graph.filter_by_node_type("event")
# Resource edge label should be preserved (it's a direct edge)
resource_edge = next(e for e in filtered.edges if e.target == "resource1")
assert resource_edge.label == "db"
def test_filter_by_node_type_no_matching_types() -> None:
"""Test filtering with types that don't exist in graph."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[WorkflowGraphEdge(source="step1", target="step2")],
)
filtered = graph.filter_by_node_type("nonexistent")
# Graph should be unchanged
assert len(filtered.nodes) == 2
assert len(filtered.edges) == 1
def test_filter_by_node_type_preserves_description() -> None:
"""Test that the workflow description is preserved."""
graph = WorkflowGraph(
nodes=[WorkflowStepNode(id="step1", label="Step 1")],
edges=[],
description="My workflow description",
)
filtered = graph.filter_by_node_type("event")
assert filtered.description == "My workflow description"
def test_filter_by_node_type_deduplicates_edges() -> None:
"""Test that duplicate edges are not created."""
graph = WorkflowGraph(
nodes=[
WorkflowStepNode(id="step1", label="Step 1"),
WorkflowEventNode(
id="EventA",
label="EventA",
event_type="EventA",
event_types=["EventA"],
),
WorkflowEventNode(
id="EventB",
label="EventB",
event_type="EventB",
event_types=["EventB"],
),
WorkflowStepNode(id="step2", label="Step 2"),
],
edges=[
# Both events lead to step2 from step1
WorkflowGraphEdge(source="step1", target="EventA"),
WorkflowGraphEdge(source="step1", target="EventB"),
WorkflowGraphEdge(source="EventA", target="step2"),
WorkflowGraphEdge(source="EventB", target="step2"),
],
)
filtered = graph.filter_by_node_type("event")
# Should only have one edge: step1 -> step2 (deduplicated)
assert len(filtered.edges) == 1
assert filtered.edges[0].source == "step1"
assert filtered.edges[0].target == "step2"