Consolidate retry config for OpenAI

Signed-off-by: Thomas Vitale <ThomasVitale@users.noreply.github.com>
This commit is contained in:
Thomas Vitale
2024-07-23 18:22:39 +02:00
parent 6270d627f1
commit 6c88654404
7 changed files with 146 additions and 160 deletions

View File

@@ -22,7 +22,6 @@ import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.openai.api.OpenAiAudioApi;
import org.springframework.ai.openai.api.OpenAiAudioApi.SpeechRequest.AudioResponseFormat;
import org.springframework.ai.openai.api.common.OpenAiApiException;
import org.springframework.ai.openai.audio.speech.Speech;
import org.springframework.ai.openai.audio.speech.SpeechModel;
import org.springframework.ai.openai.audio.speech.SpeechPrompt;
@@ -30,18 +29,18 @@ import org.springframework.ai.openai.audio.speech.SpeechResponse;
import org.springframework.ai.openai.audio.speech.StreamingSpeechModel;
import org.springframework.ai.openai.metadata.audio.OpenAiAudioSpeechResponseMetadata;
import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import reactor.core.publisher.Flux;
import java.time.Duration;
/**
* OpenAI audio speech client implementation for backed by {@link OpenAiAudioApi}.
*
* @author Ahmed Yousri
* @author Hyunjoon Choi
* @author Thomas Vitale
* @see OpenAiAudioApi
* @since 1.0.0-M1
*/
@@ -63,11 +62,7 @@ public class OpenAiAudioSpeechModel implements SpeechModel, StreamingSpeechModel
/**
* The retry template used to retry the OpenAI Audio API calls.
*/
public final RetryTemplate retryTemplate = RetryTemplate.builder()
.maxAttempts(10)
.retryOn(OpenAiApiException.class)
.exponentialBackoff(Duration.ofMillis(2000), 5, Duration.ofMillis(3 * 60000))
.build();
private final RetryTemplate retryTemplate;
/**
* Low-level access to the OpenAI Audio API.
@@ -98,10 +93,25 @@ public class OpenAiAudioSpeechModel implements SpeechModel, StreamingSpeechModel
* options.
*/
public OpenAiAudioSpeechModel(OpenAiAudioApi audioApi, OpenAiAudioSpeechOptions options) {
this(audioApi, options, RetryUtils.DEFAULT_RETRY_TEMPLATE);
}
/**
* Initializes a new instance of the OpenAiAudioSpeechModel class with the provided
* OpenAiAudioApi and options.
* @param audioApi The OpenAiAudioApi to use for speech synthesis.
* @param options The OpenAiAudioSpeechOptions containing the speech synthesis
* options.
* @param retryTemplate The retry template.
*/
public OpenAiAudioSpeechModel(OpenAiAudioApi audioApi, OpenAiAudioSpeechOptions options,
RetryTemplate retryTemplate) {
Assert.notNull(audioApi, "OpenAiAudioApi must not be null");
Assert.notNull(options, "OpenAiSpeechOptions must not be null");
Assert.notNull(options, "RetryTemplate must not be null");
this.audioApi = audioApi;
this.defaultOptions = options;
this.retryTemplate = retryTemplate;
}
@Override
@@ -113,40 +123,43 @@ public class OpenAiAudioSpeechModel implements SpeechModel, StreamingSpeechModel
@Override
public SpeechResponse call(SpeechPrompt speechPrompt) {
return this.retryTemplate.execute(ctx -> {
OpenAiAudioApi.SpeechRequest speechRequest = createRequest(speechPrompt);
OpenAiAudioApi.SpeechRequest speechRequest = createRequestBody(speechPrompt);
ResponseEntity<byte[]> speechEntity = this.retryTemplate
.execute(ctx -> this.audioApi.createSpeech(speechRequest));
ResponseEntity<byte[]> speechEntity = this.audioApi.createSpeech(speechRequest);
var speech = speechEntity.getBody();
var speech = speechEntity.getBody();
if (speech == null) {
logger.warn("No speech response returned for speechRequest: {}", speechRequest);
return new SpeechResponse(new Speech(new byte[0]));
}
if (speech == null) {
logger.warn("No speech response returned for speechRequest: {}", speechRequest);
return new SpeechResponse(new Speech(new byte[0]));
}
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(speechEntity);
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(speechEntity);
return new SpeechResponse(new Speech(speech), new OpenAiAudioSpeechResponseMetadata(rateLimits));
});
return new SpeechResponse(new Speech(speech), new OpenAiAudioSpeechResponseMetadata(rateLimits));
}
/**
* Streams the audio response for the given speech prompt.
* @param prompt The speech prompt containing the text and options for speech
* @param speechPrompt The speech prompt containing the text and options for speech
* synthesis.
* @return A Flux of SpeechResponse objects containing the streamed audio and
* metadata.
*/
@Override
public Flux<SpeechResponse> stream(SpeechPrompt prompt) {
return this.audioApi.stream(this.createRequestBody(prompt))
.map(entity -> new SpeechResponse(new Speech(entity.getBody()), new OpenAiAudioSpeechResponseMetadata(
OpenAiResponseHeaderExtractor.extractAiResponseHeaders(entity))));
public Flux<SpeechResponse> stream(SpeechPrompt speechPrompt) {
OpenAiAudioApi.SpeechRequest speechRequest = createRequest(speechPrompt);
Flux<ResponseEntity<byte[]>> speechEntity = this.retryTemplate
.execute(ctx -> this.audioApi.stream(speechRequest));
return speechEntity.map(entity -> new SpeechResponse(new Speech(entity.getBody()),
new OpenAiAudioSpeechResponseMetadata(OpenAiResponseHeaderExtractor.extractAiResponseHeaders(entity))));
}
private OpenAiAudioApi.SpeechRequest createRequestBody(SpeechPrompt request) {
private OpenAiAudioApi.SpeechRequest createRequest(SpeechPrompt request) {
OpenAiAudioSpeechOptions options = this.defaultOptions;
if (request.getOptions() != null) {

View File

@@ -56,6 +56,7 @@ import org.springframework.util.Assert;
*
* @author Michael Lavelle
* @author Christian Tzolov
* @author Thomas Vitale
* @see OpenAiAudioApi
* @since 0.8.1
*/
@@ -65,7 +66,7 @@ public class OpenAiAudioTranscriptionModel implements Model<AudioTranscriptionPr
private final OpenAiAudioTranscriptionOptions defaultOptions;
public final RetryTemplate retryTemplate;
private final RetryTemplate retryTemplate;
private final OpenAiAudioApi audioApi;
@@ -80,8 +81,7 @@ public class OpenAiAudioTranscriptionModel implements Model<AudioTranscriptionPr
.withModel(OpenAiAudioApi.WhisperModel.WHISPER_1.getValue())
.withResponseFormat(OpenAiAudioApi.TranscriptResponseFormat.JSON)
.withTemperature(0.7f)
.build(),
RetryUtils.DEFAULT_RETRY_TEMPLATE);
.build());
}
/**
@@ -119,74 +119,71 @@ public class OpenAiAudioTranscriptionModel implements Model<AudioTranscriptionPr
}
@Override
public AudioTranscriptionResponse call(AudioTranscriptionPrompt request) {
public AudioTranscriptionResponse call(AudioTranscriptionPrompt transcriptionPrompt) {
return this.retryTemplate.execute(ctx -> {
Resource audioResource = transcriptionPrompt.getInstructions();
Resource audioResource = request.getInstructions();
OpenAiAudioApi.TranscriptionRequest request = createRequest(transcriptionPrompt);
OpenAiAudioApi.TranscriptionRequest requestBody = createRequestBody(request);
if (request.responseFormat().isJsonType()) {
if (requestBody.responseFormat().isJsonType()) {
ResponseEntity<StructuredResponse> transcriptionEntity = this.retryTemplate
.execute(ctx -> this.audioApi.createTranscription(request, StructuredResponse.class));
ResponseEntity<StructuredResponse> transcriptionEntity = this.audioApi.createTranscription(requestBody,
StructuredResponse.class);
var transcription = transcriptionEntity.getBody();
if (transcription == null) {
logger.warn("No transcription returned for request: {}", audioResource);
return new AudioTranscriptionResponse(null);
}
AudioTranscription transcript = new AudioTranscription(transcription.text());
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(transcriptionEntity);
return new AudioTranscriptionResponse(transcript,
OpenAiAudioTranscriptionResponseMetadata.from(transcriptionEntity.getBody())
.withRateLimit(rateLimits));
var transcription = transcriptionEntity.getBody();
if (transcription == null) {
logger.warn("No transcription returned for request: {}", audioResource);
return new AudioTranscriptionResponse(null);
}
else {
ResponseEntity<String> transcriptionEntity = this.audioApi.createTranscription(requestBody,
String.class);
AudioTranscription transcript = new AudioTranscription(transcription.text());
var transcription = transcriptionEntity.getBody();
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(transcriptionEntity);
if (transcription == null) {
logger.warn("No transcription returned for request: {}", audioResource);
return new AudioTranscriptionResponse(null);
}
return new AudioTranscriptionResponse(transcript,
OpenAiAudioTranscriptionResponseMetadata.from(transcriptionEntity.getBody())
.withRateLimit(rateLimits));
AudioTranscription transcript = new AudioTranscription(transcription);
}
else {
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(transcriptionEntity);
ResponseEntity<String> transcriptionEntity = this.retryTemplate
.execute(ctx -> this.audioApi.createTranscription(request, String.class));
return new AudioTranscriptionResponse(transcript,
OpenAiAudioTranscriptionResponseMetadata.from(transcriptionEntity.getBody())
.withRateLimit(rateLimits));
var transcription = transcriptionEntity.getBody();
if (transcription == null) {
logger.warn("No transcription returned for request: {}", audioResource);
return new AudioTranscriptionResponse(null);
}
});
AudioTranscription transcript = new AudioTranscription(transcription);
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(transcriptionEntity);
return new AudioTranscriptionResponse(transcript,
OpenAiAudioTranscriptionResponseMetadata.from(transcriptionEntity.getBody())
.withRateLimit(rateLimits));
}
}
OpenAiAudioApi.TranscriptionRequest createRequestBody(AudioTranscriptionPrompt request) {
OpenAiAudioApi.TranscriptionRequest createRequest(AudioTranscriptionPrompt transcriptionPrompt) {
OpenAiAudioTranscriptionOptions options = this.defaultOptions;
if (request.getOptions() != null) {
if (request.getOptions() instanceof OpenAiAudioTranscriptionOptions runtimeOptions) {
if (transcriptionPrompt.getOptions() != null) {
if (transcriptionPrompt.getOptions() instanceof OpenAiAudioTranscriptionOptions runtimeOptions) {
options = this.merge(runtimeOptions, options);
}
else {
throw new IllegalArgumentException("Prompt options are not of type TranscriptionOptions: "
+ request.getOptions().getClass().getSimpleName());
+ transcriptionPrompt.getOptions().getClass().getSimpleName());
}
}
return OpenAiAudioApi.TranscriptionRequest.builder()
.withFile(toBytes(request.getInstructions()))
.withFile(toBytes(transcriptionPrompt.getInstructions()))
.withResponseFormat(options.getResponseFormat())
.withPrompt(options.getPrompt())
.withTemperature(options.getTemperature())

View File

@@ -69,8 +69,7 @@ public class OpenAiEmbeddingModel extends AbstractEmbeddingModel {
*/
public OpenAiEmbeddingModel(OpenAiApi openAiApi, MetadataMode metadataMode) {
this(openAiApi, metadataMode,
OpenAiEmbeddingOptions.builder().withModel(OpenAiApi.DEFAULT_EMBEDDING_MODEL).build(),
RetryUtils.DEFAULT_RETRY_TEMPLATE);
OpenAiEmbeddingOptions.builder().withModel(OpenAiApi.DEFAULT_EMBEDDING_MODEL).build());
}
/**
@@ -110,42 +109,45 @@ public class OpenAiEmbeddingModel extends AbstractEmbeddingModel {
return this.embed(document.getFormattedContent(this.metadataMode));
}
@SuppressWarnings("unchecked")
@Override
public EmbeddingResponse call(EmbeddingRequest request) {
return this.retryTemplate.execute(ctx -> {
org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<List<String>> apiRequest = createRequest(request);
org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<List<String>> apiRequest = (this.defaultOptions != null)
? new org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<>(request.getInstructions(),
this.defaultOptions.getModel(), this.defaultOptions.getEncodingFormat(),
this.defaultOptions.getDimensions(), this.defaultOptions.getUser())
: new org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<>(request.getInstructions(),
OpenAiApi.DEFAULT_EMBEDDING_MODEL);
EmbeddingList<OpenAiApi.Embedding> apiEmbeddingResponse = this.retryTemplate
.execute(ctx -> this.openAiApi.embeddings(apiRequest).getBody());
if (request.getOptions() != null && !EmbeddingOptions.EMPTY.equals(request.getOptions())) {
apiRequest = ModelOptionsUtils.merge(request.getOptions(), apiRequest,
org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest.class);
}
if (apiEmbeddingResponse == null) {
logger.warn("No embeddings returned for request: {}", request);
return new EmbeddingResponse(List.of());
}
EmbeddingList<OpenAiApi.Embedding> apiEmbeddingResponse = this.openAiApi.embeddings(apiRequest).getBody();
var metadata = new EmbeddingResponseMetadata(apiEmbeddingResponse.model(),
OpenAiUsage.from(apiEmbeddingResponse.usage()));
if (apiEmbeddingResponse == null) {
logger.warn("No embeddings returned for request: {}", request);
return new EmbeddingResponse(List.of());
}
List<Embedding> embeddings = apiEmbeddingResponse.data()
.stream()
.map(e -> new Embedding(e.embedding(), e.index()))
.toList();
var metadata = new EmbeddingResponseMetadata(apiEmbeddingResponse.model(),
OpenAiUsage.from(apiEmbeddingResponse.usage()));
return new EmbeddingResponse(embeddings, metadata);
}
List<Embedding> embeddings = apiEmbeddingResponse.data()
.stream()
.map(e -> new Embedding(e.embedding(), e.index()))
.toList();
@SuppressWarnings("unchecked")
private OpenAiApi.EmbeddingRequest<List<String>> createRequest(EmbeddingRequest request) {
org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<List<String>> apiRequest = (this.defaultOptions != null)
? new org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<>(request.getInstructions(),
this.defaultOptions.getModel(), this.defaultOptions.getEncodingFormat(),
this.defaultOptions.getDimensions(), this.defaultOptions.getUser())
: new org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<>(request.getInstructions(),
OpenAiApi.DEFAULT_EMBEDDING_MODEL);
return new EmbeddingResponse(embeddings, metadata);
if (request.getOptions() != null && !EmbeddingOptions.EMPTY.equals(request.getOptions())) {
apiRequest = ModelOptionsUtils.merge(request.getOptions(), apiRequest,
org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest.class);
}
});
return apiRequest;
}
}

View File

@@ -41,6 +41,7 @@ import java.util.List;
* @author Mark Pollack
* @author Christian Tzolov
* @author Hyunjoon Choi
* @author Thomas Vitale
* @since 0.8.0
*/
public class OpenAiImageModel implements ImageModel {
@@ -90,30 +91,32 @@ public class OpenAiImageModel implements ImageModel {
@Override
public ImageResponse call(ImagePrompt imagePrompt) {
return this.retryTemplate.execute(ctx -> {
String instructions = imagePrompt.getInstructions().get(0).getText();
OpenAiImageApi.OpenAiImageRequest imageRequest = createRequest(imagePrompt);
OpenAiImageApi.OpenAiImageRequest imageRequest = new OpenAiImageApi.OpenAiImageRequest(instructions,
OpenAiImageApi.DEFAULT_IMAGE_MODEL);
ResponseEntity<OpenAiImageApi.OpenAiImageResponse> imageResponseEntity = this.retryTemplate
.execute(ctx -> this.openAiImageApi.createImage(imageRequest));
if (this.defaultOptions != null) {
imageRequest = ModelOptionsUtils.merge(this.defaultOptions, imageRequest,
OpenAiImageApi.OpenAiImageRequest.class);
}
return convertResponse(imageResponseEntity, imageRequest);
}
if (imagePrompt.getOptions() != null) {
imageRequest = ModelOptionsUtils.merge(toOpenAiImageOptions(imagePrompt.getOptions()), imageRequest,
OpenAiImageApi.OpenAiImageRequest.class);
}
private OpenAiImageApi.OpenAiImageRequest createRequest(ImagePrompt imagePrompt) {
String instructions = imagePrompt.getInstructions().get(0).getText();
// Make the request
ResponseEntity<OpenAiImageApi.OpenAiImageResponse> imageResponseEntity = this.openAiImageApi
.createImage(imageRequest);
OpenAiImageApi.OpenAiImageRequest imageRequest = new OpenAiImageApi.OpenAiImageRequest(instructions,
OpenAiImageApi.DEFAULT_IMAGE_MODEL);
// Convert to org.springframework.ai.model derived ImageResponse data type
return convertResponse(imageResponseEntity, imageRequest);
});
if (this.defaultOptions != null) {
imageRequest = ModelOptionsUtils.merge(this.defaultOptions, imageRequest,
OpenAiImageApi.OpenAiImageRequest.class);
}
if (imagePrompt.getOptions() != null) {
imageRequest = ModelOptionsUtils.merge(toOpenAiImageOptions(imagePrompt.getOptions()), imageRequest,
OpenAiImageApi.OpenAiImageRequest.class);
}
return imageRequest;
}
private ImageResponse convertResponse(ResponseEntity<OpenAiImageApi.OpenAiImageResponse> imageResponseEntity,

View File

@@ -1,31 +0,0 @@
/*
* 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.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

@@ -44,7 +44,7 @@ public class TranscriptionRequestTests {
.withTemperature(66.6f)
.build());
var request = client.createRequestBody(
var request = client.createRequest(
new AudioTranscriptionPrompt(new DefaultResourceLoader().getResource("classpath:/test.png")));
assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
@@ -68,16 +68,16 @@ public class TranscriptionRequestTests {
.withTemperature(66.6f)
.build());
var request = client.createRequestBody(
new AudioTranscriptionPrompt(new DefaultResourceLoader().getResource("classpath:/test.png"),
OpenAiAudioTranscriptionOptions.builder()
.withModel("RUNTIME_MODEL")
.withResponseFormat(TranscriptResponseFormat.JSON)
.withLanguage("bg")
.withPrompt("Prompt2")
.withGranularityType(GranularityType.SEGMENT)
.withTemperature(99.9f)
.build()));
var request = client
.createRequest(new AudioTranscriptionPrompt(new DefaultResourceLoader().getResource("classpath:/test.png"),
OpenAiAudioTranscriptionOptions.builder()
.withModel("RUNTIME_MODEL")
.withResponseFormat(TranscriptResponseFormat.JSON)
.withLanguage("bg")
.withPrompt("Prompt2")
.withGranularityType(GranularityType.SEGMENT)
.withTemperature(99.9f)
.build()));
assertThat(request.model()).isEqualTo("RUNTIME_MODEL");
assertThat(request.responseFormat()).isEqualByComparingTo(TranscriptResponseFormat.JSON);

View File

@@ -49,6 +49,7 @@ import org.springframework.web.reactive.function.client.WebClient;
/**
* @author Christian Tzolov
* @author Stefan Vassilev
* @author Thomas Vitale
*/
@AutoConfiguration(after = { RestClientAutoConfiguration.class, WebClientAutoConfiguration.class,
SpringAiRetryAutoConfiguration.class })
@@ -174,8 +175,9 @@ public class OpenAiAutoConfiguration {
@ConditionalOnProperty(prefix = OpenAiAudioSpeechProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true",
matchIfMissing = true)
public OpenAiAudioSpeechModel openAiAudioSpeechClient(OpenAiConnectionProperties commonProperties,
OpenAiAudioSpeechProperties speechProperties, RestClient.Builder restClientBuilder,
WebClient.Builder webClientBuilder, ResponseErrorHandler responseErrorHandler) {
OpenAiAudioSpeechProperties speechProperties, RetryTemplate retryTemplate,
RestClient.Builder restClientBuilder, WebClient.Builder webClientBuilder,
ResponseErrorHandler responseErrorHandler) {
String apiKey = StringUtils.hasText(speechProperties.getApiKey()) ? speechProperties.getApiKey()
: commonProperties.getApiKey();
@@ -192,7 +194,7 @@ public class OpenAiAutoConfiguration {
responseErrorHandler);
OpenAiAudioSpeechModel openAiSpeechModel = new OpenAiAudioSpeechModel(openAiAudioApi,
speechProperties.getOptions());
speechProperties.getOptions(), retryTemplate);
return openAiSpeechModel;
}