diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java index f85eb7be5..13057cb1a 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java @@ -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 speechEntity = this.retryTemplate + .execute(ctx -> this.audioApi.createSpeech(speechRequest)); - ResponseEntity 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 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 stream(SpeechPrompt speechPrompt) { + + OpenAiAudioApi.SpeechRequest speechRequest = createRequest(speechPrompt); + + Flux> 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) { diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java index b9b4b19c1..fbf51bb78 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java @@ -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 { + 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 transcriptionEntity = this.retryTemplate + .execute(ctx -> this.audioApi.createTranscription(request, StructuredResponse.class)); - ResponseEntity 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 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 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()) diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java index 6f48ef601..7a160d01b 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java @@ -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> apiRequest = createRequest(request); - org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest> 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 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 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 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 embeddings = apiEmbeddingResponse.data() - .stream() - .map(e -> new Embedding(e.embedding(), e.index())) - .toList(); + @SuppressWarnings("unchecked") + private OpenAiApi.EmbeddingRequest> createRequest(EmbeddingRequest request) { + org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest> 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; } } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java index 7c4267fbe..d9cd72374 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java @@ -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 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 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 imageResponseEntity, 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 deleted file mode 100644 index bc5cc0007..000000000 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/common/OpenAiApiException.java +++ /dev/null @@ -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); - } - -} \ No newline at end of file diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java index 2d1c38f29..96a95ba4e 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java @@ -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); diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java index 6c3340404..51bb1505c 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java @@ -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; }