Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -52,6 +54,7 @@
* @author Jonghoon Park
* @author Ilayaperumal Gopinathan
* @author Sebastien Deleuze
* @author guan xu
*/
public final class OpenAiAudioSpeechModel implements TextToSpeechModel {

Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -170,6 +175,20 @@ public Flux<TextToSpeechResponse> 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
*/
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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());
Expand All @@ -146,14 +152,16 @@ public Flux<AudioTranscriptionResponse> 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<TranscriptionStreamEvent> chunks = Flux.create(sink -> this.openAiClientAsync.audio()
.transcriptions()
.createStreaming(params)
.createStreaming(params, requestOptions)
.subscribe(sink::next)
.onCompleteFuture()
.whenComplete((unused, throwable) -> {
Expand Down Expand Up @@ -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();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -123,6 +124,7 @@
* @author Eric Bottard
* @author Taewoong Kim
* @author Jewoo Shin
* @author guan xu
*/
public final class OpenAiChatModel implements ChatModel {

Expand Down Expand Up @@ -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)
Expand All @@ -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<ChatCompletion.Choice> choices = chatCompletion.choices();
if (choices.isEmpty()) {
Expand Down Expand Up @@ -268,7 +271,8 @@ public Flux<ChatResponse> stream(Prompt prompt) {
*/
private Flux<ChatResponse> internalStream(Prompt prompt) {
return Flux.deferContextual(contextView -> {
ChatCompletionCreateParams request = createRequest(prompt, true);
ChatCompletionCreateParams request = this.createRequest(prompt, true);
RequestOptions requestOptions = this.buildRequestOptions(prompt);
ConcurrentHashMap<String, String> roleMap = new ConcurrentHashMap<>();
ConcurrentHashMap<String, String> reasoningMap = new ConcurrentHashMap<>();
final ChatModelObservationContext observationContext = ChatModelObservationContext.builder()
Expand All @@ -293,9 +297,9 @@ private Flux<ChatResponse> internalStream(Prompt prompt) {
}

// Convert from AsyncStreamResponse<ChatCompletionChunk> to Flux<CCC>
Flux<ChatCompletionChunk> chunks = Flux.<ChatCompletionChunk>create(sink -> this.openAiClientAsync.chat()
Flux<ChatCompletionChunk> chunks = Flux.create(sink -> this.openAiClientAsync.chat()
.completions()
.createStreaming(request)
.createStreaming(request, requestOptions)
.subscribe(sink::next)
.onCompleteFuture()
.whenComplete((unused, throwable) -> {
Expand Down Expand Up @@ -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<String, String> toolCallAdditionalPropertiesFromMetadata(AssistantMessage assistantMessage) {
Object value = assistantMessage.getMetadata().get(TOOL_CALL_ADDITIONAL_PROPERTIES_METADATA_KEY);
if (!(value instanceof Map<?, ?> rawMap)) {
Expand Down Expand Up @@ -1476,8 +1497,8 @@ public Builder httpClientBuilderCustomizers(List<OpenAiHttpClientBuilderCustomiz
* @return the configured chat model
*/
public OpenAiChatModel build() {
OpenAiChatOptions resolvedOptions = this.options != null ? this.options
: OpenAiChatOptions.builder().build();
OpenAiChatOptions resolvedOptions = Objects.requireNonNullElseGet(this.options,
() -> OpenAiChatOptions.builder().build());
ObservationRegistry resolvedObservationRegistry = Objects.requireNonNullElse(this.observationRegistry,
ObservationRegistry.NOOP);

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -55,6 +56,7 @@
* @author Thomas Vitale
* @author Christian Tzolov
* @author Josh Long
* @author guan xu
*/
public class OpenAiEmbeddingModel extends AbstractEmbeddingModel {

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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())
Expand All @@ -243,14 +247,29 @@ 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);
return embeddingResponse;
}));
}

/**
* 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<Embedding> data = generateEmbeddingList(response.data());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -52,6 +53,7 @@
* @author Hyunjoon Choi
* @author Christian Tzolov
* @author Mark Pollack
* @author guan xu
*/
public class OpenAiImageModel implements ImageModel {

Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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);
Expand All @@ -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");
Expand Down Expand Up @@ -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;
Expand Down
Loading
Loading