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); + } + + } + +}