fix: Fix broken agent graphs

Fixes #973, #1170

PiperOrigin-RevId: 767921051
This commit is contained in:
Selcuk Gun
2025-06-05 22:46:39 -07:00
committed by Copybara-Service
parent 6ed635190c
commit 3b1f2ae9bf
+37 -35
View File
@@ -64,11 +64,11 @@ async def build_graph(
if isinstance(tool_or_agent, BaseAgent): if isinstance(tool_or_agent, BaseAgent):
# Added Workflow Agent checks for different agent types # Added Workflow Agent checks for different agent types
if isinstance(tool_or_agent, SequentialAgent): if isinstance(tool_or_agent, SequentialAgent):
return tool_or_agent.name + f' (Sequential Agent)' return tool_or_agent.name + ' (Sequential Agent)'
elif isinstance(tool_or_agent, LoopAgent): elif isinstance(tool_or_agent, LoopAgent):
return tool_or_agent.name + f' (Loop Agent)' return tool_or_agent.name + ' (Loop Agent)'
elif isinstance(tool_or_agent, ParallelAgent): elif isinstance(tool_or_agent, ParallelAgent):
return tool_or_agent.name + f' (Parallel Agent)' return tool_or_agent.name + ' (Parallel Agent)'
else: else:
return tool_or_agent.name return tool_or_agent.name
elif isinstance(tool_or_agent, BaseTool): elif isinstance(tool_or_agent, BaseTool):
@@ -144,49 +144,53 @@ async def build_graph(
) )
return False return False
def build_cluster(child: graphviz.Digraph, agent: BaseAgent, name: str): async def build_cluster(child: graphviz.Digraph, agent: BaseAgent, name: str):
if isinstance(agent, LoopAgent) and parent_agent: if isinstance(agent, LoopAgent):
# Draw the edge from the parent agent to the first sub-agent # Draw the edge from the parent agent to the first sub-agent
draw_edge(parent_agent.name, agent.sub_agents[0].name) if parent_agent:
draw_edge(parent_agent.name, agent.sub_agents[0].name)
length = len(agent.sub_agents) length = len(agent.sub_agents)
currLength = 0 curr_length = 0
# Draw the edges between the sub-agents # Draw the edges between the sub-agents
for sub_agent_int_sequential in agent.sub_agents: for sub_agent_int_sequential in agent.sub_agents:
build_graph(child, sub_agent_int_sequential, highlight_pairs) await build_graph(child, sub_agent_int_sequential, highlight_pairs)
# Draw the edge between the current sub-agent and the next one # Draw the edge between the current sub-agent and the next one
# If it's the last sub-agent, draw an edge to the first one to indicating a loop # If it's the last sub-agent, draw an edge to the first one to indicating a loop
draw_edge( draw_edge(
agent.sub_agents[currLength].name, agent.sub_agents[curr_length].name,
agent.sub_agents[ agent.sub_agents[
0 if currLength == length - 1 else currLength + 1 0 if curr_length == length - 1 else curr_length + 1
].name, ].name,
) )
currLength += 1 curr_length += 1
elif isinstance(agent, SequentialAgent) and parent_agent: elif isinstance(agent, SequentialAgent):
# Draw the edge from the parent agent to the first sub-agent # Draw the edge from the parent agent to the first sub-agent
draw_edge(parent_agent.name, agent.sub_agents[0].name) if parent_agent:
draw_edge(parent_agent.name, agent.sub_agents[0].name)
length = len(agent.sub_agents) length = len(agent.sub_agents)
currLength = 0 curr_length = 0
# Draw the edges between the sub-agents # Draw the edges between the sub-agents
for sub_agent_int_sequential in agent.sub_agents: for sub_agent_int_sequential in agent.sub_agents:
build_graph(child, sub_agent_int_sequential, highlight_pairs) await build_graph(child, sub_agent_int_sequential, highlight_pairs)
# Draw the edge between the current sub-agent and the next one # Draw the edge between the current sub-agent and the next one
# If it's the last sub-agent, don't draw an edge to avoid a loop # If it's the last sub-agent, don't draw an edge to avoid a loop
draw_edge( if curr_length != length - 1:
agent.sub_agents[currLength].name, draw_edge(
agent.sub_agents[currLength + 1].name, agent.sub_agents[curr_length].name,
) if currLength != length - 1 else None agent.sub_agents[curr_length + 1].name,
currLength += 1 )
curr_length += 1
elif isinstance(agent, ParallelAgent) and parent_agent: elif isinstance(agent, ParallelAgent):
# Draw the edge from the parent agent to every sub-agent # Draw the edge from the parent agent to every sub-agent
for sub_agent in agent.sub_agents: for sub_agent in agent.sub_agents:
build_graph(child, sub_agent, highlight_pairs) await build_graph(child, sub_agent, highlight_pairs)
draw_edge(parent_agent.name, sub_agent.name) if parent_agent:
draw_edge(parent_agent.name, sub_agent.name)
else: else:
for sub_agent in agent.sub_agents: for sub_agent in agent.sub_agents:
build_graph(child, sub_agent, highlight_pairs) await build_graph(child, sub_agent, highlight_pairs)
draw_edge(agent.name, sub_agent.name) draw_edge(agent.name, sub_agent.name)
child.attr( child.attr(
@@ -196,21 +200,20 @@ async def build_graph(
fontcolor=light_gray, fontcolor=light_gray,
) )
def draw_node(tool_or_agent: Union[BaseAgent, BaseTool]): async def draw_node(tool_or_agent: Union[BaseAgent, BaseTool]):
name = get_node_name(tool_or_agent) name = get_node_name(tool_or_agent)
shape = get_node_shape(tool_or_agent) shape = get_node_shape(tool_or_agent)
caption = get_node_caption(tool_or_agent) caption = get_node_caption(tool_or_agent)
asCluster = should_build_agent_cluster(tool_or_agent) as_cluster = should_build_agent_cluster(tool_or_agent)
child = None
if highlight_pairs: if highlight_pairs:
for highlight_tuple in highlight_pairs: for highlight_tuple in highlight_pairs:
if name in highlight_tuple: if name in highlight_tuple:
# if in highlight, draw highlight node # if in highlight, draw highlight node
if asCluster: if as_cluster:
cluster = graphviz.Digraph( cluster = graphviz.Digraph(
name='cluster_' + name name='cluster_' + name
) # adding "cluster_" to the name makes the graph render as a cluster subgraph ) # adding "cluster_" to the name makes the graph render as a cluster subgraph
build_cluster(cluster, agent, name) await build_cluster(cluster, agent, name)
graph.subgraph(cluster) graph.subgraph(cluster)
else: else:
graph.node( graph.node(
@@ -224,12 +227,12 @@ async def build_graph(
) )
return return
# if not in highlight, draw non-highlight node # if not in highlight, draw non-highlight node
if asCluster: if as_cluster:
cluster = graphviz.Digraph( cluster = graphviz.Digraph(
name='cluster_' + name name='cluster_' + name
) # adding "cluster_" to the name makes the graph render as a cluster subgraph ) # adding "cluster_" to the name makes the graph render as a cluster subgraph
build_cluster(cluster, agent, name) await build_cluster(cluster, agent, name)
graph.subgraph(cluster) graph.subgraph(cluster)
else: else:
@@ -264,10 +267,9 @@ async def build_graph(
else: else:
graph.edge(from_name, to_name, arrowhead='none', color=light_gray) graph.edge(from_name, to_name, arrowhead='none', color=light_gray)
draw_node(agent) await draw_node(agent)
for sub_agent in agent.sub_agents: for sub_agent in agent.sub_agents:
await build_graph(graph, sub_agent, highlight_pairs, agent)
build_graph(graph, sub_agent, highlight_pairs, agent)
if not should_build_agent_cluster( if not should_build_agent_cluster(
sub_agent sub_agent
) and not should_build_agent_cluster( ) and not should_build_agent_cluster(
@@ -276,7 +278,7 @@ async def build_graph(
draw_edge(agent.name, sub_agent.name) draw_edge(agent.name, sub_agent.name)
if isinstance(agent, LlmAgent): if isinstance(agent, LlmAgent):
for tool in await agent.canonical_tools(): for tool in await agent.canonical_tools():
draw_node(tool) await draw_node(tool)
draw_edge(agent.name, get_node_name(tool)) draw_edge(agent.name, get_node_name(tool))