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.
This commit is contained in:
Christian Tzolov
2024-03-01 22:17:02 +01:00
parent 50547188f7
commit b58ef6d5a3
9 changed files with 236 additions and 100 deletions

View File

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

View File

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

View File

@@ -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 <T extends Object, E extends Throwable> void onError(RetryContext context,

View File

@@ -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<HttpHeaders> 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:
* <a href="https://platform.openai.com/docs/models/gpt-4-and-gpt-4-turbo">GPT-4 and GPT-4 Turbo</a> and
* <a href="https://platform.openai.com/docs/models/gpt-3-5-turbo">GPT-3.5 Turbo</a>.
*/
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 doesnt 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:
* <a href="https://platform.openai.com/docs/models/embeddings">Embeddings</a>.
*/
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;
}
}

View File

@@ -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<HttpHeaders> 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.
* <a href="https://platform.openai.com/docs/models/dall-e">DALL·E</a>
*/
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 (

View File

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

View File

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

View File

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

View File

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