From 2493cd135ab82beaf1aeee06eeed1475bb727b87 Mon Sep 17 00:00:00 2001 From: Pierre Gentile Date: Wed, 19 Aug 2026 14:30:59 +0200 Subject: [PATCH] feat: configure similarity top-K in the VertexAiRagRetrieval tool --- .../tools/retrieval/VertexAiRagRetrieval.java | 37 ++++++- .../retrieval/VertexAiRagRetrievalTest.java | 96 +++++++++++++++++-- 2 files changed, 124 insertions(+), 9 deletions(-) diff --git a/core/src/main/java/com/google/adk/tools/retrieval/VertexAiRagRetrieval.java b/core/src/main/java/com/google/adk/tools/retrieval/VertexAiRagRetrieval.java index a2720aae5..5dc579d4d 100644 --- a/core/src/main/java/com/google/adk/tools/retrieval/VertexAiRagRetrieval.java +++ b/core/src/main/java/com/google/adk/tools/retrieval/VertexAiRagRetrieval.java @@ -23,6 +23,7 @@ import com.google.adk.utils.ModelNameUtils; import com.google.cloud.aiplatform.v1.RagContexts; import com.google.cloud.aiplatform.v1.RagQuery; +import com.google.cloud.aiplatform.v1.RagRetrievalConfig; import com.google.cloud.aiplatform.v1.RetrieveContextsRequest; import com.google.cloud.aiplatform.v1.RetrieveContextsRequest.VertexRagStore.RagResource; import com.google.cloud.aiplatform.v1.RetrieveContextsResponse; @@ -47,7 +48,8 @@ * A retrieval tool that fetches context from Vertex AI RAG. * *

