update azure-openai do beta6 and resolve the compatibility issues
This commit is contained in:
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
2
pom.xml
2
pom.xml
@@ -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>
|
||||
|
||||
Reference in New Issue
Block a user