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 06cd07a601..3c2d841a76 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 @@ -23,7 +23,9 @@ import java.util.Objects; import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; import com.openai.core.http.Headers; +import com.openai.core.http.HttpResponse; import com.openai.models.audio.speech.SpeechCreateParams; import com.openai.models.audio.speech.SpeechModel; import io.micrometer.observation.ObservationRegistry; @@ -52,6 +54,7 @@ * @author Jonghoon Park * @author Ilayaperumal Gopinathan * @author Sebastien Deleuze + * @author guan xu */ public final class OpenAiAudioSpeechModel implements TextToSpeechModel { @@ -139,7 +142,9 @@ public TextToSpeechResponse call(TextToSpeechPrompt prompt) { SpeechCreateParams params = paramsBuilder.build(); - com.openai.core.http.HttpResponse httpResponse = this.openAiClient.audio().speech().create(params); + RequestOptions requestOptions = this.buildRequestOptions(mergedOptions); + + HttpResponse httpResponse = this.openAiClient.audio().speech().create(params, requestOptions); Headers headers = httpResponse.headers(); byte[] audioBytes; @@ -170,6 +175,20 @@ public Flux stream(TextToSpeechPrompt prompt) { return Flux.just(call(prompt)); } + /** + * Creates a RequestOptions instance from the given audio speech options. + * @param options the audio speech options + * @return a RequestOptions instance + */ + private RequestOptions buildRequestOptions(OpenAiAudioSpeechOptions options) { + Assert.notNull(options, "Options cannot be null"); + RequestOptions.Builder requestOptionsBuilder = RequestOptions.builder(); + if (options.getTimeout() != null) { + requestOptionsBuilder.timeout(options.getTimeout()); + } + return requestOptionsBuilder.build(); + } + /** * @since 2.0.0 */ 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 2c3687f629..441dacdc38 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 @@ -26,6 +26,7 @@ import com.openai.client.OpenAIClient; import com.openai.client.OpenAIClientAsync; import com.openai.core.MultipartField; +import com.openai.core.RequestOptions; import com.openai.models.audio.transcriptions.TranscriptionCreateParams; import com.openai.models.audio.transcriptions.TranscriptionCreateResponse; import com.openai.models.audio.transcriptions.TranscriptionStreamEvent; @@ -84,7 +85,8 @@ public Builder mutate() { } private OpenAiAudioTranscriptionModel(Builder builder) { - this.options = builder.options != null ? builder.options : OpenAiAudioTranscriptionOptions.builder().build(); + this.options = Objects.requireNonNullElseGet(builder.options, + () -> OpenAiAudioTranscriptionOptions.builder().build()); this.openAiClient = Objects.requireNonNullElseGet(builder.openAiClient, () -> OpenAiSetup.setupSyncClient(this.options.getBaseUrl(), this.options.getApiKey(), this.options.getCredential(), this.options.getMicrosoftDeploymentName(), @@ -123,12 +125,16 @@ public AudioTranscriptionResponse call(AudioTranscriptionPrompt transcriptionPro byte[] audioBytes = toBytes(audioResource); String filename = getFilename(audioResource); - TranscriptionCreateParams params = buildParams(mergedOptions, audioBytes, filename); + TranscriptionCreateParams params = this.buildParams(mergedOptions, audioBytes, filename); if (logger.isTraceEnabled()) { logger.trace("OpenAiAudioTranscriptionModel call with model: " + mergedOptions.getModel()); } - TranscriptionCreateResponse response = this.openAiClient.audio().transcriptions().create(params); + RequestOptions requestOptions = this.buildRequestOptions(mergedOptions); + + TranscriptionCreateResponse response = this.openAiClient.audio() + .transcriptions() + .create(params, requestOptions); String text = extractText(response); AudioTranscription transcript = new AudioTranscription(text); return new AudioTranscriptionResponse(transcript, new AudioTranscriptionResponseMetadata()); @@ -146,14 +152,16 @@ public Flux stream(AudioTranscriptionPrompt transcri byte[] audioBytes = toBytes(audioResource); String filename = getFilename(audioResource); - TranscriptionCreateParams params = buildParams(mergedOptions, audioBytes, filename); + TranscriptionCreateParams params = this.buildParams(mergedOptions, audioBytes, filename); if (logger.isTraceEnabled()) { logger.trace("OpenAiAudioTranscriptionModel stream with model: " + mergedOptions.getModel()); } + RequestOptions requestOptions = this.buildRequestOptions(mergedOptions); + Flux chunks = Flux.create(sink -> this.openAiClientAsync.audio() .transcriptions() - .createStreaming(params) + .createStreaming(params, requestOptions) .subscribe(sink::next) .onCompleteFuture() .whenComplete((unused, throwable) -> { @@ -206,6 +214,20 @@ private TranscriptionCreateParams buildParams(OpenAiAudioTranscriptionOptions op return builder.build(); } + /** + * Creates a RequestOptions instance from the given transcription options. + * @param options the transcription options + * @return a RequestOptions instance + */ + private RequestOptions buildRequestOptions(OpenAiAudioTranscriptionOptions options) { + Assert.notNull(options, "Options cannot be null"); + RequestOptions.Builder requestOptionsBuilder = RequestOptions.builder(); + if (options.getTimeout() != null) { + requestOptionsBuilder.timeout(options.getTimeout()); + } + return requestOptionsBuilder.build(); + } + private static String extractText(TranscriptionCreateResponse response) { if (response.isTranscription()) { return response.asTranscription().text(); 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 7add8ccc47..ce42194f3f 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 @@ -36,6 +36,7 @@ import com.openai.client.OpenAIClient; import com.openai.client.OpenAIClientAsync; import com.openai.core.JsonValue; +import com.openai.core.RequestOptions; import com.openai.errors.OpenAIInvalidDataException; import com.openai.models.FunctionDefinition; import com.openai.models.FunctionParameters; @@ -123,6 +124,7 @@ * @author Eric Bottard * @author Taewoong Kim * @author Jewoo Shin + * @author guan xu */ public final class OpenAiChatModel implements ChatModel { @@ -204,7 +206,8 @@ public ChatResponse call(Prompt prompt) { */ private ChatResponse internalCall(Prompt prompt, @Nullable ChatResponse previousChatResponse) { - ChatCompletionCreateParams request = createRequest(prompt, false); + ChatCompletionCreateParams request = this.createRequest(prompt, false); + RequestOptions requestOptions = this.buildRequestOptions(prompt); ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) @@ -216,7 +219,7 @@ private ChatResponse internalCall(Prompt prompt, @Nullable ChatResponse previous this.observationRegistry) .observe(() -> { - ChatCompletion chatCompletion = this.openAiClient.chat().completions().create(request); + ChatCompletion chatCompletion = this.openAiClient.chat().completions().create(request, requestOptions); List choices = chatCompletion.choices(); if (choices.isEmpty()) { @@ -268,7 +271,8 @@ public Flux stream(Prompt prompt) { */ private Flux internalStream(Prompt prompt) { return Flux.deferContextual(contextView -> { - ChatCompletionCreateParams request = createRequest(prompt, true); + ChatCompletionCreateParams request = this.createRequest(prompt, true); + RequestOptions requestOptions = this.buildRequestOptions(prompt); ConcurrentHashMap roleMap = new ConcurrentHashMap<>(); ConcurrentHashMap reasoningMap = new ConcurrentHashMap<>(); final ChatModelObservationContext observationContext = ChatModelObservationContext.builder() @@ -293,9 +297,9 @@ private Flux internalStream(Prompt prompt) { } // Convert from AsyncStreamResponse to Flux - Flux chunks = Flux.create(sink -> this.openAiClientAsync.chat() + Flux chunks = Flux.create(sink -> this.openAiClientAsync.chat() .completions() - .createStreaming(request) + .createStreaming(request, requestOptions) .subscribe(sink::next) .onCompleteFuture() .whenComplete((unused, throwable) -> { @@ -916,6 +920,23 @@ else if (json.equals("required")) { return builder.build(); } + /** + * Creates a RequestOptions instance from the given prompt. + * @param prompt the prompt containing messages and options + * @return a RequestOptions instance + */ + private RequestOptions buildRequestOptions(Prompt prompt) { + Assert.notNull(prompt, "Prompt cannot be null"); + Assert.isInstanceOf(OpenAiChatOptions.class, prompt.getOptions(), + "Prompt options must be OpenAiChatOptions type"); + OpenAiChatOptions chatOptions = (OpenAiChatOptions) prompt.getOptions(); + RequestOptions.Builder requestOptionsBuilder = RequestOptions.builder(); + if (chatOptions.getTimeout() != null) { + requestOptionsBuilder.timeout(chatOptions.getTimeout()); + } + return requestOptionsBuilder.build(); + } + private Map toolCallAdditionalPropertiesFromMetadata(AssistantMessage assistantMessage) { Object value = assistantMessage.getMetadata().get(TOOL_CALL_ADDITIONAL_PROPERTIES_METADATA_KEY); if (!(value instanceof Map rawMap)) { @@ -1476,8 +1497,8 @@ public Builder httpClientBuilderCustomizers(List OpenAiChatOptions.builder().build()); ObservationRegistry resolvedObservationRegistry = Objects.requireNonNullElse(this.observationRegistry, ObservationRegistry.NOOP); 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 c5e60de34d..01210e3ab7 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 @@ -21,6 +21,7 @@ import java.util.Objects; import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; import com.openai.models.embeddings.CreateEmbeddingResponse; import com.openai.models.embeddings.EmbeddingCreateParams; import io.micrometer.observation.ObservationRegistry; @@ -55,6 +56,7 @@ * @author Thomas Vitale * @author Christian Tzolov * @author Josh Long + * @author guan xu */ public class OpenAiEmbeddingModel extends AbstractEmbeddingModel { @@ -185,7 +187,7 @@ public static Builder builder() { } private OpenAiEmbeddingModel(Builder builder) { - this.options = builder.options != null ? builder.options : OpenAiEmbeddingOptions.builder().build(); + this.options = Objects.requireNonNullElseGet(builder.options, () -> OpenAiEmbeddingOptions.builder().build()); this.metadataMode = Objects.requireNonNullElse(builder.metadataMode, MetadataMode.EMBED); this.observationRegistry = Objects.requireNonNullElse(builder.observationRegistry, ObservationRegistry.NOOP); this.openAiClient = Objects.requireNonNullElseGet(builder.openAiClient, @@ -233,6 +235,8 @@ public EmbeddingResponse call(EmbeddingRequest embeddingRequest) { + embeddingCreateParams); } + RequestOptions requestOptions = this.buildRequestOptions(options); + var observationContext = EmbeddingModelObservationContext.builder() .embeddingRequest(embeddingRequestWithMergedOptions) .provider(AiProvider.OPENAI.value()) @@ -243,7 +247,8 @@ public EmbeddingResponse call(EmbeddingRequest embeddingRequest) { .observation(this.observationConvention, DEFAULT_OBSERVATION_CONVENTION, () -> observationContext, this.observationRegistry) .observe(() -> { - CreateEmbeddingResponse response = this.openAiClient.embeddings().create(embeddingCreateParams); + CreateEmbeddingResponse response = this.openAiClient.embeddings() + .create(embeddingCreateParams, requestOptions); var embeddingResponse = generateEmbeddingResponse(response); observationContext.setResponse(embeddingResponse); @@ -251,6 +256,20 @@ public EmbeddingResponse call(EmbeddingRequest embeddingRequest) { })); } + /** + * Creates a RequestOptions instance from the given embedding options. + * @param options the embedding options + * @return a RequestOptions instance + */ + private RequestOptions buildRequestOptions(OpenAiEmbeddingOptions options) { + Assert.notNull(options, "Options cannot be null"); + RequestOptions.Builder requestOptionsBuilder = RequestOptions.builder(); + if (options.getTimeout() != null) { + requestOptionsBuilder.timeout(options.getTimeout()); + } + return requestOptionsBuilder.build(); + } + private EmbeddingResponse generateEmbeddingResponse(CreateEmbeddingResponse response) { List data = generateEmbeddingList(response.data()); 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 166b597b75..c1b04b7afa 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 @@ -21,6 +21,7 @@ import java.util.Objects; import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; import com.openai.models.images.ImageGenerateParams; import io.micrometer.observation.ObservationRegistry; import org.apache.commons.logging.Log; @@ -52,6 +53,7 @@ * @author Hyunjoon Choi * @author Christian Tzolov * @author Mark Pollack + * @author guan xu */ public class OpenAiImageModel implements ImageModel { @@ -158,7 +160,7 @@ public static Builder builder() { } private OpenAiImageModel(Builder builder) { - this.options = builder.options != null ? builder.options : OpenAiImageOptions.builder().build(); + this.options = Objects.requireNonNullElseGet(builder.options, () -> OpenAiImageOptions.builder().build()); this.observationRegistry = Objects.requireNonNullElse(builder.observationRegistry, ObservationRegistry.NOOP); this.openAiClient = Objects.requireNonNullElseGet(builder.openAiClient, () -> OpenAiSetup.setupSyncClient(this.options.getBaseUrl(), this.options.getApiKey(), @@ -187,6 +189,8 @@ public ImageResponse call(ImagePrompt imagePrompt) { ImageGenerateParams imageGenerateParams = options.toOpenAiImageGenerateParams(imagePrompt); + RequestOptions requestOptions = this.buildRequestOptions(options); + if (logger.isTraceEnabled()) { logger.trace("OpenAiImageOptions call " + options.getModel() + " with the following options : " + imageGenerateParams); @@ -202,7 +206,7 @@ public ImageResponse call(ImagePrompt imagePrompt) { .observation(this.observationConvention, DEFAULT_OBSERVATION_CONVENTION, () -> observationContext, this.observationRegistry) .observe(() -> { - var images = this.openAiClient.images().generate(imageGenerateParams); + var images = this.openAiClient.images().generate(imageGenerateParams, requestOptions); if (images.data().isEmpty() && images.data().get().isEmpty()) { throw new IllegalArgumentException("Image generation failed: no image returned"); @@ -239,6 +243,20 @@ public void setObservationConvention(ImageModelObservationConvention observation this.observationConvention = observationConvention; } + /** + * Creates a RequestOptions instance from the given image options. + * @param options the image options + * @return a RequestOptions instance + */ + private RequestOptions buildRequestOptions(OpenAiImageOptions options) { + Assert.notNull(options, "Options cannot be null"); + RequestOptions.Builder requestOptionsBuilder = RequestOptions.builder(); + if (options.getTimeout() != null) { + requestOptionsBuilder.timeout(options.getTimeout()); + } + return requestOptionsBuilder.build(); + } + public static final class Builder { private @Nullable OpenAIClient openAiClient; diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModerationModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModerationModel.java index f84396108e..965b2675d4 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModerationModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModerationModel.java @@ -18,8 +18,10 @@ import java.util.ArrayList; import java.util.List; +import java.util.Objects; import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; import com.openai.models.moderations.ModerationCreateParams; import com.openai.models.moderations.ModerationCreateResponse; import io.micrometer.observation.ObservationRegistry; @@ -32,11 +34,11 @@ import org.springframework.ai.moderation.Generation; import org.springframework.ai.moderation.Moderation; import org.springframework.ai.moderation.ModerationModel; -import org.springframework.ai.moderation.ModerationOptions; import org.springframework.ai.moderation.ModerationPrompt; import org.springframework.ai.moderation.ModerationResponse; import org.springframework.ai.moderation.ModerationResult; import org.springframework.ai.openai.http.okhttp.OpenAiHttpClientBuilderCustomizer; +import org.springframework.ai.openai.setup.OpenAiSetup; import org.springframework.util.Assert; /** @@ -49,6 +51,7 @@ * @author Ilayaperumal Gopinathan * @author Sebastien Deleuze * @author Thomas Vitale + * @author guan xu */ public final class OpenAiModerationModel implements ModerationModel { @@ -59,23 +62,19 @@ public final class OpenAiModerationModel implements ModerationModel { private final OpenAiModerationOptions options; private OpenAiModerationModel(Builder builder) { - if (builder.options == null) { - this.options = OpenAiModerationOptions.builder() - .model(OpenAiModerationOptions.DEFAULT_MODERATION_MODEL) - .build(); - } - else { - this.options = builder.options; - } - - this.openAiClient = java.util.Objects.requireNonNullElseGet(builder.openAiClient, - () -> org.springframework.ai.openai.setup.OpenAiSetup.setupSyncClient(this.options.getBaseUrl(), - this.options.getApiKey(), this.options.getCredential(), - this.options.getMicrosoftDeploymentName(), this.options.getMicrosoftFoundryServiceVersion(), - this.options.getOrganizationId(), this.options.isMicrosoftFoundry(), - this.options.isGitHubModels(), this.options.getModel(), this.options.getTimeout(), - this.options.getMaxRetries(), this.options.getProxy(), this.options.getCustomHeaders(), - ObservationRegistry.NOOP, null, builder.httpClientCustomizers)); + this.options = Objects.requireNonNullElseGet(builder.options, + () -> OpenAiModerationOptions.builder() + .model(OpenAiModerationOptions.DEFAULT_MODERATION_MODEL) + .build()); + + this.openAiClient = Objects.requireNonNullElseGet(builder.openAiClient, + () -> OpenAiSetup.setupSyncClient(this.options.getBaseUrl(), this.options.getApiKey(), + this.options.getCredential(), this.options.getMicrosoftDeploymentName(), + this.options.getMicrosoftFoundryServiceVersion(), this.options.getOrganizationId(), + this.options.isMicrosoftFoundry(), this.options.isGitHubModels(), this.options.getModel(), + this.options.getTimeout(), this.options.getMaxRetries(), this.options.getProxy(), + this.options.getCustomHeaders(), ObservationRegistry.NOOP, null, + builder.httpClientCustomizers)); } public static Builder builder() { @@ -90,7 +89,11 @@ public Builder mutate() { public ModerationResponse call(ModerationPrompt moderationPrompt) { String text = moderationPrompt.getInstructions().getText(); - OpenAiModerationOptions options = merge(moderationPrompt.getOptions(), this.options); + // Merge request options with default options + OpenAiModerationOptions options = OpenAiModerationOptions.builder() + .from(this.options) + .merge(moderationPrompt.getOptions()) + .build(); ModerationCreateParams.Builder builder = ModerationCreateParams.builder() .input(ModerationCreateParams.Input.ofString(text)); @@ -107,11 +110,27 @@ public ModerationResponse call(ModerationPrompt moderationPrompt) { ModerationCreateParams params = builder.build(); - ModerationCreateResponse response = this.openAiClient.moderations().create(params); + RequestOptions requestOptions = this.buildRequestOptions(options); + + ModerationCreateResponse response = this.openAiClient.moderations().create(params, requestOptions); return convertResponse(response); } + /** + * Creates a RequestOptions instance from the given moderation options. + * @param options the moderation options + * @return a RequestOptions instance + */ + private RequestOptions buildRequestOptions(OpenAiModerationOptions options) { + Assert.notNull(options, "Options cannot be null"); + RequestOptions.Builder requestOptionsBuilder = RequestOptions.builder(); + if (options.getTimeout() != null) { + requestOptionsBuilder.timeout(options.getTimeout()); + } + return requestOptionsBuilder.build(); + } + private ModerationResponse convertResponse(ModerationCreateResponse response) { if (response == null) { logger.warn("No moderation response returned"); @@ -167,10 +186,6 @@ private ModerationResponse convertResponse(ModerationCreateResponse response) { return new ModerationResponse(new Generation(moderation)); } - private static OpenAiModerationOptions merge(@Nullable ModerationOptions source, OpenAiModerationOptions target) { - return OpenAiModerationOptions.builder().from(target).merge(source).build(); - } - public OpenAiModerationOptions getOptions() { return this.options; } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiChatModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiChatModelTests.java index fbb6dcc6ba..ceac721df8 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiChatModelTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiChatModelTests.java @@ -17,6 +17,7 @@ package org.springframework.ai.openai; import java.nio.charset.StandardCharsets; +import java.time.Duration; import java.util.HashMap; import java.util.LinkedHashMap; import java.util.List; @@ -31,6 +32,7 @@ import com.openai.client.OpenAIClient; import com.openai.client.OpenAIClientAsync; import com.openai.core.JsonValue; +import com.openai.core.RequestOptions; import com.openai.core.http.AsyncStreamResponse; import com.openai.models.FunctionDefinition; import com.openai.models.FunctionParameters; @@ -54,6 +56,7 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.ValueSource; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import reactor.core.publisher.Flux; @@ -77,6 +80,7 @@ import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; /** @@ -110,25 +114,26 @@ void preserveUnmappedRootResponseMetadata() { ChatCompletionService chatCompletionService = mock(ChatCompletionService.class); when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("gen-1888888888-XYZabc123NewId") - .created(1777799928) - .model("moonshotai/kimi-k2.5-0127") - .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.STOP) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("hello") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List.of()) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("gen-1888888888-XYZabc123NewId") + .created(1777799928) + .model("moonshotai/kimi-k2.5-0127") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .build()) .build()) - .build()) - .additionalProperties(additionalProperties) - .build()); + .additionalProperties(additionalProperties) + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("test-model").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -155,33 +160,34 @@ void preserveUnmappedToolCallAdditionalProperties() { ChatCompletionService chatCompletionService = mock(ChatCompletionService.class); when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("chatcmpl-test") - .created(1777799928) - .model("gemini-3.5-flash") - .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.TOOL_CALLS) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List - .of(ChatCompletionMessageToolCall.ofFunction(ChatCompletionMessageFunctionToolCall.builder() - .id("call_1") - .function(ChatCompletionMessageFunctionToolCall.Function.builder() - .name("get_current_weather") - .arguments("{\"location\":\"Seoul\"}") - .build()) - .putAdditionalProperty("extra_content", - JsonValue.from(Map.of("google", Map.of("thought_signature", "signature-123")))) - .build()))) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("chatcmpl-test") + .created(1777799928) + .model("gemini-3.5-flash") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.TOOL_CALLS) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List + .of(ChatCompletionMessageToolCall.ofFunction(ChatCompletionMessageFunctionToolCall.builder() + .id("call_1") + .function(ChatCompletionMessageFunctionToolCall.Function.builder() + .name("get_current_weather") + .arguments("{\"location\":\"Seoul\"}") + .build()) + .putAdditionalProperty("extra_content", + JsonValue.from(Map.of("google", Map.of("thought_signature", "signature-123")))) + .build()))) + .build()) .build()) - .build()) - .build()); + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("gemini-3.5-flash").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -208,7 +214,8 @@ void preserveUnmappedToolCallAdditionalPropertiesFromStream() { ChatCompletionServiceAsync chatCompletionServiceAsync = mock(ChatCompletionServiceAsync.class); when(this.openAiClientAsync.chat()).thenReturn(chatServiceAsync); when(chatServiceAsync.completions()).thenReturn(chatCompletionServiceAsync); - when(chatCompletionServiceAsync.createStreaming(any(ChatCompletionCreateParams.class))) + when(chatCompletionServiceAsync.createStreaming(any(ChatCompletionCreateParams.class), + any(RequestOptions.class))) .thenReturn(asyncStreamResponse( ChatCompletionChunk.builder() .id("chatcmpl-stream-test") @@ -286,7 +293,8 @@ void mergeStreamToolCallWhenIdNameAndArgumentsArriveInSeparateChunks() { ChatCompletionServiceAsync chatCompletionServiceAsync = mock(ChatCompletionServiceAsync.class); when(this.openAiClientAsync.chat()).thenReturn(chatServiceAsync); when(chatServiceAsync.completions()).thenReturn(chatCompletionServiceAsync); - when(chatCompletionServiceAsync.createStreaming(any(ChatCompletionCreateParams.class))) + when(chatCompletionServiceAsync.createStreaming(any(ChatCompletionCreateParams.class), + any(RequestOptions.class))) .thenReturn(asyncStreamResponse(ChatCompletionChunk.builder() .id("chatcmpl-stream-test") .created(1777799928) @@ -591,25 +599,26 @@ void reasoningContentFromReasoningContentProperty() { when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("gen-1888888888-XYZabc123NewId") - .created(1777799928) - .model("deepseek-reasoner") - .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.STOP) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("hello") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List.of()) - .putAdditionalProperty("reasoning_content", JsonValue.from("Test reasoning content")) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("gen-1888888888-XYZabc123NewId") + .created(1777799928) + .model("deepseek-reasoner") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .putAdditionalProperty("reasoning_content", JsonValue.from("Test reasoning content")) + .build()) .build()) - .build()) - .build()); + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("deepseek-reasoner").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -631,25 +640,26 @@ void reasoningContentFromReasoningProperty() { when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("gen-1888888888-XYZabc123NewId") - .created(1777799928) - .model("test-reasoner") - .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.STOP) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("hello") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List.of()) - .putAdditionalProperty("reasoning", JsonValue.from("Test reasoning content")) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("gen-1888888888-XYZabc123NewId") + .created(1777799928) + .model("test-reasoner") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .putAdditionalProperty("reasoning", JsonValue.from("Test reasoning content")) + .build()) .build()) - .build()) - .build()); + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("test-reasoner").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -671,24 +681,25 @@ void reasoningContentEmptyWhenNeitherPropertyPresent() { when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("gen-1888888888-XYZabc123NewId") - .created(1777799928) - .model("test-model") - .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.STOP) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("hello") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List.of()) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("gen-1888888888-XYZabc123NewId") + .created(1777799928) + .model("test-model") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .build()) .build()) - .build()) - .build()); + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("test-model").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -708,24 +719,25 @@ void createdFieldPassedThroughInMetadata() { ChatCompletionService chatCompletionService = mock(ChatCompletionService.class); when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("test-id") - .created(1234567890L) - .model("test-model") - .usage(CompletionUsage.builder().promptTokens(10).completionTokens(20).totalTokens(30).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.STOP) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("hello") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List.of()) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("test-id") + .created(1234567890L) + .model("test-model") + .usage(CompletionUsage.builder().promptTokens(10).completionTokens(20).totalTokens(30).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .build()) .build()) - .build()) - .build()); + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("test-model").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -885,7 +897,8 @@ private Flux streamResponses(List chunks) { ChatCompletionServiceAsync chatCompletionServiceAsync = mock(ChatCompletionServiceAsync.class); when(this.openAiClientAsync.chat()).thenReturn(chatServiceAsync); when(chatServiceAsync.completions()).thenReturn(chatCompletionServiceAsync); - when(chatCompletionServiceAsync.createStreaming(any(ChatCompletionCreateParams.class))) + when(chatCompletionServiceAsync.createStreaming(any(ChatCompletionCreateParams.class), + any(RequestOptions.class))) .thenReturn(asyncStreamResponseOf(chunks)); OpenAiChatOptions options = OpenAiChatOptions.builder().model("deepseek-reasoner").build(); @@ -1062,24 +1075,25 @@ void metadataDoesNotContainOptionalValues() { ChatCompletionService chatCompletionService = mock(ChatCompletionService.class); when(this.openAiClient.chat()).thenReturn(chatService); when(chatService.completions()).thenReturn(chatCompletionService); - when(chatCompletionService.create(any(ChatCompletionCreateParams.class))).thenReturn(ChatCompletion.builder() - .id("test-id") - .created(1777799928) - .model("test-model") - .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) - .addChoice(ChatCompletion.Choice.builder() - .finishReason(ChatCompletion.Choice.FinishReason.STOP) - .index(0) - .logprobs(Optional.empty()) - .message(ChatCompletionMessage.builder() - .content("hello") - .refusal(Optional.empty()) - .role(JsonValue.from("assistant")) - .annotations(List.of()) - .toolCalls(List.of()) + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("test-id") + .created(1777799928) + .model("test-model") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .build()) .build()) - .build()) - .build()); + .build()); OpenAiChatOptions options = OpenAiChatOptions.builder().model("test-model").build(); OpenAiChatModel chatModel = OpenAiChatModel.builder() @@ -1399,4 +1413,48 @@ void intermediateStreamingChunksHaveNullFinishReason() { assertThat(responses.get(2).getResult().getMetadata().getFinishReason()).isEqualTo("STOP"); } + @Test + void testPropagatesTimeoutFromRequestOptions() { + Duration expectedTimeout = Duration.ofSeconds(30); + + ChatService chatService = mock(ChatService.class); + ChatCompletionService chatCompletionService = mock(ChatCompletionService.class); + when(this.openAiClient.chat()).thenReturn(chatService); + when(chatService.completions()).thenReturn(chatCompletionService); + when(chatCompletionService.create(any(ChatCompletionCreateParams.class), any(RequestOptions.class))) + .thenReturn(ChatCompletion.builder() + .id("test-id") + .created(1777799928) + .model("test-model") + .usage(CompletionUsage.builder().promptTokens(1).completionTokens(1).totalTokens(2).build()) + .addChoice(ChatCompletion.Choice.builder() + .finishReason(ChatCompletion.Choice.FinishReason.STOP) + .index(0) + .logprobs(Optional.empty()) + .message(ChatCompletionMessage.builder() + .content("hello") + .refusal(Optional.empty()) + .role(JsonValue.from("assistant")) + .annotations(List.of()) + .toolCalls(List.of()) + .build()) + .build()) + .build()); + + OpenAiChatOptions options = OpenAiChatOptions.builder().model("test-model").timeout(expectedTimeout).build(); + OpenAiChatModel chatModel = OpenAiChatModel.builder() + .openAiClient(this.openAiClient) + .openAiClientAsync(this.openAiClientAsync) + .options(options) + .build(); + + chatModel.call(new Prompt("hi", options)); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(RequestOptions.class); + verify(chatCompletionService).create(any(ChatCompletionCreateParams.class), argumentCaptor.capture()); + RequestOptions value = argumentCaptor.getValue(); + assertThat(value.getTimeout()).isNotNull(); + assertThat(value.getTimeout().request()).isEqualTo(expectedTimeout); + } + } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/OpenAiAudioSpeechModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/OpenAiAudioSpeechModelTests.java index 25aa828d52..829e180765 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/OpenAiAudioSpeechModelTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/OpenAiAudioSpeechModelTests.java @@ -16,24 +16,38 @@ package org.springframework.ai.openai.audio; +import java.io.ByteArrayInputStream; +import java.time.Duration; + import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; +import com.openai.core.http.HttpResponse; +import com.openai.models.audio.speech.SpeechCreateParams; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.ai.audio.tts.TextToSpeechOptions; +import org.springframework.ai.audio.tts.TextToSpeechPrompt; import org.springframework.ai.openai.OpenAiAudioSpeechModel; import org.springframework.ai.openai.OpenAiAudioSpeechOptions; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; /** * Unit tests for OpenAiAudioSpeechModel. * * @author Ilayaperumal Gopinathan * @author Sebastien Deleuze + * @author guan xu */ @ExtendWith(MockitoExtension.class) class OpenAiAudioSpeechModelTests { @@ -101,7 +115,7 @@ void testConstructorWithAllParameters() { @Test void testOptions() { OpenAiAudioSpeechModel model = OpenAiAudioSpeechModel.builder().openAiClient(this.mockClient).build(); - OpenAiAudioSpeechOptions options = (OpenAiAudioSpeechOptions) model.getOptions(); + OpenAiAudioSpeechOptions options = model.getOptions(); assertThat(options.getModel()).isEqualTo("gpt-4o-mini-tts"); assertThat(options.getVoice()).isEqualTo("alloy"); @@ -210,7 +224,7 @@ void testOptionsMerging() { .build(); // Verify that default options are set - OpenAiAudioSpeechOptions defaults = (OpenAiAudioSpeechOptions) model.getOptions(); + OpenAiAudioSpeechOptions defaults = model.getOptions(); assertThat(defaults.getModel()).isEqualTo("tts-1"); assertThat(defaults.getVoice()).isEqualTo("alloy"); assertThat(defaults.getSpeed()).isEqualTo(1.0); @@ -242,7 +256,7 @@ void testBuilderWithDefaults() { assertThat(model.getOptions()).isNotNull(); assertThat(model.getOptions()).isInstanceOf(OpenAiAudioSpeechOptions.class); - OpenAiAudioSpeechOptions defaults = (OpenAiAudioSpeechOptions) model.getOptions(); + OpenAiAudioSpeechOptions defaults = model.getOptions(); assertThat(defaults.getModel()).isEqualTo("gpt-4o-mini-tts"); assertThat(defaults.getVoice()).isEqualTo("alloy"); assertThat(defaults.getResponseFormat()).isEqualTo("mp3"); @@ -285,8 +299,31 @@ void testBuilderWithPartialOptions() { OpenAiAudioSpeechModel model = OpenAiAudioSpeechModel.builder().openAiClient(this.mockClient).build(); assertThat(model).isNotNull(); - OpenAiAudioSpeechOptions defaults = (OpenAiAudioSpeechOptions) model.getOptions(); + OpenAiAudioSpeechOptions defaults = model.getOptions(); assertThat(defaults.getModel()).isEqualTo("gpt-4o-mini-tts"); } + @Test + void testPropagatesTimeoutFromRequestOptions() { + Duration expectedTimeout = Duration.ofSeconds(30); + + OpenAIClient mockClient = mock(OpenAIClient.class, RETURNS_DEEP_STUBS); + HttpResponse mockResponse = mock(HttpResponse.class); + when(mockResponse.body()).thenReturn(new ByteArrayInputStream(new byte[0])); + when(mockClient.audio().speech().create(any(SpeechCreateParams.class), any(RequestOptions.class))) + .thenReturn(mockResponse); + + OpenAiAudioSpeechModel model = OpenAiAudioSpeechModel.builder().openAiClient(mockClient).build(); + + OpenAiAudioSpeechOptions options = OpenAiAudioSpeechOptions.builder().timeout(expectedTimeout).build(); + + model.call(new TextToSpeechPrompt("hi", options)); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(RequestOptions.class); + verify(mockClient.audio().speech()).create(any(SpeechCreateParams.class), argumentCaptor.capture()); + RequestOptions value = argumentCaptor.getValue(); + assertThat(value.getTimeout()).isNotNull(); + assertThat(value.getTimeout().request()).isEqualTo(expectedTimeout); + } + } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java index 97d3114c13..23124ec586 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java @@ -43,7 +43,7 @@ class TranscriptionModelTests { @Test - void transcrbeRequestReturnsResponseCorrectly() { + void transcribeRequestReturnsResponseCorrectly() { Resource mockAudioFile = Mockito.mock(Resource.class); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingModelTests.java new file mode 100644 index 0000000000..7653a6f51a --- /dev/null +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingModelTests.java @@ -0,0 +1,71 @@ +/* + * Copyright 2023-present 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.embedding; + +import java.time.Duration; +import java.util.List; + +import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; +import com.openai.models.embeddings.CreateEmbeddingResponse; +import com.openai.models.embeddings.EmbeddingCreateParams; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; + +import org.springframework.ai.embedding.EmbeddingRequest; +import org.springframework.ai.openai.OpenAiEmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingOptions; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +/** + * Unit tests for {@link OpenAiEmbeddingModel}. + * + * @author guan xu + */ +class OpenAiEmbeddingModelTests { + + @Test + void testPropagatesTimeoutFromRequestOptions() { + Duration expectedTimeout = Duration.ofSeconds(30); + + OpenAIClient mockClient = mock(OpenAIClient.class, RETURNS_DEEP_STUBS); + CreateEmbeddingResponse mockResponse = mock(CreateEmbeddingResponse.class); + when(mockResponse.data()).thenReturn(List.of()); + when(mockResponse.usage()).thenReturn(mock(CreateEmbeddingResponse.Usage.class)); + when(mockClient.embeddings().create(any(EmbeddingCreateParams.class), any(RequestOptions.class))) + .thenReturn(mockResponse); + + OpenAiEmbeddingModel model = OpenAiEmbeddingModel.builder().openAiClient(mockClient).build(); + + OpenAiEmbeddingOptions options = OpenAiEmbeddingOptions.builder().timeout(expectedTimeout).build(); + + model.call(new EmbeddingRequest(List.of("hi"), options)); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(RequestOptions.class); + verify(mockClient.embeddings()).create(any(EmbeddingCreateParams.class), argumentCaptor.capture()); + RequestOptions value = argumentCaptor.getValue(); + assertThat(value.getTimeout()).isNotNull(); + assertThat(value.getTimeout().request()).isEqualTo(expectedTimeout); + } + +} diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelTests.java new file mode 100644 index 0000000000..d1672e570a --- /dev/null +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelTests.java @@ -0,0 +1,73 @@ +/* + * Copyright 2023-present 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.image; + +import java.time.Duration; +import java.util.List; +import java.util.Optional; + +import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; +import com.openai.models.images.Image; +import com.openai.models.images.ImageGenerateParams; +import com.openai.models.images.ImagesResponse; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; + +import org.springframework.ai.image.ImagePrompt; +import org.springframework.ai.openai.OpenAiImageModel; +import org.springframework.ai.openai.OpenAiImageOptions; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +/** + * Unit tests for {@link OpenAiImageModel}. + * + * @author guan xu + */ +class OpenAiImageModelTests { + + @Test + void testPropagatesTimeoutFromRequestOptions() { + Duration expectedTimeout = Duration.ofSeconds(30); + + OpenAIClient mockClient = mock(OpenAIClient.class, RETURNS_DEEP_STUBS); + ImagesResponse mockResponse = mock(ImagesResponse.class); + when(mockResponse.data()) + .thenReturn(Optional.of(List.of(Image.builder().url("https://example.com/image.png").build()))); + when(mockClient.images().generate(any(ImageGenerateParams.class), any(RequestOptions.class))) + .thenReturn(mockResponse); + + OpenAiImageModel model = OpenAiImageModel.builder().openAiClient(mockClient).build(); + + OpenAiImageOptions options = OpenAiImageOptions.builder().timeout(expectedTimeout).build(); + + model.call(new ImagePrompt("A small dog", options)); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(RequestOptions.class); + verify(mockClient.images()).generate(any(ImageGenerateParams.class), argumentCaptor.capture()); + RequestOptions value = argumentCaptor.getValue(); + assertThat(value.getTimeout()).isNotNull(); + assertThat(value.getTimeout().request()).isEqualTo(expectedTimeout); + } + +} diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/moderation/OpenAiModerationModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/moderation/OpenAiModerationModelTests.java index f33f38d707..970f043266 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/moderation/OpenAiModerationModelTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/moderation/OpenAiModerationModelTests.java @@ -17,19 +17,30 @@ package org.springframework.ai.openai.moderation; import java.time.Duration; +import java.util.List; import java.util.Map; import com.openai.client.OpenAIClient; +import com.openai.core.RequestOptions; +import com.openai.models.moderations.ModerationCreateParams; +import com.openai.models.moderations.ModerationCreateResponse; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.ai.moderation.ModerationOptions; +import org.springframework.ai.moderation.ModerationPrompt; import org.springframework.ai.openai.OpenAiModerationModel; import org.springframework.ai.openai.OpenAiModerationOptions; import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; /** * Unit tests for OpenAiModerationModel. @@ -236,4 +247,25 @@ void testOptionsBuilderMergeCustomHeaders() { .containsEntry("merged-header2", "merged-value2"); } + @Test + void testPropagatesTimeoutFromRequestOptions() { + Duration expectedTimeout = Duration.ofSeconds(30); + + OpenAIClient mockClient = mock(OpenAIClient.class, RETURNS_DEEP_STUBS); + when(mockClient.moderations().create(any(ModerationCreateParams.class), any(RequestOptions.class))).thenReturn( + ModerationCreateResponse.builder().id("TEST_ID").model("TEST_MODEL").results(List.of()).build()); + + OpenAiModerationModel model = OpenAiModerationModel.builder().openAiClient(mockClient).build(); + + OpenAiModerationOptions options = OpenAiModerationOptions.builder().timeout(expectedTimeout).build(); + + model.call(new ModerationPrompt("hi", options)); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(RequestOptions.class); + verify(mockClient.moderations()).create(any(ModerationCreateParams.class), argumentCaptor.capture()); + RequestOptions value = argumentCaptor.getValue(); + assertThat(value.getTimeout()).isNotNull(); + assertThat(value.getTimeout().request()).isEqualTo(expectedTimeout); + } + } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transcription/OpenAiAudioTranscriptionModelTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transcription/OpenAiAudioTranscriptionModelTests.java index b6aa5d6bc7..b1a044c429 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transcription/OpenAiAudioTranscriptionModelTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transcription/OpenAiAudioTranscriptionModelTests.java @@ -16,6 +16,7 @@ package org.springframework.ai.openai.transcription; +import java.time.Duration; import java.util.List; import java.util.Optional; import java.util.concurrent.CompletableFuture; @@ -23,9 +24,11 @@ import com.openai.client.OpenAIClient; import com.openai.client.OpenAIClientAsync; +import com.openai.core.RequestOptions; import com.openai.core.http.AsyncStreamResponse; import com.openai.models.audio.AudioResponseFormat; import com.openai.models.audio.transcriptions.Transcription; +import com.openai.models.audio.transcriptions.TranscriptionCreateParams; import com.openai.models.audio.transcriptions.TranscriptionCreateResponse; import com.openai.models.audio.transcriptions.TranscriptionStreamEvent; import com.openai.models.audio.transcriptions.TranscriptionTextDeltaEvent; @@ -34,16 +37,20 @@ import com.openai.services.blocking.AudioService; import com.openai.services.blocking.audio.TranscriptionService; import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; import org.springframework.ai.audio.transcription.AudioTranscriptionPrompt; import org.springframework.ai.audio.transcription.AudioTranscriptionResponse; import org.springframework.ai.openai.OpenAiAudioTranscriptionModel; import org.springframework.ai.openai.OpenAiAudioTranscriptionOptions; +import org.springframework.core.io.ByteArrayResource; import org.springframework.core.io.ClassPathResource; import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; /** @@ -402,13 +409,43 @@ void streamTranscribeWithResourceAndOptions() { assertThat(text).isEqualTo("Hello, streamed transcription result"); } + @Test + void testPropagatesTimeoutFromRequestOptions() { + Duration expectedTimeout = Duration.ofSeconds(30); + + OpenAIClient mockClient = mock(OpenAIClient.class, RETURNS_DEEP_STUBS); + TranscriptionCreateResponse mockResponse = mock(TranscriptionCreateResponse.class); + when(mockClient.audio() + .transcriptions() + .create(any(TranscriptionCreateParams.class), any(RequestOptions.class))).thenReturn(mockResponse); + + OpenAiAudioTranscriptionModel model = OpenAiAudioTranscriptionModel.builder() + .openAiClient(mockClient) + .openAiClientAsync(mock(OpenAIClientAsync.class)) + .build(); + + OpenAiAudioTranscriptionOptions options = OpenAiAudioTranscriptionOptions.builder() + .timeout(expectedTimeout) + .build(); + + model.call(new AudioTranscriptionPrompt(new ByteArrayResource(new byte[] { 1 }), options)); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(RequestOptions.class); + verify(mockClient.audio().transcriptions()).create(any(TranscriptionCreateParams.class), + argumentCaptor.capture()); + RequestOptions value = argumentCaptor.getValue(); + assertThat(value.getTimeout()).isNotNull(); + assertThat(value.getTimeout().request()).isEqualTo(expectedTimeout); + } + private OpenAIClient createMockClient(TranscriptionCreateResponse mockResponse) { OpenAIClient client = mock(OpenAIClient.class); AudioService audioService = mock(AudioService.class); TranscriptionService transcriptionService = mock(TranscriptionService.class); when(client.audio()).thenReturn(audioService); when(audioService.transcriptions()).thenReturn(transcriptionService); - when(transcriptionService.create(any())).thenReturn(mockResponse); + when(transcriptionService.create(any(TranscriptionCreateParams.class), any(RequestOptions.class))) + .thenReturn(mockResponse); return client; } @@ -418,7 +455,8 @@ private OpenAIClientAsync createMockAsyncClient(AsyncStreamResponse streamTranscribe(Resource resource) { */ default Flux streamTranscribe(Resource resource, @Nullable AudioTranscriptionOptions options) { AudioTranscriptionPrompt prompt = new AudioTranscriptionPrompt(resource, options); - return stream(prompt) + return this.stream(prompt) .map(response -> Optional.ofNullable(response.getResult()).map(AudioTranscription::getOutput).orElse("")); }