diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClient.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClient.java index 1c3bf2866..38edc65b9 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClient.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClient.java @@ -24,12 +24,15 @@ import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.models.ChatChoice; import com.azure.ai.openai.models.ChatCompletions; import com.azure.ai.openai.models.ChatCompletionsOptions; -import com.azure.ai.openai.models.ChatMessage; -import com.azure.ai.openai.models.ChatRole; -import com.azure.ai.openai.models.PromptFilterResult; - +import com.azure.ai.openai.models.ChatRequestAssistantMessage; +import com.azure.ai.openai.models.ChatRequestMessage; +import com.azure.ai.openai.models.ChatRequestSystemMessage; +import com.azure.ai.openai.models.ChatRequestUserMessage; +import com.azure.ai.openai.models.ChatResponseMessage; +import com.azure.ai.openai.models.ContentFilterResultsForPrompt; import org.slf4j.Logger; import org.slf4j.LoggerFactory; + import org.springframework.ai.azure.openai.metadata.AzureOpenAiGenerationMetadata; import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; @@ -48,6 +51,7 @@ import org.springframework.util.Assert; * @author Mark Pollack * @author Ueibin Kim * @author John Blum + * @author Christian Tzolov * @see ChatClient * @see com.azure.ai.openai.OpenAIClient */ @@ -85,7 +89,7 @@ public class AzureOpenAiChatClient implements ChatClient { @Override public String generate(String text) { - ChatMessage azureChatMessage = new ChatMessage(ChatRole.USER, text); + ChatRequestMessage azureChatMessage = new ChatRequestUserMessage(text); ChatCompletionsOptions options = new ChatCompletionsOptions(List.of(azureChatMessage)); options.setTemperature(this.getTemperature()); @@ -98,7 +102,7 @@ public class AzureOpenAiChatClient implements ChatClient { StringBuilder stringBuilder = new StringBuilder(); for (ChatChoice choice : chatCompletions.getChoices()) { - ChatMessage message = choice.getMessage(); + ChatResponseMessage message = choice.getMessage(); if (message != null && message.getContent() != null) { stringBuilder.append(message.getContent()); } @@ -110,43 +114,52 @@ public class AzureOpenAiChatClient implements ChatClient { @Override public ChatResponse generate(Prompt prompt) { - List messages = prompt.getMessages(); - List azureMessages = new ArrayList<>(); - - for (Message message : messages) { - String messageType = message.getMessageType().getValue(); - ChatRole chatRole = ChatRole.fromString(messageType); - azureMessages.add(new ChatMessage(chatRole, message.getContent())); - } + List azureMessages = prompt.getMessages().stream().map(this::fromSpringAiMessage).toList(); ChatCompletionsOptions options = new ChatCompletionsOptions(azureMessages); + options.setTemperature(this.getTemperature()); options.setModel(this.getModel()); + logger.trace("Azure ChatCompletionsOptions: {}", options); ChatCompletions chatCompletions = this.openAIClient.getChatCompletions(this.getModel(), options); + logger.trace("Azure ChatCompletions: {}", chatCompletions); - List generations = new ArrayList<>(); - - for (ChatChoice choice : chatCompletions.getChoices()) { - ChatMessage choiceMessage = choice.getMessage(); - Generation generation = new Generation(choiceMessage.getContent()) - .withChoiceMetadata(generateChoiceMetadata(choice)); - generations.add(generation); - } + List generations = chatCompletions.getChoices() + .stream() + .map(choice -> new Generation(choice.getMessage().getContent()) + .withChoiceMetadata(generateChoiceMetadata(choice))) + .toList(); return new ChatResponse(generations, AzureOpenAiGenerationMetadata.from(chatCompletions)) .withPromptMetadata(generatePromptMetadata(chatCompletions)); } + private ChatRequestMessage fromSpringAiMessage(Message message) { + + switch (message.getMessageType()) { + case USER: + return new ChatRequestUserMessage(message.getContent()); + case SYSTEM: + return new ChatRequestSystemMessage(message.getContent()); + case ASSISTANT: + return new ChatRequestAssistantMessage(message.getContent()); + default: + throw new IllegalArgumentException("Unknown message type " + message.getMessageType()); + } + + } + private ChoiceMetadata generateChoiceMetadata(ChatChoice choice) { return ChoiceMetadata.from(String.valueOf(choice.getFinishReason()), choice.getContentFilterResults()); } private PromptMetadata generatePromptMetadata(ChatCompletions chatCompletions) { - List promptFilterResults = nullSafeList(chatCompletions.getPromptFilterResults()); + List promptFilterResults = nullSafeList( + chatCompletions.getPromptFilterResults()); return PromptMetadata.of(promptFilterResults.stream() .map(promptFilterResult -> PromptFilterMetadata.from(promptFilterResult.getPromptIndex(), @@ -158,4 +171,4 @@ public class AzureOpenAiChatClient implements ChatClient { return list != null ? list : Collections.emptyList(); } -} +} \ No newline at end of file diff --git a/spring-ai-test/src/main/java/org/springframework/ai/test/config/MockAiTestConfiguration.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAiTestConfiguration.java similarity index 99% rename from spring-ai-test/src/main/java/org/springframework/ai/test/config/MockAiTestConfiguration.java rename to models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAiTestConfiguration.java index 1ce3754a6..6e3a402a0 100644 --- a/spring-ai-test/src/main/java/org/springframework/ai/test/config/MockAiTestConfiguration.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAiTestConfiguration.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.test.config; +package org.springframework.ai.azure.openai; import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.status; diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java index 724b13269..a92943ee7 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java @@ -16,13 +16,10 @@ package org.springframework.ai.azure.openai; -import static org.springframework.ai.test.config.MockAiTestConfiguration.SPRING_AI_API_PATH; - import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.OpenAIClientBuilder; import org.springframework.ai.azure.openai.client.AzureOpenAiChatClient; -import org.springframework.ai.test.config.MockAiTestConfiguration; import org.springframework.boot.SpringBootConfiguration; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Import; @@ -58,7 +55,7 @@ public class MockAzureOpenAiTestConfiguration { @Bean OpenAIClient microsoftAzureOpenAiClient(MockWebServer webServer) { - HttpUrl baseUrl = webServer.url(SPRING_AI_API_PATH); + HttpUrl baseUrl = webServer.url(MockAiTestConfiguration.SPRING_AI_API_PATH); return new OpenAIClientBuilder().endpoint(baseUrl.toString()).buildClient(); } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClientMetadataTests.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClientMetadataTests.java index 08f28b29c..853b4529c 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClientMetadataTests.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/client/AzureOpenAiChatClientMetadataTests.java @@ -21,7 +21,10 @@ import static org.assertj.core.api.Assertions.assertThat; import java.nio.charset.StandardCharsets; import com.azure.ai.openai.models.ContentFilterResult; -import com.azure.ai.openai.models.ContentFilterResults; +import com.azure.ai.openai.models.ContentFilterResultDetailsForPrompt; +import com.azure.ai.openai.models.ContentFilterResultsForChoice; +import com.azure.ai.openai.models.ContentFilterResultsForPrompt; +// import com.azure.ai.openai.models.ContentFilterResultsForPrompt; import com.azure.ai.openai.models.ContentFilterSeverity; import org.junit.jupiter.api.Test; @@ -57,6 +60,7 @@ import org.springframework.web.context.request.WebRequest; * Unit Tests for {@link AzureOpenAiChatClient} asserting AI metadata. * * @author John Blum + * @author Christian Tzolov * @since 0.7.0 */ @SpringBootTest @@ -98,7 +102,8 @@ class AzureOpenAiChatClientMetadataTests { assertThat(promptFilterMetadata).isNotNull(); assertThat(promptFilterMetadata.getPromptIndex()).isZero(); - assertContentFilterResults(promptFilterMetadata.getContentFilterMetadata(), ContentFilterSeverity.HIGH); + assertContentFilterResultsForPrompt(promptFilterMetadata.getContentFilterMetadata(), + ContentFilterSeverity.HIGH); } private void assertGenerationMetadata(ChatResponse response) { @@ -126,11 +131,22 @@ class AzureOpenAiChatClientMetadataTests { assertContentFilterResults(choiceMetadata.getContentFilterMetadata()); } - private void assertContentFilterResults(ContentFilterResults contentFilterResults) { + private void assertContentFilterResultsForPrompt(ContentFilterResultDetailsForPrompt contentFilterResultForPrompt, + ContentFilterSeverity selfHarmSeverity) { + + assertThat(contentFilterResultForPrompt).isNotNull(); + assertContentFilterResult(contentFilterResultForPrompt.getHate()); + assertContentFilterResult(contentFilterResultForPrompt.getSelfHarm(), selfHarmSeverity); + assertContentFilterResult(contentFilterResultForPrompt.getSexual()); + assertContentFilterResult(contentFilterResultForPrompt.getViolence()); + + } + + private void assertContentFilterResults(ContentFilterResultsForChoice contentFilterResults) { assertContentFilterResults(contentFilterResults, ContentFilterSeverity.SAFE); } - private void assertContentFilterResults(ContentFilterResults contentFilterResults, + private void assertContentFilterResults(ContentFilterResultsForChoice contentFilterResults, ContentFilterSeverity selfHarmSeverity) { assertThat(contentFilterResults).isNotNull(); diff --git a/pom.xml b/pom.xml index 91449c941..11750e8fa 100644 --- a/pom.xml +++ b/pom.xml @@ -96,7 +96,7 @@ 3.2.0 6.1.1 4.0.2 - 1.0.0-beta.3 + 1.0.0-beta.6 0.6.1 4.31.1 2.22.0