Skip to content

Commit f866804

Browse files
kvmiloscopybara-github
authored andcommitted
fix: run approved tool confirmations for agents under a ParallelAgent
PiperOrigin-RevId: 993677550
1 parent 52e51c6 commit f866804

4 files changed

Lines changed: 389 additions & 19 deletions

File tree

‎core/src/main/java/com/google/adk/runner/Runner.java‎

Lines changed: 76 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,6 @@
2727
import com.google.adk.agents.InvocationContext;
2828
import com.google.adk.agents.LiveRequestQueue;
2929
import com.google.adk.agents.LlmAgent;
30-
import com.google.adk.agents.ParallelAgent;
3130
import com.google.adk.agents.Role;
3231
import com.google.adk.agents.RunConfig;
3332
import com.google.adk.agents.SequentialAgent;
@@ -416,13 +415,14 @@ private Single<Event> createUserMessageEvent(
416415
InvocationContext invocationContext,
417416
boolean saveInputBlobsAsArtifacts,
418417
@Nullable Map<String, Object> stateDelta) {
418+
// As on the resumable path, a function response takes the branch of the call it answers.
419419
return appendNewMessageToSession(
420420
session,
421421
newMessage,
422422
invocationContext,
423423
saveInputBlobsAsArtifacts,
424424
stateDelta,
425-
/* branch= */ null);
425+
matchingFunctionCallEvent(session, newMessage).flatMap(Event::branch).orElse(null));
426426
}
427427

