Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 60 additions & 8 deletions core/src/main/java/com/google/adk/flows/llmflows/BaseLlmFlow.java
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
import com.google.common.collect.ImmutableList;
import com.google.common.collect.Iterables;
import com.google.genai.types.FunctionResponse;
import com.google.genai.types.GenerateContentConfig;
import io.opentelemetry.api.trace.Span;
import io.opentelemetry.api.trace.StatusCode;
import io.opentelemetry.context.Context;
Expand All @@ -56,7 +57,9 @@
import io.reactivex.rxjava3.disposables.Disposable;
import io.reactivex.rxjava3.observers.DisposableCompletableObserver;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import java.util.concurrent.atomic.AtomicBoolean;
Expand Down Expand Up @@ -124,15 +127,17 @@ private Flowable<Event> preprocess(
RequestProcessor getRequestProcessorFromTools(LlmAgent agent) {
return (context, request) -> {
ReadonlyContext readonlyContext = new ReadonlyContext(context);
// In-model built-in names, recorded as each tool's request processor runs. They are
// collected here because this is the point at which both the assembled config tools and
// the tools' own contributions are visible, so no separate pass over the toolsets is
// needed to find them.
Set<String> declarationlessConfigNames = new HashSet<>();
List<BiFunction<LlmRequest.Builder, ToolContext, Completable>> processors = new ArrayList<>();

for (Object toolOrToolset : agent.toolsUnion()) {
if (toolOrToolset instanceof BaseTool baseTool) {
processors.add(
(builder, ctx) -> {
Completable c = baseTool.processLlmRequest(builder, ctx);
return c == null ? Completable.complete() : c;
});
(builder, ctx) -> applyTool(baseTool, builder, ctx, declarationlessConfigNames));
} else if (toolOrToolset instanceof BaseToolset baseToolset) {
// First apply the toolset's own request processor, then unwrap all tools from the toolset
// and apply each individual tool's request processor sequentially.
Expand All @@ -143,10 +148,7 @@ RequestProcessor getRequestProcessorFromTools(LlmAgent agent) {
return toolsetProcessor
.andThen(baseToolset.getTools(readonlyContext))
.concatMapCompletable(
b -> {
Completable tc = b.processLlmRequest(builder, ctx);
return tc == null ? Completable.complete() : tc;
});
b -> applyTool(b, builder, ctx, declarationlessConfigNames));
});
} else {
throw new IllegalArgumentException(
Expand All @@ -159,12 +161,62 @@ RequestProcessor getRequestProcessorFromTools(LlmAgent agent) {
ToolContext toolContext = ToolContext.builder(context).build();
return Flowable.fromIterable(processors)
.concatMapCompletable(f -> f.apply(builder, toolContext))
.andThen(
Completable.fromAction(
() -> {
Map<String, BaseTool> declared = builder.build().tools();
for (String name : declarationlessConfigNames) {
if (declared.containsKey(name)) {
throw new IllegalArgumentException("Duplicate tool name: " + name);
}
}
}))
.andThen(
Single.fromCallable(
() -> RequestProcessingResult.create(builder.build(), ImmutableList.of())));
};
}

/**
Comment thread
kvmilos marked this conversation as resolved.
* Runs one tool's request processor, recording the tool's name when it adds an entry to the
* request's config tools without adding one to {@link LlmRequest#tools()}.
*
* <p>Those are the in-model built-ins such as {@code google_search}: they never reach {@link
* LlmRequest#tools()}, so a function tool of the same name sits beside them and the model chooses
* between two things answering to one name. A tool that adds no config entry, such as {@code
* ExampleTool}, is not one of them and is left alone.
*
* <p>The condition is the two counts changing differently rather than the absence of a
* declaration. {@code VertexAiRagRetrieval} on Vertex declares a function yet still adds to the
* config without reaching {@code tools()}, so testing for a missing declaration misses it; and a
* tool that only contributed the shared function-declarations entry would be counted by a
* config-count test alone.
*/
private static Completable applyTool(
BaseTool tool,
LlmRequest.Builder builder,
ToolContext toolContext,
Set<String> declarationlessConfigNames) {
int configToolsBefore = configToolCount(builder);
int requestToolsBefore = builder.build().tools().size();
Completable result = tool.processLlmRequest(builder, toolContext);
Completable nonNullResult = result == null ? Completable.complete() : result;
return nonNullResult.doOnComplete(
() -> {
// Grew the config's tools but contributed nothing to the request's own tool map: an
// in-model built-in. Both counts are compared after the processor has run.
if (configToolCount(builder) > configToolsBefore
&& builder.build().tools().size() == requestToolsBefore) {
declarationlessConfigNames.add(tool.name());
}
});
}

/** Number of tools currently set on the request's config. */
private static int configToolCount(LlmRequest.Builder builder) {
return builder.build().config().flatMap(GenerateContentConfig::tools).map(List::size).orElse(0);
}

/**
* Post-processes the LLM response after receiving it from the LLM. Executes all registered {@link
* ResponseProcessor} instances. Emits events for the model response and any subsequent function
Expand Down
180 changes: 180 additions & 0 deletions core/src/test/java/com/google/adk/flows/llmflows/BaseLlmFlowTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
import com.google.adk.testing.TestLlm;
import com.google.adk.tools.BaseTool;
import com.google.adk.tools.BaseToolset;
import com.google.adk.tools.GoogleSearchTool;
import com.google.adk.tools.ToolContext;
import com.google.common.collect.ImmutableList;
import com.google.common.collect.ImmutableMap;
Expand Down Expand Up @@ -983,6 +984,185 @@ public void close() {}
.containsExactly("toolset-instruction\n\ntool-instruction");
}

@Test
public void getRequestProcessorFromTools_rejectsDeclarationlessNameCollision_inModelFirst() {
assertDeclarationlessCollisionRejected(
ImmutableList.of(GoogleSearchTool.INSTANCE, googleSearchFunctionTool()));
}

@Test
public void getRequestProcessorFromTools_rejectsDeclarationlessNameCollision_functionToolFirst() {
assertDeclarationlessCollisionRejected(
ImmutableList.of(googleSearchFunctionTool(), GoogleSearchTool.INSTANCE));
}

/** A function tool answering to the same name as the built-in search tool. */
private static BaseTool googleSearchFunctionTool() {
return new BaseTool("google_search", "function search") {
@Override
public Optional<FunctionDeclaration> declaration() {
return Optional.of(FunctionDeclaration.builder().name("google_search").build());
}
};
}

@Test
public void getRequestProcessorFromTools_allowsDeclarationlessToolAddingNoConfigEntry() {
// A declaration-less tool that adds nothing to the request's config tools is not a built-in
// answer, so it does not collide with a function tool of the same name. ExampleTool is the
// real-world case; a plain declaration-less BaseTool stands in for it here.
BaseTool declarationless = new BaseTool("lookup", "adds no config entry") {};

BaseTool functionTool =
new BaseTool("lookup", "function lookup") {
@Override
public Optional<FunctionDeclaration> declaration() {
return Optional.of(FunctionDeclaration.builder().name("lookup").build());
}
};

LlmAgent agent =
createTestAgentBuilder(createTestLlm(LlmResponse.builder().build()))
.tools(ImmutableList.of(declarationless, functionTool))
.build();

InvocationContext invocationContext = createInvocationContext(agent);
BaseLlmFlow baseLlmFlow = createBaseLlmFlowWithoutProcessors();
RequestProcessor requestProcessor = baseLlmFlow.getRequestProcessorFromTools(agent);

LlmRequest processedRequest =
requestProcessor
.processRequest(invocationContext, LlmRequest.builder().build())
.map(RequestProcessingResult::updatedRequest)
.blockingGet();

// The function tool is the only one that reaches the request's tools.
assertThat(processedRequest.tools()).containsKey("lookup");
}

private void assertDeclarationlessCollisionRejected(List<BaseTool> tools) {
// GoogleSearchTool.INSTANCE makes no network calls, so it is safe to use directly here.
// The caller passes the ordered tool list so the order under test is visible at the call.
LlmAgent agent =
createTestAgentBuilder(createTestLlm(LlmResponse.builder().build())).tools(tools).build();

InvocationContext invocationContext = createInvocationContext(agent);
BaseLlmFlow baseLlmFlow = createBaseLlmFlowWithoutProcessors();
RequestProcessor requestProcessor = baseLlmFlow.getRequestProcessorFromTools(agent);

IllegalArgumentException thrown =
assertThrows(
IllegalArgumentException.class,
() ->
requestProcessor
.processRequest(invocationContext, LlmRequest.builder().build())
.blockingGet());
assertThat(thrown).hasMessageThat().isEqualTo("Duplicate tool name: google_search");
}

@Test
public void getRequestProcessorFromTools_allowsBuiltInBesideDifferentlyNamedFunctionTool() {
// Positive control for the two rejection tests above. Those tests would all still pass if the
// guard rejected EVERY built-in, because each of them expects an exception. This one fails in
// that case: the names differ, so nothing collides and the request must go through.
BaseTool inModel = GoogleSearchTool.INSTANCE;

BaseTool functionTool =
new BaseTool("lookup", "function lookup") {
@Override
public Optional<FunctionDeclaration> declaration() {
return Optional.of(FunctionDeclaration.builder().name("lookup").build());
}
};

LlmAgent agent =
createTestAgentBuilder(createTestLlm(LlmResponse.builder().build()))
.tools(ImmutableList.of(inModel, functionTool))
.build();

InvocationContext invocationContext = createInvocationContext(agent);
BaseLlmFlow baseLlmFlow = createBaseLlmFlowWithoutProcessors();
RequestProcessor requestProcessor = baseLlmFlow.getRequestProcessorFromTools(agent);

LlmRequest processedRequest =
requestProcessor
.processRequest(invocationContext, LlmRequest.builder().build())
.map(RequestProcessingResult::updatedRequest)
.blockingGet();

// The differently named function tool reaches the request's tools; the built-in does not.
assertThat(processedRequest.tools()).containsKey("lookup");
}

@Test
public void getRequestProcessorFromTools_rejectsCollisionFromToolsetServedBuiltIn() {
// Covers the toolset branch of applyTool. A toolset that serves a built-in must be counted as
// a built-in exactly like an agent-level one, so a same-named agent-level function tool still
// collides. Without this test a toolset branch that skips applyTool passes every other test.
BaseToolset toolset =
new BaseToolset() {
@Override
public Flowable<BaseTool> getTools(ReadonlyContext readonlyContext) {
return Flowable.just(GoogleSearchTool.INSTANCE);
}

@Override
public void close() {}
};

BaseTool agentLevelFunctionTool =
new BaseTool("google_search", "function search") {
@Override
public Optional<FunctionDeclaration> declaration() {
return Optional.of(FunctionDeclaration.builder().name("google_search").build());
}
};

LlmAgent agent =
createTestAgentBuilder(createTestLlm(LlmResponse.builder().build()))
.tools(toolset, agentLevelFunctionTool)
.build();

InvocationContext invocationContext = createInvocationContext(agent);
BaseLlmFlow baseLlmFlow = createBaseLlmFlowWithoutProcessors();
RequestProcessor requestProcessor = baseLlmFlow.getRequestProcessorFromTools(agent);

IllegalArgumentException thrown =
assertThrows(
IllegalArgumentException.class,
() ->
requestProcessor
.processRequest(invocationContext, LlmRequest.builder().build())
.blockingGet());
assertThat(thrown).hasMessageThat().isEqualTo("Duplicate tool name: google_search");
}

@Test
public void getRequestProcessorFromTools_allowsTwoDeclarationlessToolsSharingAName() {
// Both are declaration-less, so nothing is dispatched by name and no name is taken. Two
// default-named ExampleTools must keep working.
BaseTool first = new BaseTool("same_name", "first") {};
BaseTool second = new BaseTool("same_name", "second") {};

LlmAgent agent =
createTestAgentBuilder(createTestLlm(LlmResponse.builder().build()))
.tools(ImmutableList.of(first, second))
.build();

InvocationContext invocationContext = createInvocationContext(agent);
BaseLlmFlow baseLlmFlow = createBaseLlmFlowWithoutProcessors();
RequestProcessor requestProcessor = baseLlmFlow.getRequestProcessorFromTools(agent);

LlmRequest processedRequest =
requestProcessor
.processRequest(invocationContext, LlmRequest.builder().build())
.map(RequestProcessingResult::updatedRequest)
.blockingGet();

// Neither tool declares anything, so neither contributes an entry to the request's tools.
assertThat(processedRequest.tools()).isEmpty();
}

@Test
public void getRequestProcessorFromTools_throwsOnUnsupportedType() {
LlmAgent agent =
Expand Down
Loading