Add support for partially compliant OpenAI APIs

Some LLM providers, such as Groq and OpenRouter, are marketed as OpenAI API compatible.
 However, they often lack full support for the API specification.

 This PR allows to use the simpler message format if no media artifacts are assigned.
This commit is contained in:
Mariusz Bernacki
2024-06-13 20:10:45 +02:00
committed by Christian Tzolov
parent 053fcb0153
commit a6bed95358
2 changed files with 126 additions and 18 deletions

View File

@@ -15,14 +15,24 @@
*/
package org.springframework.ai.openai;
import java.util.ArrayList;
import java.util.Base64;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
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.model.StreamingChatModel;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.ModelOptionsUtils;
@@ -45,17 +55,8 @@ import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.MimeType;
import reactor.core.publisher.Flux;
import java.util.ArrayList;
import java.util.Base64;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
import reactor.core.publisher.Flux;
/**
* {@link ChatModel} and {@link StreamingChatModel} implementation for {@literal OpenAI}
@@ -69,6 +70,7 @@ import java.util.concurrent.ConcurrentHashMap;
* @author Jemin Huh
* @author Grogdunn
* @author Hyunjoon Choi
* @author Mariusz Bernacki
* @see ChatModel
* @see StreamingChatModel
* @see OpenAiApi
@@ -253,18 +255,23 @@ public class OpenAiChatModel extends
Set<String> functionsForThisRequest = new HashSet<>();
List<ChatCompletionMessage> chatCompletionMessages = prompt.getInstructions().stream().map(m -> {
// Add text content.
List<MediaContent> contents = new ArrayList<>(List.of(new MediaContent(m.getContent())));
if (!CollectionUtils.isEmpty(m.getMedia())) {
// Add media content.
contents.addAll(m.getMedia()
Object content;
if (CollectionUtils.isEmpty(m.getMedia())) {
content = m.getContent();
}
else {
List<MediaContent> contentList = new ArrayList<>(List.of(new MediaContent(m.getContent())));
contentList.addAll(m.getMedia()
.stream()
.map(media -> new MediaContent(
new MediaContent.ImageUrl(this.fromMediaData(media.getMimeType(), media.getData()))))
.toList());
content = contentList;
}
return new ChatCompletionMessage(contents, ChatCompletionMessage.Role.valueOf(m.getMessageType().name()));
return new ChatCompletionMessage(content, ChatCompletionMessage.Role.valueOf(m.getMessageType().name()));
}).toList();
ChatCompletionRequest request = new ChatCompletionRequest(chatCompletionMessages, stream);

View File

@@ -0,0 +1,101 @@
/*
* Copyright 2023 - 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.openai.chat;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.MethodSource;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.ai.chat.messages.UserMessage;
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.model.StreamingChatModel;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.ai.openai.api.OpenAiApi;
import reactor.core.publisher.Flux;
import java.util.List;
import java.util.stream.Collectors;
import java.util.stream.Stream;
import static org.assertj.core.api.Assertions.assertThat;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
public class OpenAiCompatibleChatModelIT {
List<Message> conversation = List.of(new SystemMessage("You are a helpful assistant."),
new UserMessage("Are you familiar with pirates from the Golden Age of Piracy?"),
new AssistantMessage("Aye, I be well-versed in the legends of the Golden Age of Piracy!"),
new UserMessage("Tell me about 3 most famous ones."));
static OpenAiChatOptions forModelName(String modelName) {
return OpenAiChatOptions.builder().withModel(modelName).build();
};
static Stream<ChatModel> openAiCompatibleApis() {
Stream.Builder<ChatModel> builder = Stream.builder();
builder.add(new OpenAiChatModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")), forModelName("gpt-3.5-turbo")));
if (System.getenv("GROQ_API_KEY") != null) {
builder.add(new OpenAiChatModel(new OpenAiApi("https://api.groq.com/openai", System.getenv("GROQ_API_KEY")),
forModelName("llama3-8b-8192")));
}
if (System.getenv("OPEN_ROUTER_API_KEY") != null) {
builder.add(new OpenAiChatModel(
new OpenAiApi("https://openrouter.ai/api", System.getenv("OPEN_ROUTER_API_KEY")),
forModelName("meta-llama/llama-3-8b-instruct")));
}
return builder.build();
}
@ParameterizedTest
@MethodSource("openAiCompatibleApis")
void chatCompletion(ChatModel chatModel) {
Prompt prompt = new Prompt(conversation);
ChatResponse response = chatModel.call(prompt);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard");
}
@ParameterizedTest
@MethodSource("openAiCompatibleApis")
void streamCompletion(StreamingChatModel streamingChatModel) {
Prompt prompt = new Prompt(conversation);
Flux<ChatResponse> flux = streamingChatModel.stream(prompt);
List<ChatResponse> responses = flux.collectList().block();
assertThat(responses).hasSizeGreaterThan(1);
String stitchedResponseContent = responses.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
assertThat(stitchedResponseContent).contains("Blackbeard");
}
}