update azure-openai do beta6 and resolve the compatibility issues

This commit is contained in:
Christian Tzolov
2023-12-22 11:55:38 +01:00
parent aebfd719a7
commit c8de949123
5 changed files with 60 additions and 34 deletions

View File

@@ -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<Message> messages = prompt.getMessages();
List<ChatMessage> 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<ChatRequestMessage> 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<Generation> 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<Generation> 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<PromptFilterResult> promptFilterResults = nullSafeList(chatCompletions.getPromptFilterResults());
List<ContentFilterResultsForPrompt> 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();
}
}
}

View File

@@ -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;

View File

@@ -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();
}

View File

@@ -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();

View File

@@ -96,7 +96,7 @@
<spring-boot.version>3.2.0</spring-boot.version>
<spring-framework.version>6.1.1</spring-framework.version>
<stringtemplate.version>4.0.2</stringtemplate.version>
<azure-open-ai-client.version>1.0.0-beta.3</azure-open-ai-client.version>
<azure-open-ai-client.version>1.0.0-beta.6</azure-open-ai-client.version>
<jtokkit.version>0.6.1</jtokkit.version>
<victools.version>4.31.1</victools.version>
<bedrockruntime.version>2.22.0</bedrockruntime.version>