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
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -47,14 +48,16 @@
* A retrieval tool that fetches context from Vertex AI RAG.
*
* <p>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);
private final VertexRagServiceClient vertexRagServiceClient;
private final String parent;
private final List<RagResource> ragResources;
private final Double vectorDistanceThreshold;
private final Integer similarityTopK;
private final VertexRagStore vertexRagStore;
private final RetrieveContextsRequest.VertexRagStore apiVertexRagStore;

Expand All @@ -65,11 +68,30 @@ public VertexAiRagRetrieval(
String parent,
@Nullable List<RagResource> ragResources,
@Nullable Double vectorDistanceThreshold) {
this(
name,
description,
vertexRagServiceClient,
parent,
ragResources,
vectorDistanceThreshold,
null);
}

public VertexAiRagRetrieval(
String name,
String description,
VertexRagServiceClient vertexRagServiceClient,
String parent,
@Nullable List<RagResource> 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();
Expand All @@ -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
Expand Down Expand Up @@ -135,10 +160,18 @@ public Single<Map<String, Object>> runAsync(Map<String, Object> 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);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -72,20 +73,26 @@ public void runAsync_withResults_returnsContexts() throws Exception {
ImmutableList<RagResource> 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);
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)
Expand All @@ -112,20 +119,26 @@ public void runAsync_noResults_returnsNoResultFoundMessage() throws Exception {
ImmutableList<RagResource> 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);
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)
Expand Down Expand Up @@ -154,14 +167,16 @@ public void processLlmRequest_gemini2Model_addVertexRagStoreToConfig() {
ImmutableList<RagResource> 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);
vectorDistanceThreshold,
similarityTopK);
LlmRequest.Builder llmRequestBuilder = LlmRequest.builder().model("gemini-2-pro");
ToolContext toolContext = buildToolContext();

Expand All @@ -181,7 +196,8 @@ public void processLlmRequest_gemini2Model_addVertexRagStoreToConfig() {
VertexRagStoreRagResource.builder()
.ragCorpus("corpus1")
.build()))
.vectorDistanceThreshold(0.5)
.vectorDistanceThreshold(vectorDistanceThreshold)
.similarityTopK(similarityTopK)
.build())
.build())
.build());
Expand Down Expand Up @@ -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<RagResource> ragResources =
ImmutableList.of(RagResource.newBuilder().setRagCorpus("corpus1").build());
Double vectorDistanceThreshold = 0.5;
Expand All @@ -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<RagResource> 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();
Expand Down