From b58ef6d5a3906492aee6501bb1a07a809cc2da3e Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Fri, 1 Mar 2024 22:17:02 +0100 Subject: [PATCH] OpenAI api improvements - Factor out OpenAiApiClientErrorException and OpenAiApiException out of OpenAiApi into standalone common package. - Move the HTTP error handling and header setup into a shared ApiUtils. - Add ChatModel and EmbeddingModel enums to OpenAiApi. - Add ImageModel enum ot OpenAiImageClient. --- .../ai/openai/OpenAiChatClient.java | 2 +- .../ai/openai/OpenAiEmbeddingClient.java | 2 +- .../ai/openai/OpenAiImageClient.java | 4 +- .../ai/openai/api/OpenAiApi.java | 154 ++++++++++++------ .../ai/openai/api/OpenAiImageApi.java | 76 ++++----- .../ai/openai/api/common/ApiUtils.java | 62 +++++++ .../common/OpenAiApiClientErrorException.java | 18 ++ .../openai/api/common/OpenAiApiException.java | 16 ++ .../vertexai/palm2/api/VertexAiPaLm2Api.java | 2 + 9 files changed, 236 insertions(+), 100 deletions(-) create mode 100644 models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/ApiUtils.java create mode 100644 models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiClientErrorException.java create mode 100644 models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiException.java diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java index 62f7bfe61..7be58affa 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java @@ -43,8 +43,8 @@ import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.Role; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ToolCall; +import org.springframework.ai.openai.api.common.OpenAiApiException; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest; -import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiException; import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata; import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor; import org.springframework.http.ResponseEntity; diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java index f5db94a4c..1ff64a99b 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java @@ -33,8 +33,8 @@ import org.springframework.ai.embedding.EmbeddingResponseMetadata; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.EmbeddingList; -import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiException; import org.springframework.ai.openai.api.OpenAiApi.Usage; +import org.springframework.ai.openai.api.common.OpenAiApiException; import org.springframework.retry.RetryCallback; import org.springframework.retry.RetryContext; import org.springframework.retry.RetryListener; diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java index 97b032123..8f28ccda4 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java @@ -30,8 +30,8 @@ import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; import org.springframework.ai.image.ImageResponseMetadata; import org.springframework.ai.model.ModelOptionsUtils; -import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiImageApi; +import org.springframework.ai.openai.api.common.OpenAiApiException; import org.springframework.ai.openai.metadata.OpenAiImageGenerationMetadata; import org.springframework.ai.openai.metadata.OpenAiImageResponseMetadata; import org.springframework.http.ResponseEntity; @@ -59,7 +59,7 @@ public class OpenAiImageClient implements ImageClient { public final RetryTemplate retryTemplate = RetryTemplate.builder() .maxAttempts(10) - .retryOn(OpenAiApi.OpenAiApiException.class) + .retryOn(OpenAiApiException.class) .exponentialBackoff(Duration.ofMillis(2000), 5, Duration.ofMillis(3 * 60000)) .withListener(new RetryListener() { public void onError(RetryContext context, diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java index 5cab9c081..58e7d146f 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java @@ -16,11 +16,8 @@ package org.springframework.ai.openai.api; -import java.io.IOException; -import java.nio.charset.StandardCharsets; import java.util.List; import java.util.Map; -import java.util.function.Consumer; import java.util.function.Predicate; import com.fasterxml.jackson.annotation.JsonInclude; @@ -30,17 +27,12 @@ import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; import org.springframework.ai.model.ModelOptionsUtils; +import org.springframework.ai.openai.api.common.ApiUtils; import org.springframework.boot.context.properties.bind.ConstructorBinding; import org.springframework.core.ParameterizedTypeReference; -import org.springframework.http.HttpHeaders; -import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; -import org.springframework.http.client.ClientHttpResponse; -import org.springframework.lang.NonNull; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; -import org.springframework.util.StreamUtils; -import org.springframework.web.client.ResponseErrorHandler; import org.springframework.web.client.RestClient; import org.springframework.web.reactive.function.client.WebClient; @@ -89,72 +81,128 @@ public class OpenAiApi { */ public OpenAiApi(String baseUrl, String openAiToken, RestClient.Builder restClientBuilder) { - Consumer jsonContentHeaders = headers -> { - headers.setBearerAuth(openAiToken); - headers.setContentType(MediaType.APPLICATION_JSON); - }; - - var responseErrorHandler = new ResponseErrorHandler() { - - @Override - public boolean hasError(@NonNull ClientHttpResponse response) throws IOException { - return response.getStatusCode().isError(); - } - - @Override - public void handleError(@NonNull ClientHttpResponse response) throws IOException { - if (response.getStatusCode().isError()) { - String error = StreamUtils.copyToString(response.getBody(), StandardCharsets.UTF_8); - String message = String.format("%s - %s", response.getStatusCode().value(), error); - if (response.getStatusCode().is4xxClientError()) { - throw new OpenAiApiClientErrorException(message); - } - throw new OpenAiApiException(message); - } - } - }; - this.restClient = restClientBuilder .baseUrl(baseUrl) - .defaultHeaders(jsonContentHeaders) - .defaultStatusHandler(responseErrorHandler) + .defaultHeaders(ApiUtils.getJsonContentHeaders(openAiToken)) + .defaultStatusHandler(ApiUtils.DEFAULT_RESPONSE_ERROR_HANDLER) .build(); this.webClient = WebClient.builder() .baseUrl(baseUrl) - .defaultHeaders(jsonContentHeaders) + .defaultHeaders(ApiUtils.getJsonContentHeaders(openAiToken)) .build(); } - /** - * Non HTTP Error related exceptions + * OpenAI Chat Completion Models: + * GPT-4 and GPT-4 Turbo and + * GPT-3.5 Turbo. */ - public static class OpenAiApiException extends RuntimeException { + enum ChatModel { + /** + * (New) GPT-4 Turbo - latest GPT-4 model intended to reduce cases + * of “laziness” where the model doesn’t complete a task. + * Returns a maximum of 4,096 output tokens. + * Context window: 128k tokens + */ + GPT_4_0125_PREVIEW("gpt-4-0125-preview"), - public OpenAiApiException(String message) { - super(message); + /** + * Currently points to gpt-4-0125-preview - model featuring improved + * instruction following, JSON mode, reproducible outputs, + * parallel function calling, and more. + * Returns a maximum of 4,096 output tokens + * Context window: 128k tokens + */ + GPT_4_TURBO_PREVIEW("gpt-4-turbo-preview"), + + /** + * GPT-4 with the ability to understand images, in addition + * to all other GPT-4 Turbo capabilities. Currently points + * to gpt-4-1106-vision-preview. + * Returns a maximum of 4,096 output tokens + * Context window: 128k tokens + */ + GPT_4_VISION_PREVIEW("gpt-4-vision-preview"), + + /** + * Currently points to gpt-4-0613. + * Snapshot of gpt-4 from June 13th 2023 with improved + * function calling support. + * Context window: 8k tokens + */ + GPT_4("gpt-4"), + + /** + * Currently points to gpt-4-32k-0613. + * Snapshot of gpt-4-32k from June 13th 2023 with improved + * function calling support. + * Context window: 32k tokens + */ + GPT_4_32K("gpt-4-32k"), + + /** + *Currently points to gpt-3.5-turbo-0125. + * model with higher accuracy at responding in requested + * formats and a fix for a bug which caused a text + * encoding issue for non-English language function calls. + * Returns a maximum of 4,096 + * Context window: 16k tokens + */ + GPT_3_5_TURBO("gpt-3.5-turbo"), + + /** + * GPT-3.5 Turbo model with improved instruction following, + * JSON mode, reproducible outputs, parallel function calling, + * and more. Returns a maximum of 4,096 output tokens. + * Context window: 16k tokens. + */ + GPT_3_5_TURBO_1106("gpt-3.5-turbo-1106"); + + public final String value; + + ChatModel(String value) { + this.value = value; } - public OpenAiApiException(String message, Throwable cause) { - super(message, cause); + public String getValue() { + return value; } } /** - * Thrown on 4xx client errors, such as 401 - Incorrect API key provided, - * 401 - You must be a member of an organization to use the API, - * 429 - Rate limit reached for requests, 429 - You exceeded your current quota - * , please check your plan and billing details. + * OpenAI Embeddings Models: + * Embeddings. */ - public static class OpenAiApiClientErrorException extends RuntimeException { + enum EmbeddingModel { - public OpenAiApiClientErrorException(String message) { - super(message); + /** + * Most capable embedding model for both english and non-english tasks. + * DIMENSION: 3072 + */ + TEXT_EMBEDDING_3_LARGE("text-embedding-3-large"), + + /** + * Increased performance over 2nd generation ada embedding model. + * DIMENSION: 1536 + */ + TEXT_EMBEDDING_3_SMALL("text-embedding-3-small"), + + /** + * Most capable 2nd generation embedding model, replacing 16 first + * generation models. + * DIMENSION: 1536 + */ + TEXT_EMBEDDING_ADA_002("text-embedding-ada-002"); + + public final String value; + + EmbeddingModel(String value) { + this.value = value; } - public OpenAiApiClientErrorException(String message, Throwable cause) { - super(message, cause); + public String getValue() { + return value; } } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiImageApi.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiImageApi.java index 92b3607cf..8858c4e5d 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiImageApi.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiImageApi.java @@ -15,25 +15,14 @@ */ package org.springframework.ai.openai.api; -import java.io.IOException; -import java.nio.charset.StandardCharsets; import java.util.List; -import java.util.function.Consumer; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonProperty; -import com.fasterxml.jackson.databind.DeserializationFeature; -import com.fasterxml.jackson.databind.ObjectMapper; -import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiClientErrorException; -import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiException; -import org.springframework.http.HttpHeaders; -import org.springframework.http.MediaType; +import org.springframework.ai.openai.api.common.ApiUtils; import org.springframework.http.ResponseEntity; -import org.springframework.http.client.ClientHttpResponse; import org.springframework.util.Assert; -import org.springframework.util.StreamUtils; -import org.springframework.web.client.ResponseErrorHandler; import org.springframework.web.client.RestClient; /** @@ -50,8 +39,6 @@ public class OpenAiImageApi { private final RestClient restClient; - private final ObjectMapper objectMapper; - /** * Create a new OpenAI Image api with base URL set to https://api.openai.com * @param openAiToken OpenAI apiKey. @@ -62,39 +49,42 @@ public class OpenAiImageApi { public OpenAiImageApi(String baseUrl, String openAiToken, RestClient.Builder restClientBuilder) { - this.objectMapper = new ObjectMapper().configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); - - Consumer jsonContentHeaders = headers -> { - headers.setBearerAuth(openAiToken); - headers.setContentType(MediaType.APPLICATION_JSON); - }; - - var responseErrorHandler = new ResponseErrorHandler() { - - @Override - public boolean hasError(ClientHttpResponse response) throws IOException { - return response.getStatusCode().isError(); - } - - @Override - public void handleError(ClientHttpResponse response) throws IOException { - if (response.getStatusCode().isError()) { - String error = StreamUtils.copyToString(response.getBody(), StandardCharsets.UTF_8); - String message = String.format("%s - %s", response.getStatusCode().value(), error); - if (response.getStatusCode().is4xxClientError()) { - throw new OpenAiApiClientErrorException(message); - } - throw new OpenAiApiException(message); - } - } - }; - this.restClient = restClientBuilder.baseUrl(baseUrl) - .defaultHeaders(jsonContentHeaders) - .defaultStatusHandler(responseErrorHandler) + .defaultHeaders(ApiUtils.getJsonContentHeaders(openAiToken)) + .defaultStatusHandler(ApiUtils.DEFAULT_RESPONSE_ERROR_HANDLER) .build(); } + /** + * OpenAI Image API model. + * DALL·E + */ + enum ImageModel { + + /** + * The latest DALL·E model released in Nov 2023. + */ + DALL_E_3("dall-e-3"), + + /** + * The previous DALL·E model released in Nov 2022. The 2nd iteration of DALL·E + * with more realistic, accurate, and 4x greater resolution images than the + * original model. + */ + DALL_E_2("dall-e-2"); + + private final String model; + + ImageModel(String model) { + this.model = model; + } + + public String model() { + return this.model; + } + + } + // @formatter:off @JsonInclude(JsonInclude.Include.NON_NULL) public record OpenAiImageRequest ( diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/ApiUtils.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/ApiUtils.java new file mode 100644 index 000000000..48d056258 --- /dev/null +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/ApiUtils.java @@ -0,0 +1,62 @@ +/* + * Copyright 2024-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.api.common; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.util.function.Consumer; + +import org.springframework.http.HttpHeaders; +import org.springframework.http.MediaType; +import org.springframework.http.client.ClientHttpResponse; +import org.springframework.lang.NonNull; +import org.springframework.util.StreamUtils; +import org.springframework.web.client.ResponseErrorHandler; + +/** + * @author Christian Tzolov + */ +public class ApiUtils { + + public static Consumer getJsonContentHeaders(String apiKey) { + return (headers) -> { + headers.setBearerAuth(apiKey); + headers.setContentType(MediaType.APPLICATION_JSON); + }; + }; + + public static final ResponseErrorHandler DEFAULT_RESPONSE_ERROR_HANDLER = new ResponseErrorHandler() { + + @Override + public boolean hasError(@NonNull ClientHttpResponse response) throws IOException { + return response.getStatusCode().isError(); + } + + @Override + public void handleError(@NonNull ClientHttpResponse response) throws IOException { + if (response.getStatusCode().isError()) { + String error = StreamUtils.copyToString(response.getBody(), StandardCharsets.UTF_8); + String message = String.format("%s - %s", response.getStatusCode().value(), error); + if (response.getStatusCode().is4xxClientError()) { + throw new OpenAiApiClientErrorException(message); + } + throw new OpenAiApiException(message); + } + } + }; + +} diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiClientErrorException.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiClientErrorException.java new file mode 100644 index 000000000..d02ea975c --- /dev/null +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiClientErrorException.java @@ -0,0 +1,18 @@ +package org.springframework.ai.openai.api.common; + +/** + * Thrown on 4xx client errors, such as 401 - Incorrect API key provided, 401 - You must + * be a member of an organization to use the API, 429 - Rate limit reached for requests, + * 429 - You exceeded your current quota , please check your plan and billing details. + */ +public class OpenAiApiClientErrorException extends RuntimeException { + + public OpenAiApiClientErrorException(String message) { + super(message); + } + + public OpenAiApiClientErrorException(String message, Throwable cause) { + super(message, cause); + } + +} \ No newline at end of file diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiException.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiException.java new file mode 100644 index 000000000..ee8d86875 --- /dev/null +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiException.java @@ -0,0 +1,16 @@ +package org.springframework.ai.openai.api.common; + +/** + * Non HTTP Error related exceptions + */ +public class OpenAiApiException extends RuntimeException { + + public OpenAiApiException(String message) { + super(message); + } + + public OpenAiApiException(String message, Throwable cause) { + super(message, cause); + } + +} \ No newline at end of file diff --git a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/api/VertexAiPaLm2Api.java b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/api/VertexAiPaLm2Api.java index a8573f0df..d8454fc81 100644 --- a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/api/VertexAiPaLm2Api.java +++ b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/api/VertexAiPaLm2Api.java @@ -84,6 +84,8 @@ import org.springframework.web.client.RestClient; * * https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models * + * https://ai.google.dev/api/rest#rest-resource:-v1.models + * * @author Christian Tzolov */ public class VertexAiPaLm2Api {