diff --git a/src/google/adk/cli/agent_graph.py b/src/google/adk/cli/agent_graph.py index a0b8a467..249e5bfb 100644 --- a/src/google/adk/cli/agent_graph.py +++ b/src/google/adk/cli/agent_graph.py @@ -64,11 +64,11 @@ async def build_graph( if isinstance(tool_or_agent, BaseAgent): # Added Workflow Agent checks for different agent types 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): - return tool_or_agent.name + f' (Loop Agent)' + return tool_or_agent.name + ' (Loop Agent)' elif isinstance(tool_or_agent, ParallelAgent): - return tool_or_agent.name + f' (Parallel Agent)' + return tool_or_agent.name + ' (Parallel Agent)' else: return tool_or_agent.name elif isinstance(tool_or_agent, BaseTool): @@ -144,49 +144,53 @@ async def build_graph( ) return False - def build_cluster(child: graphviz.Digraph, agent: BaseAgent, name: str): - if isinstance(agent, LoopAgent) and parent_agent: + async def build_cluster(child: graphviz.Digraph, agent: BaseAgent, name: str): + if isinstance(agent, LoopAgent): # 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) - currLength = 0 + curr_length = 0 # Draw the edges between the 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 # If it's the last sub-agent, draw an edge to the first one to indicating a loop draw_edge( - agent.sub_agents[currLength].name, + agent.sub_agents[curr_length].name, agent.sub_agents[ - 0 if currLength == length - 1 else currLength + 1 + 0 if curr_length == length - 1 else curr_length + 1 ].name, ) - currLength += 1 - elif isinstance(agent, SequentialAgent) and parent_agent: + curr_length += 1 + elif isinstance(agent, SequentialAgent): # 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) - currLength = 0 + curr_length = 0 # Draw the edges between the 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 # If it's the last sub-agent, don't draw an edge to avoid a loop - draw_edge( - agent.sub_agents[currLength].name, - agent.sub_agents[currLength + 1].name, - ) if currLength != length - 1 else None - currLength += 1 + if curr_length != length - 1: + draw_edge( + agent.sub_agents[curr_length].name, + agent.sub_agents[curr_length + 1].name, + ) + 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 for sub_agent in agent.sub_agents: - build_graph(child, sub_agent, highlight_pairs) - draw_edge(parent_agent.name, sub_agent.name) + await build_graph(child, sub_agent, highlight_pairs) + if parent_agent: + draw_edge(parent_agent.name, sub_agent.name) else: 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) child.attr( @@ -196,21 +200,20 @@ async def build_graph( 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) shape = get_node_shape(tool_or_agent) caption = get_node_caption(tool_or_agent) - asCluster = should_build_agent_cluster(tool_or_agent) - child = None + as_cluster = should_build_agent_cluster(tool_or_agent) if highlight_pairs: for highlight_tuple in highlight_pairs: if name in highlight_tuple: # if in highlight, draw highlight node - if asCluster: + if as_cluster: cluster = graphviz.Digraph( name='cluster_' + name ) # 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) else: graph.node( @@ -224,12 +227,12 @@ async def build_graph( ) return # if not in highlight, draw non-highlight node - if asCluster: + if as_cluster: cluster = graphviz.Digraph( name='cluster_' + name ) # 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) else: @@ -264,10 +267,9 @@ async def build_graph( else: 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: - - build_graph(graph, sub_agent, highlight_pairs, agent) + await build_graph(graph, sub_agent, highlight_pairs, agent) if not should_build_agent_cluster( sub_agent ) and not should_build_agent_cluster( @@ -276,7 +278,7 @@ async def build_graph( draw_edge(agent.name, sub_agent.name) if isinstance(agent, LlmAgent): for tool in await agent.canonical_tools(): - draw_node(tool) + await draw_node(tool) draw_edge(agent.name, get_node_name(tool))