From e268975a80880dbfcddb29adbb5651e4f4cfe55c Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Sun, 7 Apr 2024 21:57:04 +0200 Subject: [PATCH] Fix Geminie GenerativeModel handling between calls Resolves #560 --- .../ai/vertexai/gemini/VertexAiGeminiChatClient.java | 6 ++---- .../VertexAiGeminiChatClientFunctionCallingIT.java | 2 +- .../gemini/tool/FunctionCallWithFunctionBeanIT.java | 7 +++++++ .../gemini/tool/FunctionCallWithPromptFunctionIT.java | 9 +++++++++ 4 files changed, 19 insertions(+), 5 deletions(-) diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java index 1641bc9ac..bd6b0d09a 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java @@ -78,8 +78,6 @@ public class VertexAiGeminiChatClient private final GenerationConfig generationConfig; - private GenerativeModel generativeModel; - public enum GeminiMessageType { USER("user"), @@ -140,7 +138,6 @@ public class VertexAiGeminiChatClient this.vertexAI = vertexAI; this.defaultOptions = options; this.generationConfig = toGenerationConfig(options); - this.generativeModel = new GenerativeModel(options.getModel(), vertexAI); } // https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/gemini @@ -204,7 +201,8 @@ public class VertexAiGeminiChatClient Set functionsForThisRequest = new HashSet<>(); GenerationConfig generationConfig = this.generationConfig; - GenerativeModel generativeModel = this.generativeModel; + + GenerativeModel generativeModel = new GenerativeModel(this.defaultOptions.getModel(), this.vertexAI); VertexAiGeminiChatOptions updatedRuntimeOptions = null; diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java index 5d7083d79..cc9a153bd 100644 --- a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java +++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java @@ -146,7 +146,7 @@ public class VertexAiGeminiChatClientFunctionCallingIT { public void functionCallTestInferredOpenApiSchemaStream() { UserMessage userMessage = new UserMessage( - "What's the weather like in San Francisco, in Paris and in Tokyo? Use Multi-turn function calling."); + "What's the weather like in San Francisco, in Paris and in Tokyo, Japan? Use Multi-turn function calling. Provide answer for all requested locations."); List messages = new ArrayList<>(List.of(userMessage)); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java index cc2590667..d2066f0b2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java @@ -83,6 +83,13 @@ class FunctionCallWithFunctionBeanIT { assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + response = chatClient + .call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().build())); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).doesNotContain("30", "10", "15"); + }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java index a2dc42154..b654fb124 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java @@ -77,6 +77,15 @@ public class FunctionCallWithPromptFunctionIT { logger.info("Response: {}", response); assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + + // Verify that no function call is made. + response = chatClient + .call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().build())); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).doesNotContain("30", "10", "15"); + }); }