428428
private Single<Event> appendNewMessageToSession(
@@ -768,12 +768,14 @@ private Flowable<Event> runAgentForUserEvent(
768768
*/
769769
private Flowable<Event> runAgentWithUpdatedSession(
770770
InvocationContext initialContext, Session updatedSession, Event event, BaseAgent rootAgent) {
771+
BaseAgent agentToRun = this.findAgentToRun(updatedSession, rootAgent);
771772
// Create context with updated session for beforeRunCallback
772773
InvocationContext contextWithUpdatedSession =
773774
initialContext.toBuilder()
774775
.session(updatedSession)
775-
.agent(this.findAgentToRun(updatedSession, rootAgent))
776+
.agent(agentToRun)
776777
.userContent(event.content().orElseGet(Content::fromParts))
778+
.branch(routedAgentParentBranch(updatedSession, agentToRun, rootAgent))
777779
.build();
778780

779781
// If beforeRunCallback returns content, emit it and skip agent.
@@ -1055,26 +1057,79 @@ private Flowable<Event> runResumedAgent(
10551057
}
10561058

10571059
/**
1058-
* Branch to seed a resumed context with so {@code resumeAgent} runs under the same branch it
1059-
* originally did. Returns the parent branch (the resolved agent's most recent event branch minus
1060-
* its own trailing name segment, which {@link BaseAgent#runAsync} re-appends), or {@code null}
1061-
* for the root branch. Non-null only for an agent nested under a {@link ParallelAgent}.
1060+
* Branch to seed a new invocation with so a routed sub-agent runs where it ran before, or {@code
1061+
* null} when {@code agentToRun} is the root. A function response restores the branch of the call
1062+
* it answers, since a Java agent's branch depends on the transfer path that reached it; other
1063+
* routing restores the agent's latest branch, as Python does.
1064+
*/
1065+
private static @Nullable String routedAgentParentBranch(
1066+
Session session, BaseAgent agentToRun, BaseAgent rootAgent) {
1067+
if (agentToRun.equals(rootAgent)) {
1068+
return null;
1069+
}
1070+
Optional<Event> answeredCall =
1071+
Functions.findMatchingFunctionCallEvent(session.immutableEvents());
1072+
if (answeredCall.isPresent()) {
1073+
String ownSuffix = sequentialBranchSuffix(agentToRun, answeredCall.get().author());
1074+
if (ownSuffix != null) {
1075+
return parentBranch(answeredCall.get().branch().orElse(null), ownSuffix);
1076+
}
1077+
}
1078+
return resumeParentBranch(session, /* invocationId= */ null, agentToRun);
1079+
}
1080+
1081+
/**
1082+
* Branch to seed a context with so {@code resumeAgent} runs under the branch it last ran on: the
1083+
* parent branch of the newest event by {@code resumeAgent}, or by an agent it reaches through
1084+
* SequentialAgents, with a non-empty branch. A non-null {@code invocationId} limits the search to
1085+
* that invocation; {@code null} is returned for the root branch.
10621086
*/
10631087
private static @Nullable String resumeParentBranch(
1064-
Session session, String invocationId, BaseAgent resumeAgent) {
1088+
Session session, @Nullable String invocationId, BaseAgent resumeAgent) {
10651089
ImmutableList<Event> events = session.immutableEvents();
10661090
for (int i = events.size() - 1; i >= 0; i--) {
10671091
Event event = events.get(i);
1068-
if (invocationId.equals(event.invocationId())
1069-
&& resumeAgent.name().equals(event.author())
1070-
&& event.branch().isPresent()) {
1071-
String branch = event.branch().get();
1072-
String ownSegment = "." + resumeAgent.name();
1073-
if (branch.endsWith(ownSegment)) {
1074-
String parent = branch.substring(0, branch.length() - ownSegment.length());
1075-
return parent.isEmpty() ? null : parent;
1092+
if ((invocationId == null || invocationId.equals(event.invocationId()))
1093+
&& event.branch().filter(branch -> !branch.isEmpty()).isPresent()) {
1094+
String ownSuffix = sequentialBranchSuffix(resumeAgent, event.author());
1095+
if (ownSuffix != null) {
1096+
return parentBranch(event.branch().get(), ownSuffix);
1097+
}
1098+
}
1099+
}
1100+
return null;
1101+
}
1102+
1103+
/**
1104+
* Returns {@code branch} minus the trailing {@code ownSuffix} that {@link BaseAgent#runAsync}
1105+
* re-appends, or {@code null} for the root branch.
1106+
*/
1107+
private static @Nullable String parentBranch(@Nullable String branch, String ownSuffix) {
1108+
if (branch == null || branch.isEmpty() || branch.equals(ownSuffix)) {
1109+
return null;
1110+
}
1111+
if (branch.endsWith("." + ownSuffix)) {
1112+
String parent = branch.substring(0, branch.length() - ownSuffix.length() - 1);
1113+
return parent.isEmpty() ? null : parent;
1114+
}
1115+
return branch;
1116+
}
1117+
1118+
/**
1119+
* Returns the dot-joined names from {@code agent} down to the agent named {@code author} when
1120+
* that agent is {@code agent} or below it through SequentialAgents, else {@code null}. The legacy
1121+
* flow resumes such an ancestor, which records no events of its own.
1122+
*/
1123+
private static @Nullable String sequentialBranchSuffix(BaseAgent agent, @Nullable String author) {
1124+
if (agent.name().equals(author)) {
1125+
return agent.name();
1126+
}
1127+
if (agent instanceof SequentialAgent) {
1128+
for (BaseAgent subAgent : agent.subAgents()) {
1129+
String suffix = sequentialBranchSuffix(subAgent, author);
1130+
if (suffix != null) {
1131+
return agent.name() + "." + suffix;
10761132
}
1077-
return branch.equals(resumeAgent.name()) ? null : branch;
10781133
}
10791134
}
10801135
return null;
@@ -1241,11 +1296,13 @@ private InvocationContext newInvocationContextForLive(
12411296
runConfigBuilder.inputAudioTranscription(AudioTranscriptionConfig.builder().build());
12421297
}
12431298
}
1299+
BaseAgent agentToRun = findAgentToRun(session, this.agent);
12441300
InvocationContext.Builder builder =
1245-
newInvocationContextBuilder(session, findAgentToRun(session, this.agent))
1301+
newInvocationContextBuilder(session, agentToRun)
12461302
.runConfig(runConfigBuilder.build())
12471303
.userContent(Content.fromParts())
1248-
.liveRequestQueue(liveRequestQueue);
1304+
.liveRequestQueue(liveRequestQueue)
1305+
.branch(routedAgentParentBranch(session, agentToRun, this.agent));
12491306

12501307
return builder.build();
12511308
}

‎core/src/test/java/com/google/adk/runner/RunnerLegacyResumabilityTest.java‎

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,10 +17,14 @@
1717
package com.google.adk.runner;
1818

1919
import static com.google.adk.testing.ResumabilityTestUtils.answerCall;
20+
import static com.google.adk.testing.ResumabilityTestUtils.approveConfirmation;
21+
import static com.google.adk.testing.ResumabilityTestUtils.confirmingEchoFunctionTool;
2022
import static com.google.adk.testing.ResumabilityTestUtils.newSession;
2123
import static com.google.adk.testing.ResumabilityTestUtils.pendingFunctionTool;
2224
import static com.google.adk.testing.ResumabilityTestUtils.runTurn;
25+
import static com.google.adk.testing.ResumabilityTestUtils.runTurnAskingConfirmation;
2326
import static com.google.adk.testing.ResumabilityTestUtils.shimRunner;
27+
import static com.google.adk.testing.ResumabilityTestUtils.textAgent;
2428
import static com.google.adk.testing.TestUtils.createFunctionCallLlmResponse;
2529
import static com.google.adk.testing.TestUtils.createLlmResponse;
2630
import static com.google.adk.testing.TestUtils.createTestAgentBuilder;
@@ -475,6 +479,86 @@ public void runAsync_plainTextWithShim_withStateDelta_mergesStateIntoSession() {
475479
assertThat(finalSession.state()).containsAtLeastEntriesIn(stateDelta);
476480
}
477481

482+
// The shim resumes the SequentialAgent, which has no events of its own to restore a branch from.
483+
@Test
484+
public void
485+
runAsync_withToolConfirmation_inSequentialAgentUnderParallelAgent_callsTool_legacyShim() {
486+
LlmAgent childAgent =
487+
createTestAgentBuilder(
488+
createTestLlm(
489+
createFunctionCallLlmResponse(
490+
"tool_call_id", "echoTool", ImmutableMap.of("message", "hello")),
491+
createTextLlmResponse("Response after user confirmed.")))
492+
.name("child_agent")
493+
.tools(confirmingEchoFunctionTool())
494+
.build();
495+
SequentialAgent sequentialAgent =
496+
SequentialAgent.builder()
497+
.name("sequential_agent")
498+
.subAgents(ImmutableList.of(childAgent))
499+
.build();
500+
ParallelAgent rootAgent =
501+
ParallelAgent.builder()
502+
.name("parallel_agent")
503+
.subAgents(
504+
ImmutableList.of(sequentialAgent, textAgent("sibling_agent", "Sibling done.")))
505+
.build();
506+
Runner runner = shimRunner(rootAgent);
507+
Session session = newSession(runner);
508+
FunctionCall askUserConfirmationFunctionCall =
509+
runTurnAskingConfirmation(runner, session, "from user");
510+
511+
ImmutableList<Event> eventsAfterConfirmation =
512+
approveConfirmation(runner, session, askUserConfirmationFunctionCall);
513+
514+
assertThat(simplifyEvents(eventsAfterConfirmation))
515+
.containsExactly(
516+
"child_agent: FunctionResponse(name=echoTool, response={message=hello})",
517+
"child_agent: Response after user confirmed.")
518+
.inOrder();
519+
assertThat(eventsAfterConfirmation.stream().map(event -> event.branch().orElse(null)))
520+
.containsExactly(
521+
"parallel_agent.sequential_agent.child_agent",
522+
"parallel_agent.sequential_agent.child_agent");
523+
}
524+
525+
// Answering a long-running call resumes the sequence; its later sub-agent must still run.
526+
@Test
527+
public void
528+
runAsync_withLongRunningCall_inSequentialAgentUnderParallelAgent_runsNextAgent_legacyShim() {
529+
LlmAgent childAgent =
530+
createTestAgentBuilder(
531+
createTestLlm(
532+
createFunctionCallLlmResponse(
533+
"lro_call_id", "pendingTool", ImmutableMap.of("message", "draft")),
534+
createTextLlmResponse("child resumed")))
535+
.name("child_agent")
536+
.tools(pendingFunctionTool())
537+
.build();
538+
SequentialAgent sequentialAgent =
539+
SequentialAgent.builder()
540+
.name("sequential_agent")
541+
.subAgents(ImmutableList.of(childAgent, textAgent("next_agent", "next done")))
542+
.build();
543+
ParallelAgent rootAgent =
544+
ParallelAgent.builder()
545+
.name("parallel_agent")
546+
.subAgents(
547+
ImmutableList.of(sequentialAgent, textAgent("sibling_agent", "Sibling done.")))
548+
.build();
549+
Runner runner = shimRunner(rootAgent);
550+
Session session = newSession(runner);
551+
var unused = runTurn(runner, session, "from user");
552+
553+
ImmutableList<Event> eventsAfterResume =
554+
answerCall(
555+
runner, session, "lro_call_id", "pendingTool", ImmutableMap.of("result", "done"));
556+
557+
assertThat(simplifyEvents(eventsAfterResume))
558+
.containsExactly("child_agent: child resumed", "next_agent: next done")
559+
.inOrder();
560+
}
561+
478562
// ===== CL1-parity: every CL1 resumable(true) test, re-run under the text-only shim =====
479563

480564
@Test

0 commit comments

Comments
 (0)