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