diff --git a/models/spring-ai-vertex-ai-gemini/pom.xml b/models/spring-ai-vertex-ai-gemini/pom.xml
index c903c8dc0..da9eb0beb 100644
--- a/models/spring-ai-vertex-ai-gemini/pom.xml
+++ b/models/spring-ai-vertex-ai-gemini/pom.xml
@@ -84,6 +84,12 @@
test
+
+ io.micrometer
+ micrometer-observation-test
+ test
+
+
diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java
index b4a980e46..cfd98c87e 100644
--- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java
+++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java
@@ -37,6 +37,10 @@ import org.springframework.ai.chat.model.AbstractToolCallSupport;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
+import org.springframework.ai.chat.observation.ChatModelObservationContext;
+import org.springframework.ai.chat.observation.ChatModelObservationConvention;
+import org.springframework.ai.chat.observation.ChatModelObservationDocumentation;
+import org.springframework.ai.chat.observation.DefaultChatModelObservationConvention;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.ChatModelDescription;
@@ -46,6 +50,7 @@ import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackContext;
import org.springframework.ai.model.function.FunctionCallingOptions;
import org.springframework.ai.retry.RetryUtils;
+import org.springframework.ai.vertexai.gemini.common.VertexAiGeminiConstants;
import org.springframework.ai.vertexai.gemini.metadata.VertexAiUsage;
import org.springframework.beans.factory.DisposableBean;
import org.springframework.lang.NonNull;
@@ -53,6 +58,8 @@ import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
+
+import io.micrometer.observation.ObservationRegistry;
import reactor.core.publisher.Flux;
import java.util.ArrayList;
@@ -68,10 +75,13 @@ import java.util.Set;
* @author luocongqiu
* @author Chris Turchin
* @author Mark Pollack
+ * @author Soby Chacko
* @since 0.8.1
*/
public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements ChatModel, DisposableBean {
+ private static final ChatModelObservationConvention DEFAULT_OBSERVATION_CONVENTION = new DefaultChatModelObservationConvention();
+
private final VertexAI vertexAI;
private final VertexAiGeminiChatOptions defaultOptions;
@@ -83,6 +93,16 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
private final GenerationConfig generationConfig;
+ /**
+ * Observation registry used for instrumentation.
+ */
+ private final ObservationRegistry observationRegistry;
+
+ /**
+ * Conventions to use for generating observations.
+ */
+ private final ChatModelObservationConvention observationConvention = DEFAULT_OBSERVATION_CONVENTION;
+
public enum GeminiMessageType {
USER("user"),
@@ -153,6 +173,13 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
public VertexAiGeminiChatModel(VertexAI vertexAI, VertexAiGeminiChatOptions options,
FunctionCallbackContext functionCallbackContext, List toolFunctionCallbacks,
RetryTemplate retryTemplate) {
+ this(vertexAI, options, functionCallbackContext, toolFunctionCallbacks, retryTemplate,
+ ObservationRegistry.NOOP);
+ }
+
+ public VertexAiGeminiChatModel(VertexAI vertexAI, VertexAiGeminiChatOptions options,
+ FunctionCallbackContext functionCallbackContext, List toolFunctionCallbacks,
+ RetryTemplate retryTemplate, ObservationRegistry observationRegistry) {
super(functionCallbackContext, options, toolFunctionCallbacks);
@@ -165,41 +192,60 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
this.defaultOptions = options;
this.generationConfig = toGenerationConfig(options);
this.retryTemplate = retryTemplate;
+ this.observationRegistry = observationRegistry;
}
// https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/gemini
@Override
public ChatResponse call(Prompt prompt) {
- return retryTemplate.execute(context -> {
- var geminiRequest = createGeminiRequest(prompt);
- GenerateContentResponse response = this.getContentResponse(geminiRequest);
+ VertexAiGeminiChatOptions vertexAiGeminiChatOptions = vertexAiGeminiChatOptions(prompt);
- List generations = response.getCandidatesList()
- .stream()
- .map(this::responseCandiateToGeneration)
- .flatMap(List::stream)
- .toList();
+ ChatModelObservationContext observationContext = ChatModelObservationContext.builder()
+ .prompt(prompt)
+ .provider(VertexAiGeminiConstants.PROVIDER_NAME)
+ .requestOptions(vertexAiGeminiChatOptions)
+ .build();
- ChatResponse chatResponse = new ChatResponse(generations, toChatResponseMetadata(response));
+ ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION
+ .observation(this.observationConvention, DEFAULT_OBSERVATION_CONVENTION, () -> observationContext,
+ this.observationRegistry)
+ .observe(() -> this.retryTemplate.execute(context -> {
- if (!isProxyToolCalls(prompt, this.defaultOptions)
- && isToolCall(chatResponse, Set.of(FinishReason.STOP.name()))) {
- var toolCallConversation = handleToolCalls(prompt, chatResponse);
- // Recursively call the call method with the tool call message
- // conversation that contains the call responses.
- return this.call(new Prompt(toolCallConversation, prompt.getOptions()));
- }
+ var geminiRequest = createGeminiRequest(prompt, vertexAiGeminiChatOptions);
+
+ GenerateContentResponse generateContentResponse = this.getContentResponse(geminiRequest);
+
+ List generations = generateContentResponse.getCandidatesList()
+ .stream()
+ .map(this::responseCandiateToGeneration)
+ .flatMap(List::stream)
+ .toList();
+
+ ChatResponse chatResponse = new ChatResponse(generations,
+ toChatResponseMetadata(generateContentResponse));
+
+ observationContext.setResponse(chatResponse);
+ return chatResponse;
+ }));
+
+ if (!isProxyToolCalls(prompt, this.defaultOptions) && isToolCall(response, Set.of(FinishReason.STOP.name()))) {
+ var toolCallConversation = handleToolCalls(prompt, response);
+ // Recursively call the call method with the tool call message
+ // conversation that contains the call responses.
+ return this.call(new Prompt(toolCallConversation, prompt.getOptions()));
+ }
+
+ return response;
- return chatResponse;
- });
}
@Override
public Flux stream(Prompt prompt) {
try {
- var request = createGeminiRequest(prompt);
+ VertexAiGeminiChatOptions vertexAiGeminiChatOptions = vertexAiGeminiChatOptions(prompt);
+ var request = createGeminiRequest(prompt, vertexAiGeminiChatOptions);
ResponseStream responseStream = request.model
.generateContentStream(request.contents);
@@ -281,10 +327,26 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
public record GeminiRequest(List contents, GenerativeModel model) {
}
+ private VertexAiGeminiChatOptions vertexAiGeminiChatOptions(Prompt prompt) {
+ VertexAiGeminiChatOptions updatedRuntimeOptions = VertexAiGeminiChatOptions.builder().build();
+ if (prompt.getOptions() != null) {
+ updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
+ VertexAiGeminiChatOptions.class);
+
+ }
+
+ updatedRuntimeOptions = ModelOptionsUtils.merge(updatedRuntimeOptions, this.defaultOptions,
+ VertexAiGeminiChatOptions.class);
+
+ return updatedRuntimeOptions;
+
+ }
+
/**
- * Tests access to the {@link #createGeminiRequest(Prompt)} method.
+ * Tests access to the {@link #createGeminiRequest(Prompt, VertexAiGeminiChatOptions)}
+ * method.
*/
- GeminiRequest createGeminiRequest(Prompt prompt) {
+ GeminiRequest createGeminiRequest(Prompt prompt, VertexAiGeminiChatOptions updatedRuntimeOptions) {
Set functionsForThisRequest = new HashSet<>();
@@ -293,13 +355,10 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
var generativeModelBuilder = new GenerativeModel.Builder().setModelName(this.defaultOptions.getModel())
.setVertexAi(this.vertexAI);
- VertexAiGeminiChatOptions updatedRuntimeOptions = VertexAiGeminiChatOptions.builder().build();
-
if (prompt.getOptions() != null) {
if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) {
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions,
FunctionCallingOptions.class, VertexAiGeminiChatOptions.class);
-
}
else {
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
@@ -312,9 +371,6 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
functionsForThisRequest.addAll(this.defaultOptions.getFunctions());
}
- updatedRuntimeOptions = ModelOptionsUtils.merge(updatedRuntimeOptions, this.defaultOptions,
- VertexAiGeminiChatOptions.class);
-
if (updatedRuntimeOptions != null) {
if (StringUtils.hasText(updatedRuntimeOptions.getModel())
diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/common/VertexAiGeminiConstants.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/common/VertexAiGeminiConstants.java
new file mode 100644
index 000000000..4369a2e79
--- /dev/null
+++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/common/VertexAiGeminiConstants.java
@@ -0,0 +1,28 @@
+/*
+ * Copyright 2024-2024 the original author or authors.
+ *
+ * Licensed under the Apache License, Version 2.0 (the "License");
+ * you may not use this file except in compliance with the License.
+ * You may obtain a copy of the License at
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.springframework.ai.vertexai.gemini.common;
+
+import org.springframework.ai.observation.conventions.AiProvider;
+
+/**
+ * @author Soby Chacko
+ */
+public class VertexAiGeminiConstants {
+
+ public static final String PROVIDER_NAME = AiProvider.VERTEX_AI.value();
+
+}
diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java
index 1d4dce04c..acbe2672f 100644
--- a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java
+++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java
@@ -53,7 +53,7 @@ public class CreateGeminiRequestTests {
var client = new VertexAiGeminiChatModel(vertexAI,
VertexAiGeminiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
- GeminiRequest request = client.createGeminiRequest(new Prompt("Test message content"));
+ GeminiRequest request = client.createGeminiRequest(new Prompt("Test message content"), null);
assertThat(request.contents()).hasSize(1);
@@ -61,8 +61,10 @@ public class CreateGeminiRequestTests {
assertThat(request.model().getModelName()).isEqualTo("DEFAULT_MODEL");
assertThat(request.model().getGenerationConfig().getTemperature()).isEqualTo(66.6f);
- request = client.createGeminiRequest(new Prompt("Test message content",
- VertexAiGeminiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()));
+ request = client.createGeminiRequest(
+ new Prompt("Test message content",
+ VertexAiGeminiChatOptions.builder().withModel("PROMPT_MODEL").withTemperature(99.9).build()),
+ null);
assertThat(request.contents()).hasSize(1);
@@ -82,7 +84,7 @@ public class CreateGeminiRequestTests {
var client = new VertexAiGeminiChatModel(vertexAI,
VertexAiGeminiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6).build());
- GeminiRequest request = client.createGeminiRequest(new Prompt(List.of(systemMessage, userMessage)));
+ GeminiRequest request = client.createGeminiRequest(new Prompt(List.of(systemMessage, userMessage)), null);
assertThat(request.model().getModelName()).isEqualTo("DEFAULT_MODEL");
assertThat(request.model().getGenerationConfig().getTemperature()).isEqualTo(66.6f);
@@ -119,7 +121,8 @@ public class CreateGeminiRequestTests {
.withDescription("Get the weather in location")
.withResponseConverter((response) -> "" + response.temp() + response.unit())
.build()))
- .build()));
+ .build()),
+ null);
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
@@ -148,7 +151,7 @@ public class CreateGeminiRequestTests {
.build()))
.build());
- var request = client.createGeminiRequest(new Prompt("Test message content"));
+ var request = client.createGeminiRequest(new Prompt("Test message content"), null);
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
@@ -164,7 +167,7 @@ public class CreateGeminiRequestTests {
// Explicitly enable the function
request = client.createGeminiRequest(new Prompt("Test message content",
- VertexAiGeminiChatOptions.builder().withFunction(TOOL_FUNCTION_NAME).build()));
+ VertexAiGeminiChatOptions.builder().withFunction(TOOL_FUNCTION_NAME).build()), null);
assertThat(request.model().getTools()).hasSize(1);
assertThat(request.model().getTools().get(0).getFunctionDeclarations(0).getName())
@@ -178,7 +181,8 @@ public class CreateGeminiRequestTests {
.withName(TOOL_FUNCTION_NAME)
.withDescription("Overridden function description")
.build()))
- .build()));
+ .build()),
+ null);
assertThat(request.model().getTools()).hasSize(1);
assertThat(request.model().getTools().get(0).getFunctionDeclarations(0).getName())
@@ -206,7 +210,7 @@ public class CreateGeminiRequestTests {
.withResponseMimeType("application/json")
.build());
- GeminiRequest request = client.createGeminiRequest(new Prompt("Test message content"));
+ GeminiRequest request = client.createGeminiRequest(new Prompt("Test message content"), null);
assertThat(request.contents()).hasSize(1);
diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiChatModelObservationIT.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiChatModelObservationIT.java
new file mode 100644
index 000000000..817d9a7cc
--- /dev/null
+++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiChatModelObservationIT.java
@@ -0,0 +1,156 @@
+/*
+ * Copyright 2024 the original author or authors.
+ *
+ * Licensed under the Apache License, Version 2.0 (the "License");
+ * you may not use this file except in compliance with the License.
+ * You may obtain a copy of the License at
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.springframework.ai.vertexai.gemini;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.util.List;
+
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
+
+import org.springframework.ai.chat.metadata.ChatResponseMetadata;
+import org.springframework.ai.chat.model.ChatResponse;
+import org.springframework.ai.chat.observation.ChatModelObservationDocumentation;
+import org.springframework.ai.chat.observation.DefaultChatModelObservationConvention;
+import org.springframework.ai.chat.prompt.Prompt;
+import org.springframework.ai.observation.conventions.AiOperationType;
+import org.springframework.ai.observation.conventions.AiProvider;
+import org.springframework.beans.factory.annotation.Autowired;
+import org.springframework.boot.SpringBootConfiguration;
+import org.springframework.boot.test.context.SpringBootTest;
+import org.springframework.context.annotation.Bean;
+import org.springframework.retry.support.RetryTemplate;
+
+import com.google.cloud.vertexai.Transport;
+import com.google.cloud.vertexai.VertexAI;
+import io.micrometer.observation.tck.TestObservationRegistry;
+import io.micrometer.observation.tck.TestObservationRegistryAssert;
+
+/**
+ * @author Soby Chacko
+ */
+@SpringBootTest
+@EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_PROJECT_ID", matches = ".*")
+@EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_LOCATION", matches = ".*")
+public class VertexAiChatModelObservationIT {
+
+ @Autowired
+ TestObservationRegistry observationRegistry;
+
+ @Autowired
+ VertexAiGeminiChatModel chatModel;
+
+ @BeforeEach
+ void beforeEach() {
+ observationRegistry.clear();
+ }
+
+ @Test
+ void observationForChatOperation() {
+
+ var options = VertexAiGeminiChatOptions.builder()
+ .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_PRO.getValue())
+ .withTemperature(0.7)
+ .withStopSequences(List.of("this-is-the-end"))
+ .withMaxOutputTokens(2048)
+ .withTopP(1.0)
+ .build();
+
+ Prompt prompt = new Prompt("Why does a raven look like a desk?", options);
+
+ ChatResponse chatResponse = chatModel.call(prompt);
+ assertThat(chatResponse.getResult().getOutput().getContent()).isNotEmpty();
+
+ ChatResponseMetadata responseMetadata = chatResponse.getMetadata();
+ assertThat(responseMetadata).isNotNull();
+
+ validate(responseMetadata);
+ }
+
+ private void validate(ChatResponseMetadata responseMetadata) {
+ TestObservationRegistryAssert.assertThat(observationRegistry)
+ .doesNotHaveAnyRemainingCurrentObservation()
+ .hasObservationWithNameEqualTo(DefaultChatModelObservationConvention.DEFAULT_NAME)
+ .that()
+ .hasLowCardinalityKeyValue(
+ ChatModelObservationDocumentation.LowCardinalityKeyNames.AI_OPERATION_TYPE.asString(),
+ AiOperationType.CHAT.value())
+ .hasLowCardinalityKeyValue(ChatModelObservationDocumentation.LowCardinalityKeyNames.AI_PROVIDER.asString(),
+ AiProvider.VERTEX_AI.value())
+ .hasLowCardinalityKeyValue(
+ ChatModelObservationDocumentation.LowCardinalityKeyNames.REQUEST_MODEL.asString(),
+ VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_PRO.getValue())
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_MAX_TOKENS.asString(), "2048")
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_STOP_SEQUENCES.asString(),
+ "[\"this-is-the-end\"]")
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_TEMPERATURE.asString(), "0.7")
+ .doesNotHaveHighCardinalityKeyValueWithKey(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_TOP_K.asString())
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_TOP_P.asString(), "1.0")
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.RESPONSE_FINISH_REASONS.asString(),
+ "[\"STOP\"]")
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
+ String.valueOf(responseMetadata.getUsage().getPromptTokens()))
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
+ String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
+ .hasHighCardinalityKeyValue(
+ ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
+ String.valueOf(responseMetadata.getUsage().getTotalTokens()))
+ .hasBeenStarted()
+ .hasBeenStopped();
+ }
+
+ @SpringBootConfiguration
+ static class Config {
+
+ @Bean
+ public TestObservationRegistry observationRegistry() {
+ return TestObservationRegistry.create();
+ }
+
+ @Bean
+ public VertexAI vertexAiApi() {
+ String projectId = System.getenv("VERTEX_AI_GEMINI_PROJECT_ID");
+ String location = System.getenv("VERTEX_AI_GEMINI_LOCATION");
+ return new VertexAI.Builder().setProjectId(projectId)
+ .setLocation(location)
+ .setTransport(Transport.REST)
+ .build();
+ }
+
+ @Bean
+ public VertexAiGeminiChatModel vertexAiEmbedding(VertexAI vertexAi,
+ TestObservationRegistry observationRegistry) {
+ return new VertexAiGeminiChatModel(vertexAi,
+ VertexAiGeminiChatOptions.builder()
+ .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_PRO)
+ .build(),
+ null, List.of(), RetryTemplate.defaultInstance(), observationRegistry);
+ }
+
+ }
+
+}