From a6bed95358ab2cd5a3b3ec9a1a614a1b7ae610aa Mon Sep 17 00:00:00 2001 From: Mariusz Bernacki Date: Thu, 13 Jun 2024 20:10:45 +0200 Subject: [PATCH] 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. --- .../ai/openai/OpenAiChatModel.java | 43 ++++---- .../chat/OpenAiCompatibleChatModelIT.java | 101 ++++++++++++++++++ 2 files changed, 126 insertions(+), 18 deletions(-) create mode 100644 models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiCompatibleChatModelIT.java diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index 277bd982d..4d807816d 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -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 functionsForThisRequest = new HashSet<>(); List chatCompletionMessages = prompt.getInstructions().stream().map(m -> { - // Add text content. - List 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 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); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiCompatibleChatModelIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiCompatibleChatModelIT.java new file mode 100644 index 000000000..4a4e30719 --- /dev/null +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiCompatibleChatModelIT.java @@ -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 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 openAiCompatibleApis() { + Stream.Builder 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 flux = streamingChatModel.stream(prompt); + + List 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"); + } + +}