This tool allows to retrieve relevant information based on a query using Vertex AI RAG - * service. It supports configuration of rag resources and a vector distance threshold. + * service. It supports configuration of rag resources, a vector distance threshold, and similarity + * top-k. */ public class VertexAiRagRetrieval extends BaseRetrievalTool { private static final Logger logger = LoggerFactory.getLogger(VertexAiRagRetrieval.class); @@ -55,6 +57,7 @@ public class VertexAiRagRetrieval extends BaseRetrievalTool { private final String parent; private final List ragResources; private final Double vectorDistanceThreshold; + private final Integer similarityTopK; private final VertexRagStore vertexRagStore; private final RetrieveContextsRequest.VertexRagStore apiVertexRagStore; @@ -65,11 +68,30 @@ public VertexAiRagRetrieval( String parent, @Nullable List ragResources, @Nullable Double vectorDistanceThreshold) { + this( + name, + description, + vertexRagServiceClient, + parent, + ragResources, + vectorDistanceThreshold, + null); + } + + public VertexAiRagRetrieval( + String name, + String description, + VertexRagServiceClient vertexRagServiceClient, + String parent, + @Nullable List ragResources, + @Nullable Double vectorDistanceThreshold, + @Nullable Integer similarityTopK) { super(name, description); this.vertexRagServiceClient = vertexRagServiceClient; this.parent = parent; this.ragResources = ragResources; this.vectorDistanceThreshold = vectorDistanceThreshold; + this.similarityTopK = similarityTopK; // For Gemini 2 VertexRagStore.Builder vertexRagStoreBuilder = VertexRagStore.builder(); @@ -86,6 +108,9 @@ public VertexAiRagRetrieval( if (this.vectorDistanceThreshold != null) { vertexRagStoreBuilder.vectorDistanceThreshold(this.vectorDistanceThreshold); } + if (this.similarityTopK != null) { + vertexRagStoreBuilder.similarityTopK(this.similarityTopK); + } this.vertexRagStore = vertexRagStoreBuilder.build(); // For runAsync @@ -135,10 +160,18 @@ public Single> runAsync(Map args, ToolContex return Single.fromCallable( () -> { logger.info("Retrieving context for query: {}", query); + + RagQuery.Builder queryBuilder = RagQuery.newBuilder().setText(query); + + if (similarityTopK != null) { + queryBuilder.setRagRetrievalConfig( + RagRetrievalConfig.newBuilder().setTopK(similarityTopK).build()); + } + RetrieveContextsRequest retrieveContextsRequest = RetrieveContextsRequest.newBuilder() .setParent(this.parent) - .setQuery(RagQuery.newBuilder().setText(query)) + .setQuery(queryBuilder.build()) .setVertexRagStore(this.apiVertexRagStore) .build(); logger.info("Request to VertexRagService: {}", retrieveContextsRequest); diff --git a/core/src/test/java/com/google/adk/tools/retrieval/VertexAiRagRetrievalTest.java b/core/src/test/java/com/google/adk/tools/retrieval/VertexAiRagRetrievalTest.java index 1b8cbf66a..86c807202 100644 --- a/core/src/test/java/com/google/adk/tools/retrieval/VertexAiRagRetrievalTest.java +++ b/core/src/test/java/com/google/adk/tools/retrieval/VertexAiRagRetrievalTest.java @@ -29,6 +29,7 @@ import com.google.adk.tools.ToolContext; import com.google.cloud.aiplatform.v1.RagContexts; import com.google.cloud.aiplatform.v1.RagQuery; +import com.google.cloud.aiplatform.v1.RagRetrievalConfig; import com.google.cloud.aiplatform.v1.RetrieveContextsRequest; import com.google.cloud.aiplatform.v1.RetrieveContextsRequest.VertexRagStore.RagResource; import com.google.cloud.aiplatform.v1.RetrieveContextsResponse; @@ -72,6 +73,7 @@ public void runAsync_withResults_returnsContexts() throws Exception { ImmutableList ragResources = ImmutableList.of(RagResource.newBuilder().setRagCorpus("corpus1").build()); Double vectorDistanceThreshold = 0.5; + Integer similarityTopK = 10; VertexAiRagRetrieval tool = new VertexAiRagRetrieval( "testTool", @@ -79,13 +81,18 @@ public void runAsync_withResults_returnsContexts() throws Exception { vertexRagServiceClient, "projects/test-project/locations/us-central1", ragResources, - vectorDistanceThreshold); + vectorDistanceThreshold, + similarityTopK); String query = "test query"; ToolContext toolContext = buildToolContext(); RetrieveContextsRequest expectedRequest = RetrieveContextsRequest.newBuilder() .setParent("projects/test-project/locations/us-central1") - .setQuery(RagQuery.newBuilder().setText(query)) + .setQuery( + RagQuery.newBuilder() + .setText(query) + .setRagRetrievalConfig( + RagRetrievalConfig.newBuilder().setTopK(similarityTopK).build())) .setVertexRagStore( com.google.cloud.aiplatform.v1.RetrieveContextsRequest.VertexRagStore.newBuilder() .addAllRagResources(ragResources) @@ -112,6 +119,7 @@ public void runAsync_noResults_returnsNoResultFoundMessage() throws Exception { ImmutableList ragResources = ImmutableList.of(RagResource.newBuilder().setRagCorpus("corpus1").build()); Double vectorDistanceThreshold = 0.5; + Integer similarityTopK = 10; VertexAiRagRetrieval tool = new VertexAiRagRetrieval( "testTool", @@ -119,13 +127,18 @@ public void runAsync_noResults_returnsNoResultFoundMessage() throws Exception { vertexRagServiceClient, "projects/test-project/locations/us-central1", ragResources, - vectorDistanceThreshold); + vectorDistanceThreshold, + similarityTopK); String query = "test query"; ToolContext toolContext = buildToolContext(); RetrieveContextsRequest expectedRequest = RetrieveContextsRequest.newBuilder() .setParent("projects/test-project/locations/us-central1") - .setQuery(RagQuery.newBuilder().setText(query)) + .setQuery( + RagQuery.newBuilder() + .setText(query) + .setRagRetrievalConfig( + RagRetrievalConfig.newBuilder().setTopK(similarityTopK).build())) .setVertexRagStore( com.google.cloud.aiplatform.v1.RetrieveContextsRequest.VertexRagStore.newBuilder() .addAllRagResources(ragResources) @@ -154,6 +167,7 @@ public void processLlmRequest_gemini2Model_addVertexRagStoreToConfig() { ImmutableList ragResources = ImmutableList.of(RagResource.newBuilder().setRagCorpus("corpus1").build()); Double vectorDistanceThreshold = 0.5; + Integer similarityTopK = 10; VertexAiRagRetrieval tool = new VertexAiRagRetrieval( "testTool", @@ -161,7 +175,8 @@ public void processLlmRequest_gemini2Model_addVertexRagStoreToConfig() { vertexRagServiceClient, "projects/test-project/locations/us-central1", ragResources, - vectorDistanceThreshold); + vectorDistanceThreshold, + similarityTopK); LlmRequest.Builder llmRequestBuilder = LlmRequest.builder().model("gemini-2-pro"); ToolContext toolContext = buildToolContext(); @@ -181,7 +196,8 @@ public void processLlmRequest_gemini2Model_addVertexRagStoreToConfig() { VertexRagStoreRagResource.builder() .ragCorpus("corpus1") .build())) - .vectorDistanceThreshold(0.5) + .vectorDistanceThreshold(vectorDistanceThreshold) + .similarityTopK(similarityTopK) .build()) .build()) .build()); @@ -212,7 +228,9 @@ public void processLlmRequest_gemini2Model_addVertexRagStoreToConfig() { } @Test - public void processLlmRequest_otherModel_doNotAddVertexRagStoreToConfig() { + public void processLlmRequest_gemini2Model_withoutSimilarityTopK_addVertexRagStoreToConfig() { + // This test's behavior depends on the GOOGLE_GENAI_USE_VERTEXAI environment variable + boolean useVertexAi = Boolean.parseBoolean(System.getenv("GOOGLE_GENAI_USE_VERTEXAI")); ImmutableList ragResources = ImmutableList.of(RagResource.newBuilder().setRagCorpus("corpus1").build()); Double vectorDistanceThreshold = 0.5; @@ -224,6 +242,70 @@ public void processLlmRequest_otherModel_doNotAddVertexRagStoreToConfig() { "projects/test-project/locations/us-central1", ragResources, vectorDistanceThreshold); + LlmRequest.Builder llmRequestBuilder = LlmRequest.builder().model("gemini-2-pro"); + ToolContext toolContext = buildToolContext(); + + tool.processLlmRequest(llmRequestBuilder, toolContext).blockingAwait(); + + if (useVertexAi) { + // Assert that VertexRagStore is added to the config without similarityTopK + assertThat(llmRequestBuilder.build().config().get().tools().get()) + .containsExactly( + Tool.builder() + .retrieval( + Retrieval.builder() + .vertexRagStore( + VertexRagStore.builder() + .ragResources( + ImmutableList.of( + VertexRagStoreRagResource.builder() + .ragCorpus("corpus1") + .build())) + .vectorDistanceThreshold(vectorDistanceThreshold) + .build()) + .build()) + .build()); + } else { + // Assert that the function declaration is added instead + assertThat(llmRequestBuilder.build().config().get().tools().get()) + .containsExactly( + Tool.builder() + .functionDeclarations( + ImmutableList.of( + FunctionDeclaration.builder() + .name("testTool") + .description("test description") + .parameters( + Schema.builder() + .properties( + ImmutableMap.of( + "query", + Schema.builder() + .description("The query to retrieve.") + .type("STRING") + .build())) + .type("OBJECT") + .build()) + .build())) + .build()); + } + } + + @Test + public void processLlmRequest_otherModel_doNotAddVertexRagStoreToConfig() { + ImmutableList ragResources = + ImmutableList.of(RagResource.newBuilder().setRagCorpus("corpus1").build()); + Double vectorDistanceThreshold = 0.5; + Integer similarityTopK = 10; + VertexAiRagRetrieval tool = + new VertexAiRagRetrieval( + "testTool", + "test description", + vertexRagServiceClient, + "projects/test-project/locations/us-central1", + ragResources, + vectorDistanceThreshold, + similarityTopK); LlmRequest.Builder llmRequestBuilder = LlmRequest.builder().model("other-model"); ToolContext toolContext = buildToolContext(); GenerateContentConfig initialConfig = GenerateContentConfig.builder().build();