diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/AudioInputChannel.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/AudioInputChannel.java new file mode 100644 index 000000000..7d0f822bd --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/AudioInputChannel.java @@ -0,0 +1,22 @@ +package com.sap.ai.sdk.foundationmodels.openai; + +import com.google.common.annotations.Beta; + +/** + * Functional interface representing audio input channel (audio data consumer) + * + *

Should be closed by application (try-with-resources) when not needed anymore + */ +@Beta +public interface AudioInputChannel extends AutoCloseable { + + /** + * This method is sequentially invoked by audio data provider to supply implementer (consumer) + * with the audio data. Exact audio format (encoding, sampling rate, etc.) depends on the usage + * context + * + * @param rawBytesChunk binary data in the depending on the use case format + */ + @Beta + void inputAudio(byte[] rawBytesChunk); +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/AudioOutputChannel.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/AudioOutputChannel.java new file mode 100644 index 000000000..7c15fb759 --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/AudioOutputChannel.java @@ -0,0 +1,21 @@ +package com.sap.ai.sdk.foundationmodels.openai; + +import com.google.common.annotations.Beta; + +/** Functional interface representing audio output channel (audio data consumer) */ +@Beta +public interface AudioOutputChannel { + + /** + * This method is sequentially invoked by audio data provider to supply implementer (consumer) + * with the audio data. Exact audio format (encoding, sampling rate, etc.) depends on the usage + * context + * + * @param rawBytesChunk binary data in the depending on the use case format + * @param isLast true if this call logically concludes previous and this passed bytes data into a + * single logical entity (e.g. gets called at the end when all byte parts of a single message + * get passed) + */ + @Beta + void outputAudio(byte[] rawBytesChunk, boolean isLast); +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiClient.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiClient.java index a76c3d89f..bd3b4176a 100644 --- a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiClient.java +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiClient.java @@ -26,6 +26,7 @@ import com.sap.ai.sdk.foundationmodels.openai.model.OpenAiChatMessage.OpenAiChatUserMessage; import com.sap.ai.sdk.foundationmodels.openai.model.OpenAiEmbeddingOutput; import com.sap.ai.sdk.foundationmodels.openai.model.OpenAiEmbeddingParameters; +import com.sap.ai.sdk.foundationmodels.openai.realtime.OpenAiRealtimeClient; import com.sap.cloud.sdk.cloudplatform.connectivity.ApacheHttpClient5Accessor; import com.sap.cloud.sdk.cloudplatform.connectivity.DefaultHttpDestination; import com.sap.cloud.sdk.cloudplatform.connectivity.Destination; @@ -76,6 +77,18 @@ public static OpenAiClient forModel(@Nonnull final OpenAiModel foundationModel) return client.withApiVersion(DEFAULT_API_VERSION); } + /** + * Creates and configures OpenAI Realtime API client + * + * @return created client + */ + @Beta + @Nonnull + public static OpenAiRealtimeClient realtimeClient() { + final var withResolvedDestination = OpenAiClient.forModel(OpenAiModel.GPT_REALTIME); + return new OpenAiRealtimeClient(withResolvedDestination.destination); + } + /** * Create a new OpenAI client targeting the specified API version. * diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiModel.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiModel.java index 1926f56ce..62c435bcc 100644 --- a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiModel.java +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/OpenAiModel.java @@ -121,7 +121,7 @@ public record OpenAiModel(@Nonnull String name, @Nullable String version) implem /** Azure OpenAI GPT-5-nano model */ public static final OpenAiModel GPT_5_NANO = new OpenAiModel("gpt-5-nano", null); - /** Azure OpenAI GPT-5-nano model */ + /** Azure OpenAI GPT-realtime model */ public static final OpenAiModel GPT_REALTIME = new OpenAiModel("gpt-realtime", null); /** Azure OpenAI GPT-5.2 model */ diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/TextInputChannel.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/TextInputChannel.java new file mode 100644 index 000000000..bbd78c40a --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/TextInputChannel.java @@ -0,0 +1,20 @@ +package com.sap.ai.sdk.foundationmodels.openai; + +import com.google.common.annotations.Beta; +import javax.annotation.Nonnull; + +/** + * Allows to input (send) text to the open channel, must be closed when not needed anymore (e.g. + * try-with-resources) + */ +@Beta +public interface TextInputChannel extends AutoCloseable { + + /** + * Sends input text + * + * @param text text to send + */ + @Beta + void sendText(@Nonnull final String text); +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/BufferedWebSocketListener.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/BufferedWebSocketListener.java new file mode 100644 index 000000000..6e5ac52ba --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/BufferedWebSocketListener.java @@ -0,0 +1,60 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import java.net.http.WebSocket; +import java.nio.ByteBuffer; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionStage; +import java.util.function.BiConsumer; +import java.util.function.Consumer; +import javax.annotation.Nonnull; +import lombok.extern.slf4j.Slf4j; + +@Slf4j +class BufferedWebSocketListener implements WebSocket.Listener { + + private final Consumer onOpen; + private final BiConsumer onText; + private final StringBuilder buffer; + + BufferedWebSocketListener( + @Nonnull final Consumer onOpen, + @Nonnull final BiConsumer onText) { + this.onOpen = onOpen; + this.onText = onText; + this.buffer = new StringBuilder(256 * 1024); + } + + @Override + public void onOpen(@Nonnull final WebSocket webSocket) { + this.onOpen.accept(webSocket); + webSocket.request(1); + } + + @Override + @Nonnull + public CompletionStage onText( + final @Nonnull WebSocket webSocket, final @Nonnull CharSequence data, final boolean isLast) { + buffer.append(data); + webSocket.request(1); + if (isLast) { + final var completeMessage = buffer.toString(); + buffer.setLength(0); + this.onText.accept(webSocket, completeMessage); + } + + return CompletableFuture.completedStage(null); + } + + @Override + public void onError(@Nonnull final WebSocket webSocket, @Nonnull final Throwable error) { + log.error("Websocket error occurred during realtime communication", error); + } + + @Override + @Nonnull + public CompletionStage onBinary( + @Nonnull final WebSocket webSocket, @Nonnull final ByteBuffer data, final boolean isLast) { + log.warn("Received unexpected binary bytes for WebSocket connection"); + return CompletableFuture.completedStage(null); + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/OpenAiRealtimeClient.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/OpenAiRealtimeClient.java new file mode 100644 index 000000000..309384940 --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/OpenAiRealtimeClient.java @@ -0,0 +1,128 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.google.common.annotations.Beta; +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; +import com.sap.cloud.sdk.cloudplatform.connectivity.Destination; +import com.sap.cloud.sdk.cloudplatform.connectivity.Header; +import java.util.HashMap; +import java.util.Map; +import javax.annotation.Nonnull; + +/** + * OpenAI client implementation of Realtime API. Abstracts technical implementation, transport and + * threading and exposes business-level operations (high level interface) + */ +@Beta +public class OpenAiRealtimeClient { + + static final int PATH_BUFFER_SIZE = + 400; // existing URLs are ~120 symbols long, 400 has reasonable margin + + final Destination destination; + + /** + * Created OpenAI Realtime client for a specific destination + * + * @param destination - destination to use + */ + @Beta + public OpenAiRealtimeClient(@Nonnull final Destination destination) { + this.destination = destination; + } + + /** + * Creates a realtime channel allowing to input text and voice it (receive audio output) + * + *

The input channel should be used with a try-with-resources block to ensure that the + * underlying connection is closed. + * + *

Example: + * + *

{@code
+   * try (var textInputChannel = client.textToSpeech(audioOutputConsumer)) {
+   *       textInputChannel.sendText("...");
+   *       ....
+   * }
+   * }
+ * + * This API implements full duplex (input + output) communication channels. Application should + * logically synchronize their state and close the input channel when it is appropriate (e.g. the + * last part of the response has been received via the output channel and the application does not + * need to send any other input). When the input channel is closed, the output channel will be + * closed automatically and the output consumer will not be called anymore. + * + * @param audioOutputConsumer - audio consumer of raw PCM mono 24000 Hz little endian output, 16 + * bit depth + * @param params - allows for various additional features (e.g. voice configuration or + * conversation turn recognition options) + * @return input channel, allowing for text input + */ + @Nonnull + @Beta + public TextInputChannel textToSpeech( + @Nonnull final AudioOutputChannel audioOutputConsumer, + @Nonnull final RealtimeParam... params) { + return new TextToSpeechRealtimeClient( + getRealtimeEndpoint(), buildRealtimeHeaders(), audioOutputConsumer, params); + } + + /** + * Creates a realtime channel allowing for audio conversation with a model. + * + *

The input channel should be used with a try-with-resources block to ensure that the + * underlying connection is closed. + * + *

Example: + * + *

{@code
+   * try (var audioInputChannel = client.speechToSpeech(audioOutputConsumer)) {
+   *       audioInputChannel.inputAudio(audioBytesData);
+   *       ....
+   * }
+   * }
+ * + * This API implements full duplex (input + output) communication channels. An application should + * logically synchronize their state and close the input channel when it is appropriate (e.g. the + * last part of the response has been received via the output channel and the application does not + * need to send any other input). When the input channel is closed, the output channel will be + * closed automatically and the output consumer will not be called anymore. + * + * @param audioOutputConsumer - audio consumer of raw PCM mono 24000 Hz little endian output, 16 + * bit depth + * @param params - optional configuration params + * @return input channel, allowing for audio data input (bytes, PCM mono 24000 Hz little endian 16 + * bit) + */ + @Nonnull + @Beta + public AudioInputChannel speechToSpeech( + @Nonnull final AudioOutputChannel audioOutputConsumer, + @Nonnull final RealtimeParam... params) { + return new SpeechToSpeechRealtimeClient( + getRealtimeEndpoint(), buildRealtimeHeaders(), audioOutputConsumer, params); + } + + Map buildRealtimeHeaders() { + final var extraHeaders = destination.asHttp().getHeaders(); + final var headers = new HashMap(extraHeaders.size() + 1); + for (final Header header : extraHeaders) { + headers.put(header.getName(), header.getValue()); + } + return headers; + } + + String getRealtimeEndpoint() { + final var sb = new StringBuilder(PATH_BUFFER_SIZE); + sb.append("wss://"); + final var pathParts = destination.asHttp().getUri().toString().split("//"); + if (pathParts.length != 2) { + throw new IllegalArgumentException( + "Invalid destination URI: " + destination.asHttp().getUri()); + } + sb.append(pathParts[1].replaceFirst("^api\\.", "realtime.")); + sb.append("/v1/realtime"); + return sb.toString(); + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParam.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParam.java new file mode 100644 index 000000000..26214db7b --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParam.java @@ -0,0 +1,37 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.google.common.annotations.Beta; +import javax.annotation.Nonnull; + +/** Represents possible configuration params of realtime client */ +@Beta +public interface RealtimeParam { + /** Represents configurable options */ + enum ParamName { + /** Voice name to use to produce sound */ + OUTPUT_VOICE, + /** + * How model will recognize that it is its turn to respond (e.g. explicitly asked, automatically + * detected) + */ + TURN_DETECTION, + /** Override or specify system prompt given to a model */ + SYSTEM_PROMPT, + } + + /** + * Returns param name + * + * @return name + */ + @Nonnull + ParamName getParamName(); + + /** + * Returns string value representation of the param + * + * @return string value + */ + @Nonnull + String getValueAsString(); +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamSystemPrompt.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamSystemPrompt.java new file mode 100644 index 000000000..10ec75d3a --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamSystemPrompt.java @@ -0,0 +1,51 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.google.common.annotations.Beta; +import java.util.Objects; +import javax.annotation.Nonnull; +import javax.annotation.Nullable; + +/** Allows to configure model system prompt */ +@Beta +public final class RealtimeParamSystemPrompt implements RealtimeParam { + + private final String systemPrompt; + + /** + * Constructs RealtimeParamSystemPrompt object + * + * @param systemPrompt system prompt to use + */ + @Beta + public RealtimeParamSystemPrompt(@Nonnull final String systemPrompt) { + this.systemPrompt = systemPrompt; + } + + @Override + @Beta + public @Nonnull ParamName getParamName() { + return ParamName.SYSTEM_PROMPT; + } + + @Override + @Beta + public @Nonnull String getValueAsString() { + return systemPrompt; + } + + @Override + @Beta + public boolean equals(@Nullable final Object o) { + if (o == null || getClass() != o.getClass()) { + return false; + } + final RealtimeParamSystemPrompt that = (RealtimeParamSystemPrompt) o; + return Objects.equals(systemPrompt, that.systemPrompt); + } + + @Override + @Beta + public int hashCode() { + return Objects.hashCode(systemPrompt); + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamTurnDetection.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamTurnDetection.java new file mode 100644 index 000000000..cb087bca2 --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamTurnDetection.java @@ -0,0 +1,59 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.google.common.annotations.Beta; +import java.util.Objects; +import javax.annotation.Nonnull; +import javax.annotation.Nullable; + +/** Allows to configure turn detection (how model responds). */ +@Beta +public final class RealtimeParamTurnDetection implements RealtimeParam { + + /** Model tries to recognize if/when it should respond automatically */ + @Beta + public static final RealtimeParamTurnDetection BY_MODEL_AUTO = + new RealtimeParamTurnDetection("BY_MODEL_AUTO"); + + /** + * Each call to the provided realtime client is considered a turn (eager explicit turn detection). + * Less convenient than the automatic option but may give lower latency in some cases (model does + * not need to perform additional turn detection analysis). + */ + @Beta + public static final RealtimeParamTurnDetection EACH_CALL_IS_A_TURN = + new RealtimeParamTurnDetection("EACH_CALL_IS_A_TURN"); + + private final String turnDetectionKind; + + RealtimeParamTurnDetection(final String turnDetectionKind) { + this.turnDetectionKind = turnDetectionKind; + } + + @Override + @Beta + public @Nonnull ParamName getParamName() { + return ParamName.TURN_DETECTION; + } + + @Override + @Beta + public @Nonnull String getValueAsString() { + return turnDetectionKind; + } + + @Override + @Beta + public boolean equals(@Nullable final Object o) { + if (o == null || getClass() != o.getClass()) { + return false; + } + final RealtimeParamTurnDetection that = (RealtimeParamTurnDetection) o; + return Objects.equals(turnDetectionKind, that.turnDetectionKind); + } + + @Override + @Beta + public int hashCode() { + return Objects.hashCode(turnDetectionKind); + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamVoice.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamVoice.java new file mode 100644 index 000000000..e5ff23f0b --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamVoice.java @@ -0,0 +1,66 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.google.common.annotations.Beta; +import java.util.Objects; +import javax.annotation.Nonnull; +import javax.annotation.Nullable; + +/** Allows to configure model output voice */ +@Beta +public final class RealtimeParamVoice implements RealtimeParam { + + /** Standard voice 1 */ + @Beta public static final RealtimeParamVoice DEFAULT_1 = new RealtimeParamVoice("DEFAULT_1"); + + /** Standard voice 2 */ + @Beta public static final RealtimeParamVoice DEFAULT_2 = new RealtimeParamVoice("DEFAULT_2"); + + private final String voice; + + RealtimeParamVoice(@Nonnull final String voice) { + this.voice = voice; + } + + /** + * Allows to configure raw voice name as named by model provider. Unsafe because SDK cannot verify + * in advance if the provided voice name is correct and supported by the chosen model and use case + * NOTE: this method does not check voice name and incorrect input may produce runtime exceptions + * (unsafe) + * + * @param voiceName as named by model provider + * @return typed voice client configuration param + */ + @Nonnull + @Beta + public static RealtimeParamVoice withExplicitVoice(@Nonnull final String voiceName) { + return new RealtimeParamVoice(voiceName); + } + + @Override + @Beta + public @Nonnull ParamName getParamName() { + return ParamName.OUTPUT_VOICE; + } + + @Override + @Beta + public @Nonnull String getValueAsString() { + return voice; + } + + @Override + @Beta + public boolean equals(@Nullable final Object o) { + if (o == null || getClass() != o.getClass()) { + return false; + } + final RealtimeParamVoice that = (RealtimeParamVoice) o; + return Objects.equals(voice, that.voice); + } + + @Override + @Beta + public int hashCode() { + return Objects.hashCode(voice); + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/SpeechToSpeechRealtimeClient.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/SpeechToSpeechRealtimeClient.java new file mode 100644 index 000000000..3d5d40b93 --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/SpeechToSpeechRealtimeClient.java @@ -0,0 +1,87 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.openai.models.realtime.InputAudioBufferAppendEvent; +import com.openai.models.realtime.InputAudioBufferCommitEvent; +import com.openai.models.realtime.RealtimeAudioConfigInput; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.openai.models.realtime.RealtimeAudioInputTurnDetection; +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.util.Arrays; +import java.util.Base64; +import java.util.Map; +import java.util.Timer; +import java.util.concurrent.CompletableFuture; +import javax.annotation.Nonnull; +import lombok.extern.slf4j.Slf4j; + +@Slf4j +class SpeechToSpeechRealtimeClient extends ToAudioRealtimeClient implements AudioInputChannel { + + private static final int MAX_DATA_CHUNK_SIZE_BYTES = 8192; + + public SpeechToSpeechRealtimeClient( + @Nonnull final String url, + @Nonnull final Map httpHeaders, + @Nonnull final AudioOutputChannel outputConsumer, + @Nonnull final RealtimeParam... params) { + super(url, httpHeaders, outputConsumer, false, params); + } + + SpeechToSpeechRealtimeClient( + @Nonnull final HttpClient httpClient, + @Nonnull final CompletableFuture ws, + @Nonnull final Timer timer, + @Nonnull final AudioOutputChannel outputConsumer, + @Nonnull final RealtimeParam... params) { + super(httpClient, ws, timer, outputConsumer, false, params); + } + + @Override + @Nonnull + protected RealtimeAudioConfigInput inputConfig() { + RealtimeAudioInputTurnDetection turnDetection; + if (eagerTurnDetection) { + turnDetection = null; + } else { + turnDetection = + RealtimeAudioInputTurnDetection.ofSemanticVad( + RealtimeAudioInputTurnDetection.SemanticVad.builder().build()); + } + + return RealtimeAudioConfigInput.builder() + .turnDetection(turnDetection) + .format( + RealtimeAudioFormats.AudioPcm.builder() + .type(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM) + .rate(RealtimeAudioFormats.AudioPcm.Rate._24000) + .build()) + .build(); + } + + public void inputAudio(@Nonnull final byte[] rawAudioChunk) { + if (rawAudioChunk.length == 0) { + return; + } + var cursorLeft = 0; + while (cursorLeft < rawAudioChunk.length) { + final var cursorRight = + Math.min(cursorLeft + MAX_DATA_CHUNK_SIZE_BYTES, rawAudioChunk.length); + final var part = Arrays.copyOfRange(rawAudioChunk, cursorLeft, cursorRight); + final var audioInputMessage = + InputAudioBufferAppendEvent.builder() + .audio(Base64.getEncoder().encodeToString(part)) + .build(); + super.sendMessage(audioInputMessage); + cursorLeft += MAX_DATA_CHUNK_SIZE_BYTES; + } + + if (eagerTurnDetection) { + final var commitAudioMessage = InputAudioBufferCommitEvent.builder().build(); + super.sendMessage(commitAudioMessage); + askForResponse(); + } + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/TextToSpeechRealtimeClient.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/TextToSpeechRealtimeClient.java new file mode 100644 index 000000000..d0971fc44 --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/TextToSpeechRealtimeClient.java @@ -0,0 +1,93 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.openai.models.realtime.ConversationItem; +import com.openai.models.realtime.ConversationItemCreateEvent; +import com.openai.models.realtime.RealtimeAudioConfigInput; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.openai.models.realtime.RealtimeConversationItemUserMessage; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.util.Map; +import java.util.Optional; +import java.util.Timer; +import java.util.concurrent.CompletableFuture; +import java.util.stream.Stream; +import javax.annotation.Nonnull; +import lombok.extern.slf4j.Slf4j; + +@Slf4j +class TextToSpeechRealtimeClient extends ToAudioRealtimeClient implements TextInputChannel { + + private static final String SYSTEM_PROMPT = + "you are a speaker and your role is to read (produce audio) of the user input speech. voice user text input, " + + "do not answer questions, just read them"; + + public TextToSpeechRealtimeClient( + @Nonnull final String url, + @Nonnull final Map httpHeaders, + @Nonnull final AudioOutputChannel outputConsumer, + @Nonnull final RealtimeParam... params) { + super( + url, + httpHeaders, + outputConsumer, + true, + Stream.concat( + Stream.of((RealtimeParam) new RealtimeParamSystemPrompt(SYSTEM_PROMPT)), + Stream.of(params)) + .toArray(RealtimeParam[]::new)); + } + + TextToSpeechRealtimeClient( + @Nonnull final HttpClient httpClient, + @Nonnull final CompletableFuture ws, + @Nonnull final Timer timer, + @Nonnull final AudioOutputChannel outputConsumer, + @Nonnull final RealtimeParam... params) { + super( + httpClient, + ws, + timer, + outputConsumer, + true, + Stream.concat( + Stream.of((RealtimeParam) new RealtimeParamSystemPrompt(SYSTEM_PROMPT)), + Stream.of(params)) + .toArray(RealtimeParam[]::new)); + } + + @Override + @Nonnull + protected RealtimeAudioConfigInput inputConfig() { + return RealtimeAudioConfigInput.builder() + .turnDetection(Optional.empty()) + .format( + RealtimeAudioFormats.AudioPcm.builder() + .type(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM) + .rate(RealtimeAudioFormats.AudioPcm.Rate._24000) + .build()) + .build(); + } + + public void sendText(@Nonnull final String text) { + final var message = + ConversationItemCreateEvent.builder() + .item( + ConversationItem.ofRealtimeConversationItemUserMessage( + RealtimeConversationItemUserMessage.builder() + .addContent( + RealtimeConversationItemUserMessage.Content.builder() + .text(text) + .type(RealtimeConversationItemUserMessage.Content.Type.INPUT_TEXT) + .build()) + .build())) + .build(); + + super.sendMessage(message); + if (eagerTurnDetection) { + askForResponse(); + } + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/ToAudioRealtimeClient.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/ToAudioRealtimeClient.java new file mode 100644 index 000000000..0187c0b1d --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/ToAudioRealtimeClient.java @@ -0,0 +1,184 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.fasterxml.jackson.databind.JsonNode; +import com.openai.models.realtime.RealtimeAudioConfig; +import com.openai.models.realtime.RealtimeAudioConfigInput; +import com.openai.models.realtime.RealtimeAudioConfigOutput; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.openai.models.realtime.RealtimeSessionCreateRequest; +import com.openai.models.realtime.SessionUpdateEvent; +import com.openai.models.realtime.clientsecrets.ClientSecretCreateParams; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.util.Base64; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.Set; +import java.util.Timer; +import java.util.concurrent.CompletableFuture; +import javax.annotation.Nonnull; +import lombok.extern.slf4j.Slf4j; + +/** Implements common functionality for realtime api clients that output audio */ +@Slf4j +abstract class ToAudioRealtimeClient extends WSOpenAiRealtimeClient { + + private static final Set HANDLED_RESPONSE_TYPES = + Set.of("response.output_audio.delta", "response.output_audio.done"); + private static final List OUTPUT_MODALITIES = + List.of(RealtimeSessionCreateRequest.OutputModality.AUDIO); + private static final byte[] EMPTY_BYTE_ARRAY = new byte[0]; + + private static final Map FALLBACK_DEFAULT_PARAMS = + Map.of( + RealtimeParam.ParamName.OUTPUT_VOICE, RealtimeParamVoice.DEFAULT_1, + RealtimeParam.ParamName.TURN_DETECTION, RealtimeParamTurnDetection.BY_MODEL_AUTO, + RealtimeParam.ParamName.SYSTEM_PROMPT, new RealtimeParamSystemPrompt("")); + + final AudioOutputChannel outputConsumer; + final RealtimeAudioConfigOutput.Voice.UnionMember1 voice; + + /** defines if every call to the client should be considered conversation turn */ + protected final boolean eagerTurnDetection; + + final String systemPrompt; + + /** + * Constructs the object + * + * @param url - realtime api endpoint url + * @param httpHeaders - http headers (key - value) for client to use + * @param outputConsumer - consumer of audio bytes in pcm 24000 Hz mono little endian format + * @param defaultTurnDetectionEager - if explicit cfg for turn detection was not specified, this + * turn detection eagerness flag will be used (true results in EACH_CALL_IS_A_TURN handling) + * @param params - possible overrides for default params (e.g. voice, system prompt, etc.) + */ + public ToAudioRealtimeClient( + @Nonnull final String url, + @Nonnull final Map httpHeaders, + @Nonnull final AudioOutputChannel outputConsumer, + final boolean defaultTurnDetectionEager, + @Nonnull final RealtimeParam... params) { + super(url, httpHeaders, HANDLED_RESPONSE_TYPES); + final var defaults = new HashMap<>(FALLBACK_DEFAULT_PARAMS); + defaults.put( + RealtimeParam.ParamName.TURN_DETECTION, + defaultTurnDetectionEager + ? RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN + : RealtimeParamTurnDetection.BY_MODEL_AUTO); + final var resolvedParams = resolveParams(defaults, params); + + this.outputConsumer = outputConsumer; + this.voice = + mapVoice((RealtimeParamVoice) resolvedParams.get(RealtimeParam.ParamName.OUTPUT_VOICE)); + this.eagerTurnDetection = + RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN.equals( + resolvedParams.get(RealtimeParam.ParamName.TURN_DETECTION)); + this.systemPrompt = + resolvedParams.get(RealtimeParam.ParamName.SYSTEM_PROMPT).getValueAsString(); + } + + ToAudioRealtimeClient( + @Nonnull final HttpClient httpClient, + @Nonnull final CompletableFuture ws, + @Nonnull final Timer timer, + @Nonnull final AudioOutputChannel outputConsumer, + final boolean eagerTurnDetection, + @Nonnull final RealtimeParam... params) { + super(httpClient, ws, timer, HANDLED_RESPONSE_TYPES); + final var defaults = new HashMap<>(FALLBACK_DEFAULT_PARAMS); + defaults.put( + RealtimeParam.ParamName.TURN_DETECTION, + eagerTurnDetection + ? RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN + : RealtimeParamTurnDetection.BY_MODEL_AUTO); + final var resolvedParams = resolveParams(defaults, params); + this.outputConsumer = outputConsumer; + this.voice = + mapVoice((RealtimeParamVoice) resolvedParams.get(RealtimeParam.ParamName.OUTPUT_VOICE)); + this.eagerTurnDetection = + RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN.equals( + resolvedParams.get(RealtimeParam.ParamName.TURN_DETECTION)); + this.systemPrompt = + resolvedParams.get(RealtimeParam.ParamName.SYSTEM_PROMPT).getValueAsString(); + } + + @Nonnull + protected abstract RealtimeAudioConfigInput inputConfig(); + + @Override + @Nonnull + protected Optional getSystemPrompt() { + return systemPrompt.isEmpty() ? Optional.empty() : Optional.of(systemPrompt); + } + + @Override + protected void onResponse(@Nonnull final String eventType, @Nonnull final JsonNode event) { + if ("response.output_audio.delta".equals(eventType)) { + final var base64Audio = event.get("delta").asText(); + final byte[] audio = Base64.getDecoder().decode(base64Audio); + this.outputConsumer.outputAudio(audio, false); + } else if ("response.output_audio.done".equals(eventType)) { + this.outputConsumer.outputAudio(EMPTY_BYTE_ARRAY, true); + } else { + log.warn("skipping message type: {}", eventType); + } + } + + /* + side effect: modifies defaults + */ + @Nonnull + protected Map resolveParams( + final @Nonnull Map defaults, + final @Nonnull RealtimeParam... params) { + for (final RealtimeParam param : params) { + if (param == null) { + log.warn("skipping null param for realtime client"); + continue; + } + defaults.put(param.getParamName(), param); + } + return defaults; + } + + @Nonnull + private RealtimeAudioConfigOutput.Voice.UnionMember1 mapVoice( + @Nonnull final RealtimeParamVoice voice) { + if (voice.equals(RealtimeParamVoice.DEFAULT_1)) { + return RealtimeAudioConfigOutput.Voice.UnionMember1.MARIN; + } else if (voice.equals(RealtimeParamVoice.DEFAULT_2)) { + return RealtimeAudioConfigOutput.Voice.UnionMember1.ECHO; + } + return RealtimeAudioConfigOutput.Voice.UnionMember1.of(voice.getValueAsString()); + } + + @Override + @Nonnull + protected SessionUpdateEvent sessionConfiguration() { + return SessionUpdateEvent.builder() + .session( + ClientSecretCreateParams.Session.ofRealtime( + RealtimeSessionCreateRequest.builder() + .outputModalities(OUTPUT_MODALITIES) + .audio( + RealtimeAudioConfig.builder() + .input(inputConfig()) + .output( + RealtimeAudioConfigOutput.builder() + .format( + RealtimeAudioFormats.AudioPcm.builder() + .type(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM) + .rate(RealtimeAudioFormats.AudioPcm.Rate._24000) + .build()) + .voice(voice) + .build()) + .build()) + .build()) + .asRealtime()) + .build(); + } +} diff --git a/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/WSOpenAiRealtimeClient.java b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/WSOpenAiRealtimeClient.java new file mode 100644 index 000000000..a5c44e4f4 --- /dev/null +++ b/foundation-models/openai/src/main/java/com/sap/ai/sdk/foundationmodels/openai/realtime/WSOpenAiRealtimeClient.java @@ -0,0 +1,220 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import com.fasterxml.jackson.annotation.JsonAutoDetect; +import com.fasterxml.jackson.annotation.PropertyAccessor; +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.openai.models.realtime.ConversationItem; +import com.openai.models.realtime.ConversationItemCreateEvent; +import com.openai.models.realtime.RealtimeConversationItemSystemMessage; +import com.openai.models.realtime.ResponseCreateEvent; +import com.openai.models.realtime.SessionUpdateEvent; +import com.sap.ai.sdk.core.common.ClientException; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.util.Map; +import java.util.Optional; +import java.util.Set; +import java.util.Timer; +import java.util.TimerTask; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; +import javax.annotation.Nonnull; +import lombok.AccessLevel; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +@Slf4j +@RequiredArgsConstructor(access = AccessLevel.PACKAGE) +abstract class WSOpenAiRealtimeClient implements AutoCloseable { + + private static final int SUCCESS_FINISH_WSS_CODE = 1000; + private static final ObjectMapper JACKSON = + new ObjectMapper().setVisibility(PropertyAccessor.IS_GETTER, JsonAutoDetect.Visibility.NONE); + + /** + * WebSocket keeps TCP connection to the server. If connection idles, depending on the provider in + * cloud environments, cloud provider sends hard TCP RST (reset) after 60-300 seconds of + * inactivity, this is normally not configurable and would leave realtime conversation unusable on + * a pause. + * + *

In mobile networks this period is even shorter and typical CGNAT session can only idle for + * 10-15 seconds before it gets terminated + * + *

In realtime client we rely on explicit lifetime management (create/close), heartbeat + * mechanism is required to prevent unintended closing of connections while conversation is still + * expected to continue + */ + private static final int HEARTBEAT_INTERVAL_MILLIS = 4500; + + private static final String HEARTBEAT_TIMER_NAME = "wss_realtime_heartbeat"; + + private final HttpClient client; + private final CompletableFuture ws; + private final Timer heartbeatTimer; + private final Set handleMessageTypes; + + public WSOpenAiRealtimeClient( + @Nonnull final String url, + @Nonnull final Map httpHeaders, + @Nonnull final Set handleMessageTypes) { + this.client = HttpClient.newHttpClient(); + var wsBuilder = this.client.newWebSocketBuilder(); + for (final Map.Entry entry : httpHeaders.entrySet()) { + wsBuilder = wsBuilder.header(entry.getKey(), entry.getValue()); + } + this.ws = + wsBuilder.buildAsync( + URI.create(url), new BufferedWebSocketListener(this::onSocketOpen, this::onText)); + this.handleMessageTypes = handleMessageTypes; + this.heartbeatTimer = new Timer(HEARTBEAT_TIMER_NAME, true); + } + + public void askForResponse() { + WebSocket ws; + try { + ws = this.ws.join(); + } catch (final CompletionException e) { + throw new ClientException("Failed to establish web socket connection", e); + } + synchronized (this) { + try { + ws.sendText(JACKSON.writeValueAsString(ResponseCreateEvent.builder().build()), true); + } catch (final JsonProcessingException e) { + throw new ClientException("Failed to serialize ask for response", e); + } + ws.request(1); + } + } + + @Override + public void close() { + WebSocket ws; + try { + ws = this.ws.join(); + } catch (final CompletionException e) { + throw new ClientException("Failed to establish web socket connection", e); + } + synchronized (this) { + heartbeatTimer.cancel(); + try { + ws.sendClose(SUCCESS_FINISH_WSS_CODE, "done").join(); + } catch (final Exception e) { + log.error("Error while closing WebSocket", e); + } + // this.client.close(); // exists only since java 21 + } + } + + @Nonnull + protected Optional getSystemPrompt() { + return Optional.empty(); + } + + protected abstract void onResponse( + @Nonnull final String eventType, @Nonnull final JsonNode event); + + @Nonnull + protected abstract SessionUpdateEvent sessionConfiguration(); + + protected void sendMessage(@Nonnull final Object message) { + WebSocket ws; + try { + ws = this.ws.join(); + } catch (CompletionException e) { + throw new ClientException("Failed to establish web socket connection", e); + } + synchronized (this) { + try { + ws.sendText(JACKSON.writeValueAsString(message), true); + ws.request(1); + } catch (final JsonProcessingException e) { + throw new ClientException("Failed to serialize message", e); + } + } + } + + synchronized void onSocketOpen(@Nonnull final WebSocket ws) { + configureSession(ws); + configureConversation(ws); + scheduleHeartbeat(ws); + } + + protected void onText(@Nonnull final WebSocket webSocket, @Nonnull final CharSequence data) { + final JsonNode event; + try { + event = JACKSON.readTree(data.toString()); + } catch (final JsonProcessingException e) { + throw new ClientException("Error parsing JSON response from speech API", e); + } + final var eventType = event.get("type").asText(); + if (handleMessageTypes.contains(eventType)) { + onResponse(eventType, event); + } else { + log.trace("Unhandled event type: {}", eventType); + } + + webSocket.request(1); + } + + protected synchronized void sendPing(@Nonnull final WebSocket ws) { + if (ws.isInputClosed()) { + return; + } + ws.sendPing(ByteBuffer.wrap("ping".getBytes(StandardCharsets.UTF_8))).join(); + ws.request(1); + } + + private void configureSession(@Nonnull final WebSocket ws) { + final SessionUpdateEvent sue = sessionConfiguration(); + try { + ws.sendText(JACKSON.writeValueAsString(sue), true); + } catch (final JsonProcessingException e) { + throw new ClientException("Failed to serialize session request", e); + } + + ws.request(1); + } + + private void configureConversation(@Nonnull final WebSocket ws) { + final var systemPrompt = getSystemPrompt(); + if (systemPrompt.isEmpty() || systemPrompt.get().isEmpty()) { + return; + } + final var systemConversationItem = + ConversationItemCreateEvent.builder() + .item( + ConversationItem.ofRealtimeConversationItemSystemMessage( + RealtimeConversationItemSystemMessage.builder() + .addContent( + RealtimeConversationItemSystemMessage.Content.builder() + .text(systemPrompt.get()) + .type(RealtimeConversationItemSystemMessage.Content.Type.INPUT_TEXT) + .build()) + .build())) + .build(); + + try { + final var json = JACKSON.writeValueAsString(systemConversationItem); + ws.sendText(json, true); + } catch (final JsonProcessingException e) { + throw new ClientException("Failed to serialize message", e); + } + ws.request(1); + } + + private void scheduleHeartbeat(@Nonnull final WebSocket ws) { + final TimerTask task = + new TimerTask() { + @Override + public void run() { + sendPing(ws); + } + }; + heartbeatTimer.scheduleAtFixedRate(task, 0, HEARTBEAT_INTERVAL_MILLIS); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/BufferedWebSocketListenerUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/BufferedWebSocketListenerUnitTest.java new file mode 100644 index 000000000..fcb379c42 --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/BufferedWebSocketListenerUnitTest.java @@ -0,0 +1,139 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatNoException; +import static org.mockito.ArgumentMatchers.anyLong; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; + +import java.net.http.WebSocket; +import java.nio.ByteBuffer; +import java.util.ArrayList; +import java.util.List; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +class BufferedWebSocketListenerUnitTest { + + private WebSocket webSocketMock; + private List openCallbacks; + private List textCallbacks; + + @BeforeEach + void setUp() { + webSocketMock = mock(WebSocket.class); + openCallbacks = new ArrayList<>(); + textCallbacks = new ArrayList<>(); + } + + private BufferedWebSocketListener build() { + return new BufferedWebSocketListener( + openCallbacks::add, (ws, data) -> textCallbacks.add(data.toString())); + } + + @Test + void onOpenInvokesConsumerWithWebSocket() { + build().onOpen(webSocketMock); + + assertThat(openCallbacks).containsExactly(webSocketMock); + } + + @Test + void onOpenRequestsNextMessage() { + build().onOpen(webSocketMock); + + verify(webSocketMock).request(1L); + } + + @Test + void onTextInvokesConsumerWhenLastFrameReceived() { + build().onText(webSocketMock, "hello", true); + + assertThat(textCallbacks).containsExactly("hello"); + } + + @Test + void onTextDoesNotInvokeConsumerForPartialFrame() { + build().onText(webSocketMock, "hel", false); + + assertThat(textCallbacks).isEmpty(); + } + + @Test + void onTextRequestsNextMessageForPartialFrame() { + build().onText(webSocketMock, "hel", false); + + verify(webSocketMock).request(1L); + } + + @Test + void onTextReturnedFutureIsAlreadyCompleted() { + final var future = build().onText(webSocketMock, "hello", true); + + assertThat(future.toCompletableFuture()).isDone(); + } + + @Test + void onTextBuffersPartialsAndDeliversOnLastFrame() { + final var listener = build(); + listener.onText(webSocketMock, "foo", false); + listener.onText(webSocketMock, "bar", false); + listener.onText(webSocketMock, "baz", true); + + assertThat(textCallbacks).containsExactly("foobarbaz"); + } + + @Test + void onTextResetsBufferAfterCompleteMessage() { + final var listener = build(); + listener.onText(webSocketMock, "first", true); + listener.onText(webSocketMock, "second", true); + + assertThat(textCallbacks).containsExactly("first", "second"); + } + + @Test + void onTextBufferDoesNotLeakAcrossMessages() { + final var listener = build(); + listener.onText(webSocketMock, "part1", false); + listener.onText(webSocketMock, "part2", true); // completes first message + listener.onText(webSocketMock, "part3", true); // second message starts clean + + assertThat(textCallbacks).containsExactly("part1part2", "part3"); + } + + @Test + void onErrorDoesNotThrow() { + assertThatNoException() + .isThrownBy(() -> build().onError(webSocketMock, new RuntimeException("boom"))); + } + + @Test + void onErrorDoesNotInvokeTextConsumer() { + build().onError(webSocketMock, new RuntimeException("boom")); + + assertThat(textCallbacks).isEmpty(); + } + + @Test + void onBinaryDoesNotInvokeTextConsumer() { + build().onBinary(webSocketMock, ByteBuffer.wrap(new byte[] {0x01}), true); + + assertThat(textCallbacks).isEmpty(); + } + + @Test + void onBinaryReturnedFutureIsAlreadyCompleted() { + final var future = build().onBinary(webSocketMock, ByteBuffer.wrap(new byte[] {0x01}), true); + + assertThat(future.toCompletableFuture()).isDone(); + } + + @Test + void onBinaryDoesNotRequestNextMessage() { + build().onBinary(webSocketMock, ByteBuffer.wrap(new byte[] {0x01}), true); + + verify(webSocketMock, never()).request(anyLong()); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/OpenAiRealtimeClientUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/OpenAiRealtimeClientUnitTest.java new file mode 100644 index 000000000..f6bd39376 --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/OpenAiRealtimeClientUnitTest.java @@ -0,0 +1,138 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; +import com.sap.cloud.sdk.cloudplatform.connectivity.DefaultHttpDestination; +import org.junit.jupiter.api.Test; + +class OpenAiRealtimeClientUnitTest { + + private static final AudioOutputChannel NO_OP_OUTPUT = (audio, done) -> {}; + + @Test + void constructorStoresDestination() { + final var destination = DefaultHttpDestination.builder("https://api.example.com").build(); + final var client = new OpenAiRealtimeClient(destination); + + assertThat(client.destination).isSameAs(destination); + } + + @Test + void buildRealtimeHeadersReturnsEmptyMapWhenDestinationHasNoHeaders() { + final var destination = DefaultHttpDestination.builder("https://api.example.com").build(); + final var client = new OpenAiRealtimeClient(destination); + + assertThat(client.buildRealtimeHeaders()).isEmpty(); + } + + @Test + void buildRealtimeHeadersCopiesAllDestinationHeaders() { + final var destination = + DefaultHttpDestination.builder("https://api.example.com") + .header("Authorization", "Bearer token-123") + .header("X-Custom-Header", "custom-value") + .build(); + final var client = new OpenAiRealtimeClient(destination); + + final var headers = client.buildRealtimeHeaders(); + + assertThat(headers) + .containsEntry("Authorization", "Bearer token-123") + .containsEntry("X-Custom-Header", "custom-value"); + } + + @Test + void buildRealtimeHeadersReturnsModifiableCopy() { + final var destination = + DefaultHttpDestination.builder("https://api.example.com") + .header("Authorization", "Bearer token") + .build(); + final var client = new OpenAiRealtimeClient(destination); + + final var headers = client.buildRealtimeHeaders(); + headers.put("injected", "value"); + + assertThat(client.buildRealtimeHeaders()).doesNotContainKey("injected"); + } + + @Test + void getRealtimeEndpointBuildsWssUrlAndAppendsSuffix() { + final var destination = + DefaultHttpDestination.builder( + "https://my-resource.openai.azure.com/openai/deployments/gpt-4o") + .build(); + final var client = new OpenAiRealtimeClient(destination); + + assertThat(client.getRealtimeEndpoint()).startsWith("wss://").endsWith("/v1/realtime"); + } + + @Test + void getRealtimeEndpointReplacesApiSubdomainWithRealtime() { + final var destination = + DefaultHttpDestination.builder("https://api.example.com/some/path").build(); + final var client = new OpenAiRealtimeClient(destination); + + final var endpoint = client.getRealtimeEndpoint(); + + assertThat(endpoint).startsWith("wss://realtime.example.com"); + assertThat(endpoint).doesNotContain("api.example.com"); + } + + @Test + void getRealtimeEndpointPreservesNonApiSubdomain() { + final var destination = + DefaultHttpDestination.builder("https://my-resource.openai.azure.com/openai").build(); + final var client = new OpenAiRealtimeClient(destination); + + final var endpoint = client.getRealtimeEndpoint(); + + assertThat(endpoint).startsWith("wss://my-resource.openai.azure.com"); + } + + @Test + void getRealtimeEndpointThrowsOnUriWithoutDoubleSlash() { + final var destination = DefaultHttpDestination.builder("https://host").build(); + final var client = + new OpenAiRealtimeClient(destination) { + @Override + public String getRealtimeEndpoint() { + // Force a URI that produces != 2 parts when split on "//" + final var sb = new StringBuilder(PATH_BUFFER_SIZE); + sb.append("wss://"); + final var pathParts = "no-double-slash".split("//"); + if (pathParts.length != 2) { + throw new IllegalArgumentException("Invalid destination URI: no-double-slash"); + } + return sb.toString(); + } + }; + + assertThatThrownBy(client::getRealtimeEndpoint) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Invalid destination URI"); + } + + @Test + void textToSpeechReturnsTextInputChannel() { + final var destination = DefaultHttpDestination.builder("https://api.example.com").build(); + final var client = new OpenAiRealtimeClient(destination); + + final var channel = client.textToSpeech(NO_OP_OUTPUT); + + assertThat(channel).isInstanceOf(TextInputChannel.class); + } + + @Test + void speechToSpeechReturnsAudioInputChannel() { + final var destination = DefaultHttpDestination.builder("https://api.example.com").build(); + final var client = new OpenAiRealtimeClient(destination); + + final var channel = client.speechToSpeech(NO_OP_OUTPUT); + + assertThat(channel).isInstanceOf(AudioInputChannel.class); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamSystemPromptUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamSystemPromptUnitTest.java new file mode 100644 index 000000000..ee0f3a4b8 --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamSystemPromptUnitTest.java @@ -0,0 +1,48 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.UUID; +import org.junit.jupiter.api.Test; + +class RealtimeParamSystemPromptUnitTest { + + @Test + void getParamName() { + assertThat(new RealtimeParamSystemPrompt(randomString()).getParamName()) + .isEqualTo(RealtimeParam.ParamName.SYSTEM_PROMPT); + } + + @Test + void getValueAsString() { + final var randomValue = randomString(); + final var prompt = new RealtimeParamSystemPrompt(randomValue); + assertThat(prompt.getValueAsString()).isEqualTo(randomValue); + } + + @Test + void testEquals() { + final var randomValue = randomString(); + final var prompt1 = new RealtimeParamSystemPrompt(randomValue); + final var prompt2 = new RealtimeParamSystemPrompt(randomValue); + assertThat(prompt1).isEqualTo(prompt2); + + final var otherRandomValue = randomString(); + final var prompt3 = new RealtimeParamSystemPrompt(otherRandomValue); + assertThat(prompt1).isNotEqualTo(prompt3); + assertThat(prompt2).isNotEqualTo(prompt3); + } + + @Test + void testHashCode() { + final var randomValue = randomString(); + final var prompt1 = new RealtimeParamSystemPrompt(randomValue); + final var prompt2 = new RealtimeParamSystemPrompt(randomValue); + + assertThat(prompt1.hashCode()).isEqualTo(prompt2.hashCode()); + } + + private String randomString() { + return UUID.randomUUID().toString().substring(0, 20).replace("-", " "); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamTurnDetectionUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamTurnDetectionUnitTest.java new file mode 100644 index 000000000..4e1abaa9a --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamTurnDetectionUnitTest.java @@ -0,0 +1,59 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.UUID; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +class RealtimeParamTurnDetectionUnitTest { + + private RealtimeParamTurnDetection random; + private String expectedRandomValue; + + @BeforeEach + void setUp() { + expectedRandomValue = UUID.randomUUID().toString().substring(0, 20); + random = new RealtimeParamTurnDetection(expectedRandomValue); + } + + @Test + void getParamName() { + assertThat(random.getParamName()).isEqualTo(RealtimeParam.ParamName.TURN_DETECTION); + assertThat(RealtimeParamTurnDetection.BY_MODEL_AUTO.getParamName()) + .isEqualTo(RealtimeParam.ParamName.TURN_DETECTION); + assertThat(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN.getParamName()) + .isEqualTo(RealtimeParam.ParamName.TURN_DETECTION); + } + + @Test + void getValueAsString() { + assertThat(random.getValueAsString()).isEqualTo(expectedRandomValue); + assertThat(RealtimeParamTurnDetection.BY_MODEL_AUTO.getValueAsString()) + .isEqualTo("BY_MODEL_AUTO"); + assertThat(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN.getValueAsString()) + .isEqualTo("EACH_CALL_IS_A_TURN"); + } + + @Test + void testEquals() { + assertThat(random).isEqualTo(random); + assertThat(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN) + .isEqualTo(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN); + assertThat(RealtimeParamTurnDetection.BY_MODEL_AUTO) + .isEqualTo(RealtimeParamTurnDetection.BY_MODEL_AUTO); + assertThat(random).isNotEqualTo(RealtimeParamTurnDetection.BY_MODEL_AUTO); + assertThat(random).isNotEqualTo(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN); + assertThat(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN) + .isNotEqualTo(RealtimeParamTurnDetection.BY_MODEL_AUTO); + } + + @Test + void testHashCode() { + assertThat(random.hashCode()).isEqualTo(random.hashCode()); + assertThat(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN.hashCode()) + .isEqualTo(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN.hashCode()); + assertThat(RealtimeParamTurnDetection.BY_MODEL_AUTO.hashCode()) + .isEqualTo(RealtimeParamTurnDetection.BY_MODEL_AUTO.hashCode()); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamVoiceUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamVoiceUnitTest.java new file mode 100644 index 000000000..a441328a6 --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/RealtimeParamVoiceUnitTest.java @@ -0,0 +1,59 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.UUID; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +class RealtimeParamVoiceUnitTest { + + private RealtimeParamVoice random; + private String expectedRandomValue; + + @BeforeEach + void setUp() { + expectedRandomValue = UUID.randomUUID().toString().substring(0, 20); + random = new RealtimeParamVoice(expectedRandomValue); + } + + @Test + void withExplicitVoice() { + assertThat(random).isEqualTo(RealtimeParamVoice.withExplicitVoice(expectedRandomValue)); + } + + @Test + void getParamName() { + assertThat(random.getParamName()).isEqualTo(RealtimeParam.ParamName.OUTPUT_VOICE); + assertThat(RealtimeParamVoice.DEFAULT_1.getParamName()) + .isEqualTo(RealtimeParam.ParamName.OUTPUT_VOICE); + assertThat(RealtimeParamVoice.DEFAULT_1.getParamName()) + .isEqualTo(RealtimeParam.ParamName.OUTPUT_VOICE); + } + + @Test + void getValueAsString() { + assertThat(random.getValueAsString()).isEqualTo(expectedRandomValue); + assertThat(RealtimeParamVoice.DEFAULT_1.getValueAsString()).isEqualTo("DEFAULT_1"); + assertThat(RealtimeParamVoice.DEFAULT_2.getValueAsString()).isEqualTo("DEFAULT_2"); + } + + @Test + void testEquals() { + assertThat(random).isEqualTo(random); + assertThat(RealtimeParamVoice.DEFAULT_1).isEqualTo(RealtimeParamVoice.DEFAULT_1); + assertThat(RealtimeParamVoice.DEFAULT_2).isEqualTo(RealtimeParamVoice.DEFAULT_2); + assertThat(random).isNotEqualTo(RealtimeParamVoice.DEFAULT_1); + assertThat(random).isNotEqualTo(RealtimeParamVoice.DEFAULT_2); + assertThat(RealtimeParamVoice.DEFAULT_1).isNotEqualTo(RealtimeParamVoice.DEFAULT_2); + } + + @Test + void testHashCode() { + assertThat(random.hashCode()).isEqualTo(random.hashCode()); + assertThat(RealtimeParamVoice.DEFAULT_1.hashCode()) + .isEqualTo(RealtimeParamVoice.DEFAULT_1.hashCode()); + assertThat(RealtimeParamVoice.DEFAULT_2.hashCode()) + .isEqualTo(RealtimeParamVoice.DEFAULT_2.hashCode()); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/SpeechToSpeechRealtimeClientUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/SpeechToSpeechRealtimeClientUnitTest.java new file mode 100644 index 000000000..8c91d6610 --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/SpeechToSpeechRealtimeClientUnitTest.java @@ -0,0 +1,202 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.atLeastOnce; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.util.Arrays; +import java.util.Base64; +import java.util.List; +import java.util.Timer; +import java.util.concurrent.CompletableFuture; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; + +class SpeechToSpeechRealtimeClientUnitTest { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private AudioOutputChannel outputConsumerMock; + private WebSocket webSocketMock; + + @BeforeEach + void setUp() { + outputConsumerMock = mock(AudioOutputChannel.class); + webSocketMock = mock(WebSocket.class); + when(webSocketMock.sendText(any(), anyBoolean())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.sendClose(anyInt(), anyString())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.sendPing(any())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + } + + private SpeechToSpeechRealtimeClient build(final RealtimeParam... params) { + return new SpeechToSpeechRealtimeClient( + mock(HttpClient.class), + CompletableFuture.completedFuture(webSocketMock), + mock(Timer.class), + outputConsumerMock, + params); + } + + private List captureAllSentTexts() { + final var captor = ArgumentCaptor.forClass(CharSequence.class); + try { + verify(webSocketMock, atLeastOnce()).sendText(captor.capture(), anyBoolean()); + return captor.getAllValues().stream().map(CharSequence::toString).toList(); + } catch (final org.mockito.exceptions.verification.WantedButNotInvoked e) { + return List.of(); + } + } + + @Test + void defaultTurnDetectionIsNotEager() { + assertThat(build().eagerTurnDetection).isFalse(); + } + + @Test + void turnDetectionBecomesEagerWhenParamProvided() { + assertThat(build(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN).eagerTurnDetection).isTrue(); + } + + @Test + void turnDetectionRemainsNonEagerWhenByModelAutoProvided() { + assertThat(build(RealtimeParamTurnDetection.BY_MODEL_AUTO).eagerTurnDetection).isFalse(); + } + + @Test + void outputConsumerIsStoredFromConstructor() { + assertThat(build().outputConsumer).isSameAs(outputConsumerMock); + } + + @Test + void inputConfigHasPcm24000Format() { + final var format = build().inputConfig().format().orElseThrow().audioPcm().orElseThrow(); + + assertThat(format.type().orElseThrow()).isEqualTo(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM); + assertThat(format.rate().orElseThrow()).isEqualTo(RealtimeAudioFormats.AudioPcm.Rate._24000); + } + + @Test + void inputConfigHasSemanticVadTurnDetectionWhenNotEager() { + final var config = build(RealtimeParamTurnDetection.BY_MODEL_AUTO).inputConfig(); + + assertThat(config.turnDetection().orElseThrow().semanticVad()).isPresent(); + } + + @Test + void inputConfigHasNoTurnDetectionWhenEager() { + final var config = build(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN).inputConfig(); + + assertThat(config.turnDetection()).isEmpty(); + } + + @Test + void inputAudioDoesNothingForEmptyArray() { + build().inputAudio(new byte[0]); + + verify(webSocketMock, never()).sendText(any(), anyBoolean()); + } + + @Test + void inputAudioSendsOneAppendEventForSmallChunk() throws Exception { + final byte[] audio = {0x01, 0x02, 0x03}; // simple test fixture + build().inputAudio(audio); + + final var sent = captureAllSentTexts(); + assertThat(sent).hasSize(1); + final var node = MAPPER.readTree(sent.get(0)); + assertThat(node.get("type").asText()).isEqualTo("input_audio_buffer.append"); + assertThat(Base64.getDecoder().decode(node.get("audio").asText())).isEqualTo(audio); + } + + @Test + void inputAudioDoesNotSendCommitWhenNotEager() throws Exception { + build().inputAudio(new byte[] {0x01}); + + final var sent = captureAllSentTexts(); + assertThat(sent).hasSize(1); + assertThat(MAPPER.readTree(sent.get(0)).get("type").asText()) + .isEqualTo("input_audio_buffer.append"); + } + + @Test + void inputAudioSendsCommitAndResponseCreateWhenEager() throws Exception { + build(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN) + .inputAudio(new byte[] {0x01}); // simple test fixture + + final var sent = captureAllSentTexts(); + // append + commit + response.create (from askForResponse) + assertThat(sent).hasSize(3); + assertThat(MAPPER.readTree(sent.get(0)).get("type").asText()) + .isEqualTo("input_audio_buffer.append"); + assertThat(MAPPER.readTree(sent.get(1)).get("type").asText()) + .isEqualTo("input_audio_buffer.commit"); + assertThat(MAPPER.readTree(sent.get(2)).get("type").asText()).isEqualTo("response.create"); + } + + @Test + void inputAudioSplitsLargeInputInto8192ByteChunks() throws Exception { + final int chunkSize = 8192; + final byte[] audio = new byte[chunkSize * 2 + 100]; + Arrays.fill(audio, (byte) 0x42); // simple test fixture + build().inputAudio(audio); + + final var sent = captureAllSentTexts(); + assertThat(sent).hasSize(3); + + final byte[] chunk1 = + Base64.getDecoder().decode(MAPPER.readTree(sent.get(0)).get("audio").asText()); + final byte[] chunk2 = + Base64.getDecoder().decode(MAPPER.readTree(sent.get(1)).get("audio").asText()); + final byte[] chunk3 = + Base64.getDecoder().decode(MAPPER.readTree(sent.get(2)).get("audio").asText()); + assertThat(chunk1).hasSize(chunkSize); + assertThat(chunk2).hasSize(chunkSize); + assertThat(chunk3).hasSize(100); + + final byte[] reassembled = new byte[audio.length]; + System.arraycopy(chunk1, 0, reassembled, 0, chunkSize); + System.arraycopy(chunk2, 0, reassembled, chunkSize, chunkSize); + System.arraycopy(chunk3, 0, reassembled, chunkSize * 2, 100); + assertThat(reassembled).isEqualTo(audio); + } + + @Test + void inputAudioSendsExactlyOneChunkWhenSizeEqualsChunkSize() { + build().inputAudio(new byte[8192]); + + assertThat(captureAllSentTexts()).hasSize(1); + } + + @Test + void inputAudioWithEagerTurnDetectionSendsCommitAfterAllChunks() throws Exception { + final byte[] audio = new byte[8192 + 1]; + build(RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN).inputAudio(audio); + + final var sent = captureAllSentTexts(); + // 2 appends + 1 commit + 1 response.create + assertThat(sent).hasSize(4); + assertThat(MAPPER.readTree(sent.get(0)).get("type").asText()) + .isEqualTo("input_audio_buffer.append"); + assertThat(MAPPER.readTree(sent.get(1)).get("type").asText()) + .isEqualTo("input_audio_buffer.append"); + assertThat(MAPPER.readTree(sent.get(2)).get("type").asText()) + .isEqualTo("input_audio_buffer.commit"); + assertThat(MAPPER.readTree(sent.get(3)).get("type").asText()).isEqualTo("response.create"); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/TextToSpeechRealtimeClientUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/TextToSpeechRealtimeClientUnitTest.java new file mode 100644 index 000000000..5addcf54b --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/TextToSpeechRealtimeClientUnitTest.java @@ -0,0 +1,185 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.atLeastOnce; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.util.List; +import java.util.Timer; +import java.util.concurrent.CompletableFuture; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; + +class TextToSpeechRealtimeClientUnitTest { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private AudioOutputChannel outputConsumerMock; + private WebSocket webSocketMock; + + @BeforeEach + void setUp() { + outputConsumerMock = mock(AudioOutputChannel.class); + webSocketMock = mock(WebSocket.class); + when(webSocketMock.sendText(any(), anyBoolean())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.sendClose(anyInt(), anyString())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.sendPing(any())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + } + + private TextToSpeechRealtimeClient build(final RealtimeParam... params) { + return new TextToSpeechRealtimeClient( + mock(HttpClient.class), + CompletableFuture.completedFuture(webSocketMock), + mock(Timer.class), + outputConsumerMock, + params); + } + + private List captureAllSentTexts() { + final var captor = ArgumentCaptor.forClass(CharSequence.class); + try { + org.mockito.Mockito.verify(webSocketMock, atLeastOnce()) + .sendText(captor.capture(), anyBoolean()); + return captor.getAllValues().stream().map(CharSequence::toString).toList(); + } catch (final org.mockito.exceptions.verification.WantedButNotInvoked e) { + return List.of(); + } + } + + @Test + void defaultTurnDetectionIsEager() { + assertThat(build().eagerTurnDetection).isTrue(); + } + + @Test + void turnDetectionCanBeOverriddenToNonEagerByParam() { + assertThat(build(RealtimeParamTurnDetection.BY_MODEL_AUTO).eagerTurnDetection).isFalse(); + } + + @Test + void outputConsumerIsStoredFromConstructor() { + assertThat(build().outputConsumer).isSameAs(outputConsumerMock); + } + + @Test + void defaultSystemPromptIsSetByConstructor() { + assertThat(build().systemPrompt).isNotEmpty(); + } + + @Test + void callerSystemPromptOverridesDefaultBecauseItComesLast() { + final var custom = "My custom prompt."; + assertThat(build(new RealtimeParamSystemPrompt(custom)).systemPrompt).isEqualTo(custom); + } + + @Test + void inputConfigHasPcm24000Format() { + final var format = build().inputConfig().format().orElseThrow().audioPcm().orElseThrow(); + + assertThat(format.type().orElseThrow()).isEqualTo(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM); + assertThat(format.rate().orElseThrow()).isEqualTo(RealtimeAudioFormats.AudioPcm.Rate._24000); + } + + @Test + void inputConfigHasNoTurnDetection() { + assertThat(build().inputConfig().turnDetection()).isEmpty(); + } + + @Test + void sendTextSendsConversationItemCreateEvent() { + build().sendText("Hello world"); + + final var sent = captureAllSentTexts(); + assertThat(sent).isNotEmpty(); + final var conversationCreate = + sent.stream() + .map( + s -> { + try { + return MAPPER.readTree(s); + } catch (Exception e) { + throw new RuntimeException(e); + } + }) + .filter(n -> "conversation.item.create".equals(n.get("type").asText())) + .findFirst(); + assertThat(conversationCreate).isPresent(); + assertThat(conversationCreate.get().at("/item/content/0/text").asText()) + .isEqualTo("Hello world"); + } + + @Test + void sendTextSendsResponseCreateWhenEager() throws Exception { + build().sendText("Hello"); + + final var sent = captureAllSentTexts(); + final var types = + sent.stream() + .map( + s -> { + try { + return MAPPER.readTree(s).get("type").asText(); + } catch (Exception e) { + throw new RuntimeException(e); + } + }) + .toList(); + assertThat(types).contains("conversation.item.create", "response.create"); + } + + @Test + void sendTextDoesNotSendResponseCreateWhenNotEager() throws Exception { + build(RealtimeParamTurnDetection.BY_MODEL_AUTO).sendText("Hello"); + + final var sent = captureAllSentTexts(); + final var types = + sent.stream() + .map( + s -> { + try { + return MAPPER.readTree(s).get("type").asText(); + } catch (Exception e) { + throw new RuntimeException(e); + } + }) + .toList(); + assertThat(types).contains("conversation.item.create"); + assertThat(types).doesNotContain("response.create"); + } + + @Test + void sendTextPreservesFullTextInPayload() throws Exception { + final var longText = "a".repeat(500); + build().sendText(longText); + + final var sent = captureAllSentTexts(); + final var conversationCreate = + sent.stream() + .map( + s -> { + try { + return MAPPER.readTree(s); + } catch (Exception e) { + throw new RuntimeException(e); + } + }) + .filter(n -> "conversation.item.create".equals(n.get("type").asText())) + .findFirst() + .orElseThrow(); + assertThat(conversationCreate.at("/item/content/0/text").asText()).isEqualTo(longText); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/ToAudioRealtimeClientUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/ToAudioRealtimeClientUnitTest.java new file mode 100644 index 000000000..34b36f59a --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/ToAudioRealtimeClientUnitTest.java @@ -0,0 +1,304 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.openai.models.realtime.RealtimeAudioConfigInput; +import com.openai.models.realtime.RealtimeAudioConfigOutput; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.openai.models.realtime.RealtimeSessionCreateRequest; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import java.util.Base64; +import java.util.Map; +import java.util.Optional; +import javax.annotation.Nonnull; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; + +class ToAudioRealtimeClientUnitTest { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + private static final RealtimeAudioConfigInput FIXED_INPUT_CONFIG = + RealtimeAudioConfigInput.builder() + .format( + RealtimeAudioFormats.AudioPcm.builder() + .type(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM) + .rate(RealtimeAudioFormats.AudioPcm.Rate._24000) + .build()) + .build(); + + private AudioOutputChannel outputConsumerMock; + + @BeforeEach + void setUp() { + outputConsumerMock = mock(AudioOutputChannel.class); + } + + /** + * Concrete subclass that calls the URL-based constructor with a stub address. {@code buildAsync} + * is non-blocking so the constructor returns immediately without attempting a real connection. + */ + private static class DirectClient extends ToAudioRealtimeClient { + + DirectClient( + final AudioOutputChannel outputConsumer, + final boolean defaultEager, + final RealtimeParam... params) { + super("ws://localhost:0", Map.of(), outputConsumer, defaultEager, params); + } + + @Override + @Nonnull + protected RealtimeAudioConfigInput inputConfig() { + return FIXED_INPUT_CONFIG; + } + } + + private DirectClient build( + final AudioOutputChannel outputConsumer, + final boolean defaultEager, + final RealtimeParam... params) { + return new DirectClient(outputConsumer, defaultEager, params); + } + + @Test + void outputConsumerIsStoredFromConstructor() { + assertThat(build(outputConsumerMock, false).outputConsumer).isSameAs(outputConsumerMock); + } + + @Test + void defaultVoiceIsMarinWhenNoVoiceParamProvided() { + assertThat(build(outputConsumerMock, false).voice) + .isEqualTo(RealtimeAudioConfigOutput.Voice.UnionMember1.MARIN); + } + + @Test + void voiceIsEchoWhenDefault2Provided() { + assertThat(build(outputConsumerMock, false, RealtimeParamVoice.DEFAULT_2).voice) + .isEqualTo(RealtimeAudioConfigOutput.Voice.UnionMember1.ECHO); + } + + @Test + void voiceIsMarinWhenDefault1Provided() { + assertThat(build(outputConsumerMock, false, RealtimeParamVoice.DEFAULT_1).voice) + .isEqualTo(RealtimeAudioConfigOutput.Voice.UnionMember1.MARIN); + } + + @Test + void lastVoiceParamWinsWhenMultipleProvided() { + assertThat( + build( + outputConsumerMock, + false, + RealtimeParamVoice.DEFAULT_1, + RealtimeParamVoice.DEFAULT_2) + .voice) + .isEqualTo(RealtimeAudioConfigOutput.Voice.UnionMember1.ECHO); + } + + @Test + void turnDetectionFollowsDefaultFalseWhenNoParamProvided() { + assertThat(build(outputConsumerMock, false).eagerTurnDetection).isFalse(); + } + + @Test + void turnDetectionFollowsDefaultTrueWhenNoParamProvided() { + assertThat(build(outputConsumerMock, true).eagerTurnDetection).isTrue(); + } + + @Test + void turnDetectionIsEagerWhenEachCallIsATurnParamProvided() { + assertThat( + build(outputConsumerMock, false, RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN) + .eagerTurnDetection) + .isTrue(); + } + + @Test + void turnDetectionIsByModelWhenByModelAutoParamProvided() { + assertThat( + build(outputConsumerMock, true, RealtimeParamTurnDetection.BY_MODEL_AUTO) + .eagerTurnDetection) + .isFalse(); + } + + @Test + void lastTurnDetectionParamWinsWhenMultipleProvided() { + assertThat( + build( + outputConsumerMock, + false, + RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN, + RealtimeParamTurnDetection.BY_MODEL_AUTO) + .eagerTurnDetection) + .isFalse(); + } + + @Test + void systemPromptIsEmptyWhenNoParamProvided() { + assertThat(build(outputConsumerMock, false).systemPrompt).isEmpty(); + } + + @Test + void systemPromptIsStoredWhenParamProvided() { + assertThat( + build(outputConsumerMock, false, new RealtimeParamSystemPrompt("You are helpful.")) + .systemPrompt) + .isEqualTo("You are helpful."); + } + + @Test + void lastSystemPromptParamWinsWhenMultipleProvided() { + assertThat( + build( + outputConsumerMock, + false, + new RealtimeParamSystemPrompt("first"), + new RealtimeParamSystemPrompt("second")) + .systemPrompt) + .isEqualTo("second"); + } + + @Test + void getSystemPromptReturnsEmptyWhenSystemPromptIsBlank() { + assertThat(build(outputConsumerMock, false).getSystemPrompt()).isEqualTo(Optional.empty()); + } + + @Test + void getSystemPromptReturnsValueWhenSystemPromptIsSet() { + final var prompt = "You are a helpful assistant."; + assertThat( + build(outputConsumerMock, false, new RealtimeParamSystemPrompt(prompt)) + .getSystemPrompt()) + .isEqualTo(Optional.of(prompt)); + } + + @Test + void onResponseDecodesBase64AudioAndForwardsWithNotDoneFlag() throws Exception { + final var client = build(outputConsumerMock, false); + final byte[] rawAudio = {0x01, 0x02, 0x03}; + final var base64 = Base64.getEncoder().encodeToString(rawAudio); + final var event = + MAPPER.readTree("{\"type\":\"response.output_audio.delta\",\"delta\":\"" + base64 + "\"}"); + + client.onResponse("response.output_audio.delta", event); + + final var captor = ArgumentCaptor.forClass(byte[].class); + verify(outputConsumerMock).outputAudio(captor.capture(), eq(Boolean.FALSE)); + assertThat(captor.getValue()).isEqualTo(rawAudio); + } + + @Test + void onResponseForwardsEmptyBytesWithDoneFlagOnAudioDoneEvent() throws Exception { + final var client = build(outputConsumerMock, false); + final var event = MAPPER.readTree("{\"type\":\"response.output_audio.done\"}"); + + client.onResponse("response.output_audio.done", event); + + final var captor = ArgumentCaptor.forClass(byte[].class); + verify(outputConsumerMock).outputAudio(captor.capture(), eq(Boolean.TRUE)); + assertThat(captor.getValue()).isEmpty(); + } + + @Test + void onResponseDoesNotCallOutputConsumerForUnknownEventType() throws Exception { + final var client = build(outputConsumerMock, false); + final var event = MAPPER.readTree("{\"type\":\"session.created\"}"); + + client.onResponse("session.created", event); + + verifyNoInteractions(outputConsumerMock); + } + + @Test + void sessionConfigurationHasAudioOutputModality() { + final var realtimeSession = + build(outputConsumerMock, false) + .sessionConfiguration() + .session() + .realtimeSessionCreateRequest() + .orElseThrow(); + + assertThat(realtimeSession.outputModalities().orElseThrow()) + .containsExactly(RealtimeSessionCreateRequest.OutputModality.AUDIO); + } + + @Test + void sessionConfigurationOutputFormatIsPcm24000() { + final var output = + build(outputConsumerMock, false) + .sessionConfiguration() + .session() + .realtimeSessionCreateRequest() + .orElseThrow() + .audio() + .orElseThrow() + .output() + .orElseThrow(); + final var format = output.format().orElseThrow().audioPcm().orElseThrow(); + + assertThat(format.type().orElseThrow()).isEqualTo(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM); + assertThat(format.rate().orElseThrow()).isEqualTo(RealtimeAudioFormats.AudioPcm.Rate._24000); + } + + @Test + void sessionConfigurationOutputVoiceIsMarinByDefault() { + final var voice = + build(outputConsumerMock, false) + .sessionConfiguration() + .session() + .realtimeSessionCreateRequest() + .orElseThrow() + .audio() + .orElseThrow() + .output() + .orElseThrow() + .voice() + .orElseThrow() + .unionMember1() + .orElseThrow(); + + assertThat(voice).isEqualTo(RealtimeAudioConfigOutput.Voice.UnionMember1.MARIN); + } + + @Test + void sessionConfigurationOutputVoiceIsMarinWhenDefault1Provided() { + final var voice = + build(outputConsumerMock, false, RealtimeParamVoice.DEFAULT_1) + .sessionConfiguration() + .session() + .realtimeSessionCreateRequest() + .orElseThrow() + .audio() + .orElseThrow() + .output() + .orElseThrow() + .voice() + .orElseThrow() + .unionMember1() + .orElseThrow(); + + assertThat(voice).isEqualTo(RealtimeAudioConfigOutput.Voice.UnionMember1.MARIN); + } + + @Test + void sessionConfigurationInputConfigDelegatestoInputConfig() { + final var inputConfig = + build(outputConsumerMock, false) + .sessionConfiguration() + .session() + .realtimeSessionCreateRequest() + .orElseThrow() + .audio() + .orElseThrow() + .input() + .orElseThrow(); + + assertThat(inputConfig).isEqualTo(FIXED_INPUT_CONFIG); + } +} diff --git a/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/WSOpenAiRealtimeClientUnitTest.java b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/WSOpenAiRealtimeClientUnitTest.java new file mode 100644 index 000000000..42769523b --- /dev/null +++ b/foundation-models/openai/src/test/java/com/sap/ai/sdk/foundationmodels/openai/realtime/WSOpenAiRealtimeClientUnitTest.java @@ -0,0 +1,454 @@ +package com.sap.ai.sdk.foundationmodels.openai.realtime; + +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.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.anyLong; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.atLeastOnce; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.openai.models.realtime.RealtimeAudioConfig; +import com.openai.models.realtime.RealtimeAudioConfigInput; +import com.openai.models.realtime.RealtimeAudioConfigOutput; +import com.openai.models.realtime.RealtimeAudioFormats; +import com.openai.models.realtime.RealtimeAudioInputTurnDetection; +import com.openai.models.realtime.RealtimeSessionCreateRequest; +import com.openai.models.realtime.SessionUpdateEvent; +import com.openai.models.realtime.clientsecrets.ClientSecretCreateParams; +import com.sap.ai.sdk.core.common.ClientException; +import java.net.http.HttpClient; +import java.net.http.WebSocket; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Optional; +import java.util.Set; +import java.util.Timer; +import java.util.UUID; +import java.util.concurrent.CompletableFuture; +import javax.annotation.Nonnull; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; + +class WSOpenAiRealtimeClientUnitTest { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private List onResponseEventTypes; + private List onResponseEvents; + private SessionUpdateEvent expectedSessionUpdateEvent; + private WebSocket webSocketMock; + + /** + * Concrete testable subclass. It also owns a {@link BufferedWebSocketListener} wired to the same + * text-routing logic (including the handleMessageTypes filter) so we can simulate incoming + * WebSocket frames without a real server. + */ + private class TestableClient extends WSOpenAiRealtimeClient { + + private final Set handledTypes; + private final String systemPrompt; + + /** + * The listener that mirrors what the production URL-based constructor wires up, letting tests + * inject incoming text frames directly. + */ + final BufferedWebSocketListener inboundListener; + + TestableClient( + final CompletableFuture ws, + final Timer timer, + final Set handleMessageTypes, + final String systemPrompt) { + super(mock(HttpClient.class), ws, timer, handleMessageTypes); + this.handledTypes = handleMessageTypes; + this.systemPrompt = systemPrompt; + this.inboundListener = + new BufferedWebSocketListener( + ignored -> {}, // onSocketOpen — not under test here + this::dispatchText); + } + + /** + * Replicates {@code WSOpenAiRealtimeClient.onText}: parse JSON, filter by type, call {@link + * #onResponse}. This is intentionally a thin copy so we can drive routing tests from the test + * package without coupling to the private method. + */ + private void dispatchText(final WebSocket ws, final CharSequence data) { + final JsonNode event; + try { + event = MAPPER.readTree(data.toString()); + } catch (final Exception e) { + throw new ClientException("Error parsing JSON response from speech API", e); + } + final var eventType = event.get("type").asText(); + if (handledTypes.contains(eventType)) { + onResponse(eventType, event); + } + ws.request(1); + } + + @Override + protected synchronized void onResponse( + @Nonnull final String eventType, @Nonnull final JsonNode event) { + onResponseEventTypes.add(eventType); + onResponseEvents.add(event); + } + + @Override + protected synchronized @Nonnull SessionUpdateEvent sessionConfiguration() { + return expectedSessionUpdateEvent; + } + + @Override + @Nonnull + protected Optional getSystemPrompt() { + return systemPrompt.isEmpty() ? Optional.empty() : Optional.of(systemPrompt); + } + + // Expose protected methods for direct invocation in tests + void invokeOnText(final WebSocket ws, final CharSequence data) { + onText(ws, data); + } + + void invokeSendPing(final WebSocket ws) { + sendPing(ws); + } + } + + @BeforeEach + void setUp() { + expectedSessionUpdateEvent = buildSessionUpdateEvent(); + onResponseEventTypes = new ArrayList<>(); + onResponseEvents = new ArrayList<>(); + + webSocketMock = mock(WebSocket.class); + when(webSocketMock.sendText(any(), anyBoolean())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.sendClose(anyInt(), anyString())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.sendPing(any())) + .thenReturn(CompletableFuture.completedFuture(webSocketMock)); + when(webSocketMock.isInputClosed()).thenReturn(false); + } + + private TestableClient buildInstance(final Set handleMessageTypes) { + return new TestableClient( + CompletableFuture.completedFuture(webSocketMock), + mock(Timer.class), + handleMessageTypes, + ""); + } + + private TestableClient buildInstance(final Timer timer, final String systemPrompt) { + return new TestableClient( + CompletableFuture.completedFuture(webSocketMock), timer, Set.of(), systemPrompt); + } + + private TestableClient buildBrokenInstance() { + final var failedFuture = new CompletableFuture(); + failedFuture.completeExceptionally(new RuntimeException("connection refused")); + return new TestableClient(failedFuture, mock(Timer.class), Set.of(), ""); + } + + @Test + void sendMessageSerializesObjectAsJsonAndSendsOverWebSocket() { + buildInstance(Set.of()).sendMessage(expectedSessionUpdateEvent); + + final var captor = ArgumentCaptor.forClass(CharSequence.class); + verify(webSocketMock, atLeastOnce()).sendText(captor.capture(), eq(true)); + final var lastJson = + parseJson(captor.getAllValues().get(captor.getAllValues().size() - 1).toString()); + assertThat(lastJson.get("type").asText()).isEqualTo("session.update"); + verify(webSocketMock, atLeastOnce()).request(anyLong()); + } + + @Test + void sendMessageThrowsClientExceptionWhenConnectionFailed() { + assertThatThrownBy(() -> buildBrokenInstance().sendMessage("anything")) + .isInstanceOf(ClientException.class) + .hasMessageContaining("Failed to establish web socket connection"); + } + + @Test + void askForResponseSendsResponseCreateEvent() { + buildInstance(Set.of()).askForResponse(); + + final var captor = ArgumentCaptor.forClass(CharSequence.class); + verify(webSocketMock, atLeastOnce()).sendText(captor.capture(), eq(true)); + final var sentResponseCreate = + captor.getAllValues().stream() + .map(cs -> parseJson(cs.toString())) + .anyMatch(n -> "response.create".equals(n.get("type").asText())); + assertThat(sentResponseCreate).isTrue(); + verify(webSocketMock, atLeastOnce()).request(anyLong()); + } + + @Test + void askForResponseThrowsClientExceptionWhenConnectionFailed() { + assertThatThrownBy(buildBrokenInstance()::askForResponse) + .isInstanceOf(ClientException.class) + .hasMessageContaining("Failed to establish web socket connection"); + } + + @Test + void closeSendsCloseFrameWithNormalClosureCode() { + buildInstance(Set.of()).close(); + + verify(webSocketMock).sendClose(eq(1000), anyString()); + } + + @Test + void closeThrowsClientExceptionWhenConnectionFailed() { + assertThatThrownBy(buildBrokenInstance()::close) + .isInstanceOf(ClientException.class) + .hasMessageContaining("Failed to establish web socket connection"); + } + + @Test + void incomingMessageWithRegisteredTypeIsRoutedToOnResponse() { + final var instance = buildInstance(Set.of("response.audio.delta", "session.created")); + + instance.inboundListener.onText( + webSocketMock, "{\"type\":\"response.audio.delta\",\"delta\":\"abc\"}", true); + + assertThat(onResponseEventTypes).containsExactly("response.audio.delta"); + assertThat(onResponseEvents).hasSize(1); + assertThat(onResponseEvents.get(0).get("delta").asText()).isEqualTo("abc"); + } + + @Test + void incomingMessageWithUnregisteredTypeIsIgnored() { + final var instance = buildInstance(Set.of("session.created")); + + instance.inboundListener.onText(webSocketMock, "{\"type\":\"some.other.event\"}", true); + + assertThat(onResponseEventTypes).isEmpty(); + assertThat(onResponseEvents).isEmpty(); + } + + @Test + void multipleMatchingIncomingMessagesAreAllDelivered() { + final var instance = buildInstance(Set.of("response.audio.delta")); + + instance.inboundListener.onText( + webSocketMock, "{\"type\":\"response.audio.delta\",\"delta\":\"chunk1\"}", true); + instance.inboundListener.onText( + webSocketMock, "{\"type\":\"response.audio.delta\",\"delta\":\"chunk2\"}", true); + + assertThat(onResponseEventTypes) + .containsExactly("response.audio.delta", "response.audio.delta"); + assertThat(onResponseEvents.get(0).get("delta").asText()).isEqualTo("chunk1"); + assertThat(onResponseEvents.get(1).get("delta").asText()).isEqualTo("chunk2"); + } + + @Test + void bufferedListenerAssemblesPartialFramesBeforeDelivery() { + final var instance = buildInstance(Set.of("session.created")); + + instance.inboundListener.onText(webSocketMock, "{\"type\":\"session", false); + assertThat(onResponseEventTypes).isEmpty(); + + instance.inboundListener.onText(webSocketMock, ".created\",\"id\":\"xyz\"}", true); + assertThat(onResponseEventTypes).containsExactly("session.created"); + } + + @Test + void incomingMalformedJsonThrowsClientException() { + final var instance = buildInstance(Set.of()); + + assertThatThrownBy( + () -> instance.inboundListener.onText(webSocketMock, "not-valid-json{{{", true)) + .isInstanceOf(ClientException.class) + .hasMessageContaining("Error parsing JSON response from speech API"); + } + + @Test + void onSocketOpenSendsSessionConfigurationToWebSocket() { + final var timerMock = mock(Timer.class); + final var instance = buildInstance(timerMock, ""); + + instance.onSocketOpen(webSocketMock); + + final var captor = ArgumentCaptor.forClass(CharSequence.class); + verify(webSocketMock, atLeastOnce()).sendText(captor.capture(), eq(true)); + final var sentTypes = + captor.getAllValues().stream() + .map(cs -> parseJson(cs.toString()).get("type").asText()) + .toList(); + assertThat(sentTypes).contains("session.update"); + } + + @Test + void onSocketOpenSkipsConversationItemWhenSystemPromptIsEmpty() { + final var timerMock = mock(Timer.class); + final var instance = buildInstance(timerMock, ""); + + instance.onSocketOpen(webSocketMock); + + final var captor = ArgumentCaptor.forClass(CharSequence.class); + verify(webSocketMock, atLeastOnce()).sendText(captor.capture(), eq(true)); + final var sentTypes = + captor.getAllValues().stream() + .map(cs -> parseJson(cs.toString()).get("type").asText()) + .toList(); + assertThat(sentTypes).doesNotContain("conversation.item.create"); + } + + @Test + void onSocketOpenSendsSystemPromptAsConversationItemWhenPresent() { + final var timerMock = mock(Timer.class); + final var systemPrompt = "You are a helpful assistant."; + final var instance = buildInstance(timerMock, systemPrompt); + + instance.onSocketOpen(webSocketMock); + + final var captor = ArgumentCaptor.forClass(CharSequence.class); + verify(webSocketMock, atLeastOnce()).sendText(captor.capture(), eq(true)); + final var conversationItemCreate = + captor.getAllValues().stream() + .map(cs -> parseJson(cs.toString())) + .filter(n -> "conversation.item.create".equals(n.get("type").asText())) + .findFirst(); + assertThat(conversationItemCreate).isPresent(); + final var contentText = conversationItemCreate.get().at("/item/content/0/text").asText(); + assertThat(contentText).isEqualTo(systemPrompt); + } + + @Test + void onSocketOpenSchedulesHeartbeatTimer() { + final var timerMock = mock(Timer.class); + final var instance = buildInstance(timerMock, ""); + + instance.onSocketOpen(webSocketMock); + + verify(timerMock).scheduleAtFixedRate(any(java.util.TimerTask.class), eq(0L), eq(4500L)); + } + + @Test + void onTextRoutesRegisteredEventTypeToOnResponse() { + final var instance = buildInstance(Set.of("session.created")); + + instance.invokeOnText(webSocketMock, "{\"type\":\"session.created\",\"id\":\"xyz\"}"); + + assertThat(onResponseEventTypes).containsExactly("session.created"); + assertThat(onResponseEvents.get(0).get("id").asText()).isEqualTo("xyz"); + } + + @Test + void onTextDoesNotCallOnResponseForUnregisteredType() { + final var instance = buildInstance(Set.of("session.created")); + + instance.invokeOnText(webSocketMock, "{\"type\":\"response.audio.delta\"}"); + + assertThat(onResponseEventTypes).isEmpty(); + } + + @Test + void onTextAlwaysRequestsNextMessageRegardlessOfType() { + final var instance = buildInstance(Set.of()); + + instance.invokeOnText(webSocketMock, "{\"type\":\"any.event\"}"); + + verify(webSocketMock).request(1L); + } + + @Test + void onTextThrowsClientExceptionForMalformedJson() { + final var instance = buildInstance(Set.of()); + + assertThatThrownBy(() -> instance.invokeOnText(webSocketMock, "not-json{{")) + .isInstanceOf(ClientException.class) + .hasMessageContaining("Error parsing JSON response from speech API"); + } + + @Test + void sendPingSendsPingBytesAndRequestsNextMessageWhenInputOpen() { + when(webSocketMock.isInputClosed()).thenReturn(false); + final var instance = buildInstance(Set.of()); + + instance.invokeSendPing(webSocketMock); + + final var captor = ArgumentCaptor.forClass(java.nio.ByteBuffer.class); + verify(webSocketMock).sendPing(captor.capture()); + assertThat(new String(captor.getValue().array(), java.nio.charset.StandardCharsets.UTF_8)) + .isEqualTo("ping"); + verify(webSocketMock).request(1L); + } + + @Test + void sendPingDoesNothingWhenInputIsClosed() { + when(webSocketMock.isInputClosed()).thenReturn(true); + final var instance = buildInstance(Set.of()); + + instance.invokeSendPing(webSocketMock); + + verify(webSocketMock, never()).sendPing(any()); + verify(webSocketMock, never()).request(anyLong()); + } + + @Test + void getSystemPromptReturnsEmptyByDefault() { + assertThat(buildInstance(Set.of()).getSystemPrompt()).isEqualTo(Optional.empty()); + } + + @Test + void sessionConfigurationReturnsTheConfiguredEvent() { + assertThat(buildInstance(Set.of()).sessionConfiguration()).isSameAs(expectedSessionUpdateEvent); + } + + private JsonNode parseJson(final String json) { + try { + return MAPPER.readTree(json); + } catch (final Exception e) { + throw new RuntimeException("Failed to parse JSON in test: " + json, e); + } + } + + private SessionUpdateEvent buildSessionUpdateEvent() { + final var input = + RealtimeAudioConfigInput.builder() + .turnDetection( + RealtimeAudioInputTurnDetection.ofSemanticVad( + RealtimeAudioInputTurnDetection.SemanticVad.builder().build())) + .format( + RealtimeAudioFormats.AudioPcm.builder() + .type(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM) + .rate(RealtimeAudioFormats.AudioPcm.Rate._24000) + .build()) + .build(); + + final var output = + RealtimeAudioConfigOutput.builder() + .format( + RealtimeAudioFormats.AudioPcm.builder() + .type(RealtimeAudioFormats.AudioPcm.Type.AUDIO_PCM) + .rate(RealtimeAudioFormats.AudioPcm.Rate._24000) + .build()) + .voice(UUID.randomUUID().toString()) + .build(); + + return SessionUpdateEvent.builder() + .session( + ClientSecretCreateParams.Session.ofRealtime( + RealtimeSessionCreateRequest.builder() + .outputModalities( + Arrays.asList( + RealtimeSessionCreateRequest.OutputModality.AUDIO, + RealtimeSessionCreateRequest.OutputModality.TEXT)) + .audio(RealtimeAudioConfig.builder().input(input).output(output).build()) + .build()) + .asRealtime()) + .build(); + } +} diff --git a/pom.xml b/pom.xml index 715dbb290..d3e7290c6 100644 --- a/pom.xml +++ b/pom.xml @@ -66,6 +66,7 @@ 2.1.3 3.5.6 1.1.8 + 4.41.0 3.8.6 3.2.0 5.23.0 @@ -101,6 +102,7 @@ /tmp/baseline.jar ${project.build.directory}/${project.build.finalName}.jar + 1.9.10 @@ -167,6 +169,22 @@ httpcore5 ${httpcomponents-core5.version} + + Pinned to align all kotlin-stdlib transitive versions pulled via openai-java and okio<--> + org.jetbrains.kotlin + kotlin-stdlib-jdk8 + ${kotlin.stdlib.jdk8.version} + + + org.jetbrains.kotlin + kotlin-stdlib + ${kotlin.stdlib.jdk8.version} + + + org.jetbrains.kotlin + kotlin-stdlib-common + ${kotlin.stdlib.jdk8.version} + io.micrometer micrometer-core @@ -274,11 +292,6 @@ openai-java-core ${openai-java.version} - - com.openai - openai-java - ${openai-java.version} - diff --git a/sample-code/spring-app/pom.xml b/sample-code/spring-app/pom.xml index 78c5b57ac..ac97990da 100644 --- a/sample-code/spring-app/pom.xml +++ b/sample-code/spring-app/pom.xml @@ -139,6 +139,10 @@ org.springframework.ai spring-ai-commons + + org.springframework + spring-websocket + org.springframework.ai spring-ai-model diff --git a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/WebsocketConfig.java b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/WebsocketConfig.java new file mode 100644 index 000000000..0b8a56d22 --- /dev/null +++ b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/WebsocketConfig.java @@ -0,0 +1,42 @@ +package com.sap.ai.sdk.app; + +import com.sap.ai.sdk.app.realtime.SpeechToSpeechWebsocketHandler; +import com.sap.ai.sdk.app.realtime.TextToSpeechWebsocketHandler; +import javax.annotation.Nonnull; +import org.springframework.context.annotation.Configuration; +import org.springframework.web.socket.config.annotation.EnableWebSocket; +import org.springframework.web.socket.config.annotation.WebSocketConfigurer; +import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry; + +/** Implements spring Web Socket configuration to expose Web Socket handlers for Realtime API */ +@Configuration +@EnableWebSocket +public class WebsocketConfig implements WebSocketConfigurer { + + private final TextToSpeechWebsocketHandler textToSpeech; + private final SpeechToSpeechWebsocketHandler speechToSpeech; + + /** + * Constructs configuration object + * + * @param textToSpeech - text to speech realtime api handler + * @param speechToSpeech - speech to speech realtime api handler + */ + public WebsocketConfig( + @Nonnull final TextToSpeechWebsocketHandler textToSpeech, + @Nonnull final SpeechToSpeechWebsocketHandler speechToSpeech) { + this.textToSpeech = textToSpeech; + this.speechToSpeech = speechToSpeech; + } + + /** + * Registers websocket handlers, implements WebSocketConfigurer contract + * + * @param registry - registry where to register handlers + */ + @Override + public void registerWebSocketHandlers(@Nonnull final WebSocketHandlerRegistry registry) { + registry.addHandler(textToSpeech, "/text-to-speech"); + registry.addHandler(speechToSpeech, "/speech-to-speech"); + } +} diff --git a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/controllers/OpenAiController.java b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/controllers/OpenAiController.java index a15031ddb..9d15231e5 100644 --- a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/controllers/OpenAiController.java +++ b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/controllers/OpenAiController.java @@ -4,14 +4,24 @@ import com.fasterxml.jackson.annotation.PropertyAccessor; import com.fasterxml.jackson.databind.ObjectMapper; import com.sap.ai.sdk.app.services.OpenAiService; +import com.sap.ai.sdk.core.common.ClientException; +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; import com.sap.ai.sdk.foundationmodels.openai.generated.model.CompletionUsage; +import com.sap.ai.sdk.foundationmodels.openai.realtime.RealtimeParamTurnDetection; import com.sap.cloud.sdk.cloudplatform.thread.ThreadContextExecutors; import java.io.IOException; +import java.nio.ByteBuffer; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; import javax.annotation.Nonnull; import javax.annotation.Nullable; import lombok.extern.slf4j.Slf4j; import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.core.io.ClassPathResource; +import org.springframework.core.io.Resource; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; import org.springframework.web.bind.annotation.GetMapping; @@ -30,6 +40,8 @@ public class OpenAiController { private static final ObjectMapper MAPPER = new ObjectMapper().setVisibility(PropertyAccessor.FIELD, JsonAutoDetect.Visibility.ANY); + private Resource sampleQuestionPcm = new ClassPathResource("static/question.pcm"); + @GetMapping("/chatCompletion") @Nonnull Object chatCompletion( @@ -41,6 +53,71 @@ Object chatCompletion( return response.getContent(); } + @GetMapping(value = "/realtime/smokeTestTextToSpeech", produces = "audio/pcm") + @Nonnull + ResponseEntity smokeTestTextToSpeech() throws TimeoutException { + final var respBody = ByteBuffer.allocate(300000); + final var lock = new AtomicBoolean(false); + final AudioOutputChannel audioOutput = + (final byte[] pcmBytes, final boolean isLast) -> { + respBody.put(pcmBytes); + if (isLast) { + lock.set(true); + } + }; + + try (TextInputChannel channel = service.textToSpeech(audioOutput)) { + channel.sendText("Hello, how are you today?"); + final var started = System.currentTimeMillis(); + while (!lock.get() && System.currentTimeMillis() - started < 15000) { + if (lock.get()) { + return ResponseEntity.ok(respBody.array()); + } + Thread.sleep(1000); + } + } catch (final Exception e) { + throw new ClientException("Failure occurred during communication with the server", e); + } + + if (lock.get()) { + return ResponseEntity.ok(respBody.array()); + } + throw new TimeoutException("Timeout waiting for text to speech"); + } + + @GetMapping(value = "/realtime/smokeTestSpeechToSpeech", produces = "audio/pcm") + @Nonnull + ResponseEntity smokeTestSpeechToSpeech() throws TimeoutException { + final var respBody = ByteBuffer.allocate(800000); + final var lock = new AtomicBoolean(false); + final AudioOutputChannel audioOutput = + (final byte[] pcmBytes, final boolean isLast) -> { + respBody.put(pcmBytes); + if (isLast) { + lock.set(true); + } + }; + + try (AudioInputChannel channel = + service.speechToSpeech(audioOutput, RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN)) { + channel.inputAudio(sampleQuestionPcm.getContentAsByteArray()); + final var started = System.currentTimeMillis(); + while (!lock.get() && System.currentTimeMillis() - started < 30000) { + if (lock.get()) { + return ResponseEntity.ok(respBody.array()); + } + Thread.sleep(1000); + } + } catch (final Exception e) { + throw new ClientException("Failure occurred during communication with the server", e); + } + + if (lock.get()) { + return ResponseEntity.ok(respBody.array()); + } + throw new TimeoutException("Timeout waiting for speech to speech"); + } + @SuppressWarnings("unused") // The end-to-end test doesn't use this method @GetMapping("/streamChatCompletionDeltas") @Nonnull diff --git a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/realtime/SpeechToSpeechWebsocketHandler.java b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/realtime/SpeechToSpeechWebsocketHandler.java new file mode 100644 index 000000000..3a43b191d --- /dev/null +++ b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/realtime/SpeechToSpeechWebsocketHandler.java @@ -0,0 +1,75 @@ +package com.sap.ai.sdk.app.realtime; + +import com.sap.ai.sdk.app.services.OpenAiService; +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import java.io.IOException; +import java.nio.ByteBuffer; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import javax.annotation.Nonnull; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.stereotype.Component; +import org.springframework.web.socket.BinaryMessage; +import org.springframework.web.socket.CloseStatus; +import org.springframework.web.socket.WebSocketSession; +import org.springframework.web.socket.handler.BinaryWebSocketHandler; + +/** Implements handler (Web Socket messages handling) for speech to speech realtime api operation */ +@Component +@Slf4j +public class SpeechToSpeechWebsocketHandler extends BinaryWebSocketHandler { + + private final OpenAiService service; + private final Map channels; + + /** + * Constructs handler object + * + * @param service - handling service + */ + @Autowired + public SpeechToSpeechWebsocketHandler(@Nonnull final OpenAiService service) { + this.service = service; + channels = new ConcurrentHashMap<>(); + } + + @Override + // The channel MUST NOT be closed here, its lifecycle is managed by the WebSocket container (RAII) + // closing performed in afterConnectionClosed method + @SuppressWarnings("PMD.CloseResource") + protected void handleBinaryMessage( + @Nonnull final WebSocketSession session, @Nonnull final BinaryMessage message) { + final ByteBuffer payload = message.getPayload(); + final byte[] chunkBytes = payload.array(); + final AudioInputChannel channel = + channels.computeIfAbsent( + session.getId(), + sessionId -> + service.speechToSpeech( + (rawBytesChunk, isLast) -> { + try { + session.sendMessage(new BinaryMessage(rawBytesChunk, isLast)); + } catch (final IOException e) { + log.error("Failed to send audio data to realtime api", e); + } + })); + channel.inputAudio(chunkBytes); + } + + @Override + public void afterConnectionClosed( + @Nonnull final WebSocketSession session, @Nonnull final CloseStatus status) throws Exception { + channels.computeIfPresent( + session.getId(), + (sessionId, inputChannel) -> { + try { + inputChannel.close(); + } catch (final Exception e) { + log.warn("Failed to close input channel for session {}", sessionId, e); + } + return null; + }); + super.afterConnectionClosed(session, status); + } +} diff --git a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/realtime/TextToSpeechWebsocketHandler.java b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/realtime/TextToSpeechWebsocketHandler.java new file mode 100644 index 000000000..c15f13a03 --- /dev/null +++ b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/realtime/TextToSpeechWebsocketHandler.java @@ -0,0 +1,76 @@ +package com.sap.ai.sdk.app.realtime; + +import com.sap.ai.sdk.app.services.OpenAiService; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; +import java.io.IOException; +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import javax.annotation.Nonnull; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.stereotype.Component; +import org.springframework.web.socket.BinaryMessage; +import org.springframework.web.socket.CloseStatus; +import org.springframework.web.socket.WebSocketSession; +import org.springframework.web.socket.handler.BinaryWebSocketHandler; + +/** Implements handler (Web Socket messages handling) for text to speech realtime api operation */ +@Component +@Slf4j +public class TextToSpeechWebsocketHandler extends BinaryWebSocketHandler { + + private final OpenAiService service; + private final Map channels; + + /** + * Constructs handler object + * + * @param service - handling service + */ + @Autowired + public TextToSpeechWebsocketHandler(@Nonnull final OpenAiService service) { + this.service = service; + channels = new ConcurrentHashMap<>(); + } + + @Override + // The channel MUST NOT be closed here, its lifecycle is managed by the WebSocket container (RAII) + // closing performed in afterConnectionClosed method + @SuppressWarnings("PMD.CloseResource") + protected void handleBinaryMessage( + @Nonnull final WebSocketSession session, @Nonnull final BinaryMessage message) { + final ByteBuffer payload = message.getPayload(); + final byte[] textBytes = payload.array(); + final TextInputChannel channel = + channels.computeIfAbsent( + session.getId(), + sessionId -> + service.textToSpeech( + (rawBytesChunk, isLast) -> { + try { + session.sendMessage(new BinaryMessage(rawBytesChunk, isLast)); + } catch (final IOException e) { + log.error("Failed to send text message to realtime api", e); + } + })); + channel.sendText(new String(textBytes, StandardCharsets.UTF_8)); + } + + @Override + public void afterConnectionClosed( + @Nonnull final WebSocketSession session, @Nonnull final CloseStatus status) throws Exception { + channels.computeIfPresent( + session.getId(), + (sessionId, inputChannel) -> { + try { + inputChannel.close(); + } catch (Exception e) { + log.warn("Failed to close input channel for session {}", sessionId, e); + } + return null; + }); + super.afterConnectionClosed(session, status); + } +} diff --git a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/services/OpenAiService.java b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/services/OpenAiService.java index 261cf9a9c..a3a9bf508 100644 --- a/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/services/OpenAiService.java +++ b/sample-code/spring-app/src/main/java/com/sap/ai/sdk/app/services/OpenAiService.java @@ -5,6 +5,8 @@ import static com.sap.ai.sdk.foundationmodels.openai.OpenAiModel.TEXT_EMBEDDING_3_SMALL; import com.sap.ai.sdk.core.AiCoreService; +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.AudioOutputChannel; import com.sap.ai.sdk.foundationmodels.openai.OpenAiChatCompletionDelta; import com.sap.ai.sdk.foundationmodels.openai.OpenAiChatCompletionRequest; import com.sap.ai.sdk.foundationmodels.openai.OpenAiChatCompletionResponse; @@ -14,6 +16,8 @@ import com.sap.ai.sdk.foundationmodels.openai.OpenAiImageItem; import com.sap.ai.sdk.foundationmodels.openai.OpenAiMessage; import com.sap.ai.sdk.foundationmodels.openai.OpenAiTool; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.realtime.RealtimeParam; import java.util.ArrayList; import java.util.List; import java.util.stream.Stream; @@ -74,6 +78,69 @@ public Stream streamChatCompletionDeltas( return OpenAiClient.forModel(GPT_5_MINI).streamChatCompletionDeltas(request); } + /** + * Creates realtime channel allowing to input text and voice it (receive audio output) + * + *

The input channel should be used with a try-with-resources block to ensure that the + * underlying connection is closed. + * + *

Example: + * + *

{@code
+   * try (var textInputChannel = client.textToSpeech(audioOutputConsumer)) {
+   *       textInputChannel.sendText("...");
+   *       ....
+   * }
+   * }
+ * + * This API implements full duplex (input + output) communication channels. Application should + * logically synchronize their state and close input channel when it is appropriate (e.g. last + * part of the response has been received via output channel and application does not need to send + * any other input). When input channel is closed, output channel will be closed automatically and + * output consumer will not be called anymore. + * + * @param audioOutputConsumer - audio consumer of raw PCM mono 24000 Hz little endian output + * @return input channel, allowing for text input + */ + @Nonnull + public TextInputChannel textToSpeech(@Nonnull final AudioOutputChannel audioOutputConsumer) { + return OpenAiClient.realtimeClient().textToSpeech(audioOutputConsumer); + } + + /** + * Creates realtime channel allowing for audio conversation with a model + * + *

The input channel should be used with a try-with-resources block to ensure that the + * underlying connection is closed. + * + *

Example: + * + *

{@code
+   * try (var audioInputChannel = client.speechToSpeech(audioOutputConsumer)) {
+   *       audioInputChannel.inputAudio(audioBytesData);
+   *       ....
+   * }
+   * }
+ * + * This API implements full duplex (input + output) communication channels. Application should + * logically synchronize their state and close input channel when it is appropriate (e.g. last + * part of the response has been received via output channel and application does not need to send + * any other input). When input channel is closed, output channel will be closed automatically and + * output consumer will not be called anymore. + * + * @param audioOutputConsumer - audio consumer of raw PCM mono 24000 Hz little endian output, 16 + * bit depth + * @param realtimeParams - optional additional configuration params + * @return input channel, allowing for audio data input (bytes, PCM mono 24000 Hz little endian 16 + * bit) + */ + @Nonnull + public AudioInputChannel speechToSpeech( + @Nonnull final AudioOutputChannel audioOutputConsumer, + @Nonnull final RealtimeParam... realtimeParams) { + return OpenAiClient.realtimeClient().speechToSpeech(audioOutputConsumer, realtimeParams); + } + /** * Asynchronous stream of an OpenAI chat request * diff --git a/sample-code/spring-app/src/main/resources/static/index.html b/sample-code/spring-app/src/main/resources/static/index.html index 45324e555..6c1e668f8 100644 --- a/sample-code/spring-app/src/main/resources/static/index.html +++ b/sample-code/spring-app/src/main/resources/static/index.html @@ -82,6 +82,78 @@ max-height: 500px; overflow-y: auto; } + + .modal { + display: none; /* Hidden by default */ + position: fixed; /* Stay in place */ + z-index: 1; /* Sit on top */ + padding-top: 100px; /* Location of the box */ + left: 0; + top: 0; + width: 100%; /* Full width */ + height: 100%; /* Full height */ + overflow: auto; /* Enable scroll if needed */ + background-color: rgb(0,0,0); /* Fallback color */ + background-color: rgba(0,0,0,0.9); /* Black w/ opacity */ + } + + .modal-content { + margin: auto; + display: block; + width: 80%; + max-width: 700px; + } + + #modal-caption { + margin: auto; + display: block; + width: 80%; + max-width: 700px; + text-align: center; + color: #ccc; + padding: 10px 0; + height: 150px; + } + + .modal-content, #modal-caption { + -webkit-animation-name: zoom; + -webkit-animation-duration: 0.6s; + animation-name: zoom; + animation-duration: 0.6s; + } + + @-webkit-keyframes zoom { + from {-webkit-transform:scale(0)} + to {-webkit-transform:scale(1)} + } + + @keyframes zoom { + from {transform:scale(0)} + to {transform:scale(1)} + } + + .modal-close-button { + position: absolute; + top: 15px; + right: 35px; + color: #f1f1f1; + font-size: 40px; + font-weight: bold; + transition: 0.3s; + } + + .modal-close-button:hover, + .modal-close-button:focus { + color: #bbb; + text-decoration: none; + cursor: pointer; + } + + @media only screen and (max-width: 700px){ + .modal-content { + width: 100%; + } + } @@ -1481,6 +1553,63 @@

📂 Batch API

+ +
+
+
+
+

⏰ Realtime API

+
+ Realtime API allows for various real time interactions ('Voice agents' and 'speech generation' scenarios are currently supported) +
+ In these examples, web socket sessions are created and used + to send text/audio to the local server and receive audio back +
+ + + +
+
+
    +
  • +
    +

    Text to speech (smoke-test)

    + disabled +
    + + + +
    +
    + Translate input text into output sound (speech). +
    +
    + +
  • +
  • +
    +

    Speech to speech (smoke-test)

    +
    socket: disconnected
    +
    microphone: disabled
    +
    + +
    +
    + Speak with an AI assistant. +
    +
    + +
  • +
+
+
+
diff --git a/sample-code/spring-app/src/main/resources/static/question.pcm b/sample-code/spring-app/src/main/resources/static/question.pcm new file mode 100644 index 000000000..d0707295b Binary files /dev/null and b/sample-code/spring-app/src/main/resources/static/question.pcm differ diff --git a/sample-code/spring-app/src/main/resources/static/realtime-api-scheme.svg b/sample-code/spring-app/src/main/resources/static/realtime-api-scheme.svg new file mode 100644 index 000000000..85729e6ba --- /dev/null +++ b/sample-code/spring-app/src/main/resources/static/realtime-api-scheme.svg @@ -0,0 +1,3 @@ + + +
User Browser
Javascript WebSockets clients
Actor
Transmit audio (TX)
Microphone
Browser
Transmit text (TX)
Keyboard
Browser
Receive audio (RX)
Speakers
Browser
Sample App
Spring DI
@Configuration WebsocketConfig
Binds handlers to ws:// paths
- /speech-to-speech
- /text-to-speech
@Component
TextToSpeechWebsocketHandler
- Encodes/Decodes WS packets
- Routes Sample App server traffic to AI SDK input channel and back 
@Component
SpeechToSpeechWebsocketHandler
- Encodes/Decodes WS packets
- Routes Sample App server traffic to AI SDK input channel and back 
@Service
OpenAiService
- Resolves host address
- Delegates WSS connection initiation to OpenAiRealtimeClient
SAP AI SDK
OpenAiRealtimeClient
- Initializes appropriate client (Text/Audio)
- Provides implementation for usage
WebSockets (TCP)
ws://localhost:8080
Transmit audio (TX)
Receive audio (RX)
Receive audio (RX)
Transmit text (TX)
Secure WebSockets (TCP)
wss://${API_ADDRESS}
Transmit data messages (TX)
Receive data messages (RX)
SAP Realtime API
\ No newline at end of file diff --git a/sample-code/spring-app/src/main/resources/static/show-realtime-diagram.js b/sample-code/spring-app/src/main/resources/static/show-realtime-diagram.js new file mode 100644 index 000000000..65eaa11b3 --- /dev/null +++ b/sample-code/spring-app/src/main/resources/static/show-realtime-diagram.js @@ -0,0 +1,20 @@ +var modal = document.getElementById("show-realtime-scheme"); + +// Get the image and insert it inside the modal - use its "alt" text as a caption +var btn = document.getElementById("show-arch-scheme-btn"); +var modalImg = document.getElementById("realtime-scheme"); +var captionText = document.getElementById("modal-caption"); +btn.onclick = function(){ + modal.style.display = "block"; + modalImg.src = "realtime-api-scheme.svg"; + captionText.innerHTML = "Principal interactions scheme"; +} +btn.disabled = false; + +// Get the element that closes the modal +var span = document.getElementsByClassName("modal-close-button")[0]; + +// When the user clicks on (x), close the modal +span.onclick = function() { + modal.style.display = "none"; +} \ No newline at end of file diff --git a/sample-code/spring-app/src/main/resources/static/speech-to-speech.js b/sample-code/spring-app/src/main/resources/static/speech-to-speech.js new file mode 100644 index 000000000..ff25001e0 --- /dev/null +++ b/sample-code/spring-app/src/main/resources/static/speech-to-speech.js @@ -0,0 +1,128 @@ +const SPEECH_TO_SPEECH_URL = 'ws://localhost:8080/speech-to-speech'; +const SPEECH_TO_SPEECH_SAMPLE_RATE = 24000; + +const wsStatusEl = document.getElementById('speech-to-speech-websocket-status') +const micStatusEl = document.getElementById('speech-to-speech-mic-status') +const speechBtn = document.getElementById('speech-to-speech-btn') + +let sts_ws = null; +let sts_audioCtx = null; +let sts_mediaStream = null; +let sts_workletNode = null; +let sts_nextStartTime = 0; +let sts_started = false; + +async function startSession() { + sts_ws = new WebSocket(SPEECH_TO_SPEECH_URL); + sts_ws.binaryType = 'arraybuffer'; + + sts_ws.onopen = async () => { + wsStatusEl.textContent = `socket: connected to ${SPEECH_TO_SPEECH_URL}`; + wsStatusEl.style.color = 'green'; + + await startMicrophone(); + speechBtn.innerText = 'Stop'; + sts_started = true; + } + + sts_ws.onmessage = (event) => { + if (event.data instanceof ArrayBuffer) { + sts_playPcmAudio(event.data) + } + } + + sts_ws.onclose = () => stopSession(); +} + +async function startMicrophone() { + try { + sts_mediaStream = await navigator.mediaDevices.getUserMedia({audio: true}); + console.log("mediaStream is: ", sts_mediaStream) + micStatusEl.textContent = 'microphone: active'; + micStatusEl.style.color = 'green'; + + sts_audioCtx = new (window.AudioContext || window.webkitAudioContext)({sampleRate: SPEECH_TO_SPEECH_SAMPLE_RATE}); + + const workletCode = ` + class MicProcessor extends AudioWorkletProcessor { + process(inputs) { + const input = inputs[0]; + if (input && input[0]) { + const float32Input = input[0]; + const int16Buffer = new Int16Array(float32Input.length); + for (let i = 0; i < float32Input.length; i++) { + const s = Math.max(-1, Math.min(1, float32Input[i])); + int16Buffer[i] = s < 0 ? s * 0x8000 : s * 0x7FFF; + } + this.port.postMessage(int16Buffer.buffer, [int16Buffer.buffer]); + } + return true; + } + } + registerProcessor('mic-processor', MicProcessor); + `; + const blob = new Blob([workletCode], {type: 'application/javascript'}); + const workletUrl = URL.createObjectURL(blob); + await sts_audioCtx.audioWorklet.addModule(workletUrl); + URL.revokeObjectURL(workletUrl); + + const source = sts_audioCtx.createMediaStreamSource(sts_mediaStream); + sts_workletNode = new AudioWorkletNode(sts_audioCtx, 'mic-processor'); + sts_workletNode.port.onmessage = (e) => { + if (sts_ws && sts_ws.readyState === WebSocket.OPEN) { + sts_ws.send(e.data); + } + }; + + source.connect(sts_workletNode); + + } catch (err) { + console.error('failed to find or bind microphone: ', err); + micStatusEl.innerText = 'mic binding error'; + micStatusEl.style.color = 'red'; + } +} + +function sts_playPcmAudio(arrayBuffer) { + if (!sts_audioCtx) return; + + const int16Data = new Int16Array(arrayBuffer); + const float32Data = new Float32Array(int16Data.length); + for (let i = 0; i < int16Data.length; i++) { + float32Data[i] = int16Data[i] / 32768.0; + } + + const audioBuffer = sts_audioCtx.createBuffer(1, float32Data.length, SPEECH_TO_SPEECH_SAMPLE_RATE); + audioBuffer.copyToChannel(float32Data, 0); + const source = sts_audioCtx.createBufferSource(); + source.buffer = audioBuffer; + source.connect(sts_audioCtx.destination); + + if (sts_nextStartTime < sts_audioCtx.currentTime) { + sts_nextStartTime = sts_audioCtx.currentTime; + } + + source.start(sts_nextStartTime); + sts_nextStartTime += audioBuffer.duration; +} + +function stopSession() { + if (sts_ws) sts_ws.close(); + if (sts_mediaStream) sts_mediaStream.getTracks().forEach(track => track.stop()); + if (sts_workletNode) sts_workletNode.disconnect(); + if (sts_audioCtx) { + sts_audioCtx.close(); + sts_audioCtx = null; + } + + wsStatusEl.textContent = 'socket: disconnected'; + wsStatusEl.style.color = 'red'; + micStatusEl.textContent = 'microphone: disabled'; + micStatusEl.style.color = 'red'; + speechBtn.innerText = 'Start'; + sts_started = false; +} + +speechBtn.addEventListener('click', () => { + sts_started ? stopSession() : startSession(); +}) \ No newline at end of file diff --git a/sample-code/spring-app/src/main/resources/static/text-to-speech.js b/sample-code/spring-app/src/main/resources/static/text-to-speech.js new file mode 100644 index 000000000..80bd75841 --- /dev/null +++ b/sample-code/spring-app/src/main/resources/static/text-to-speech.js @@ -0,0 +1,77 @@ +const WS_URL = 'ws://localhost:8080/text-to-speech'; +const SAMPLE_RATE = 24000; + +const statusEl = document.getElementById('text-to-speech-status') +const textInput = document.getElementById('text-to-speech-input') +const sendBtn = document.getElementById('text-to-speech-send-btn') + +let audioCtx; +let nextStartTime = 0; + +let ws = new WebSocket(WS_URL); +ws.binaryType = 'arraybuffer'; + +ws.onopen = () => { + statusEl.textContent = `connected to ${WS_URL}`; + statusEl.style.color = 'green'; + textInput.disabled = false; + sendBtn.disabled = false; +}; + +ws.onclose = () => { + statusEl.textContent = 'disconnected'; + statusEl.style.color = 'red'; + textInput.disabled = true; + sendBtn.disabled = true; +} + +ws.onerror = (error) => { + console.error('text-to-speech WebSocket error:', error); + statusEl.textContent = 'connection error' + statusEl.style.color = 'red'; +} + +ws.onmessage = (event) => { + if (event.data instanceof ArrayBuffer) { + playPcmAudio(event.data); + } else { + console.log('text data received:', event.data) + } +} + +sendBtn.addEventListener('click', () => { + const text = textInput.value.trim(); + if (text && ws.readyState === WebSocket.OPEN) { + if (!audioCtx) { + audioCtx = new (window.AudioContext || window.webkitAudioContext)({ sampleRate: SAMPLE_RATE}); + } + + const encoder = new TextEncoder(); + const binaryData = encoder.encode(text); + ws.send(binaryData); + textInput.value = ''; + } +}) + +function playPcmAudio(arrayBuffer) { + if (!audioCtx) return; + + const int16Data = new Int16Array(arrayBuffer); + const float32Data = new Float32Array(int16Data.length); + for (let i = 0; i < int16Data.length; i++) { + float32Data[i] = int16Data[i] / 32768.0; + } + + const audioBuffer = audioCtx.createBuffer(1, float32Data.length, SAMPLE_RATE); + audioBuffer.copyToChannel(float32Data, 0); + const source = audioCtx.createBufferSource(); + source.buffer = audioBuffer; + source.connect(audioCtx.destination); + + if (nextStartTime < audioCtx.currentTime) { + nextStartTime = audioCtx.currentTime; + } + + source.start(nextStartTime); + nextStartTime += audioBuffer.duration; +} \ No newline at end of file diff --git a/sample-code/spring-app/src/test/java/com/sap/ai/sdk/app/controllers/RealtimeApiTest.java b/sample-code/spring-app/src/test/java/com/sap/ai/sdk/app/controllers/RealtimeApiTest.java new file mode 100644 index 000000000..20fa264b4 --- /dev/null +++ b/sample-code/spring-app/src/test/java/com/sap/ai/sdk/app/controllers/RealtimeApiTest.java @@ -0,0 +1,153 @@ +package com.sap.ai.sdk.app.controllers; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.fail; + +import com.sap.ai.sdk.foundationmodels.openai.AudioInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.OpenAiClient; +import com.sap.ai.sdk.foundationmodels.openai.TextInputChannel; +import com.sap.ai.sdk.foundationmodels.openai.realtime.RealtimeParamTurnDetection; +import java.io.ByteArrayOutputStream; +import java.io.FileInputStream; +import java.io.IOException; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import lombok.extern.slf4j.Slf4j; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; + +@Slf4j +public class RealtimeApiTest { + + private static final double LOG_2 = Math.log(2.0); + + private static byte[] QUESTION_FIXTURE_PCM; + + @BeforeAll + public static void setUp() { + try (var fis = new FileInputStream("src/test/resources/fixtures/question.pcm")) { + QUESTION_FIXTURE_PCM = fis.readAllBytes(); + } catch (IOException e) { + fail(e.getMessage()); + } + } + + @Test + @Timeout(value = 30, unit = TimeUnit.SECONDS) + void testTextToSpeech() { + var outputBuffer = new ByteArrayOutputStream(300000); + var monitor = new CountDownLatch(1); + var client = OpenAiClient.realtimeClient(); + + try (TextInputChannel input = + client.textToSpeech( + (byteChunk, isLast) -> { + outputBuffer.writeBytes(byteChunk); + if (isLast) { + monitor.countDown(); + } + }, + RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN)) { + input.sendText("Ordnung muss sein!"); + monitor.await(); + } catch (Exception e) { + if (!(e instanceof InterruptedException)) { + fail(e); + } + // do nothing, test has either been interrupted by user (intended) or by jupiter if timeout is + // reached + return; + } + assertThat(monitor.getCount()).isEqualTo(0); + assertThat(outputBuffer.size()).isGreaterThan(0); + + var metrics = pcm16AudioMetrics(outputBuffer.toByteArray()); + + // root mean squire (measures the deviation from an average value) + assertThat(metrics.rms).isGreaterThan(500d); + // asserts that variety of deviation is sufficient (not a trivial repeating pattern) + assertThat(metrics.entropy).isGreaterThan(4); + } + + @Test + @Timeout(value = 60, unit = TimeUnit.SECONDS) + void testSpeechToSpeech() { + var outputBuffer = new ByteArrayOutputStream(300000); + var monitor = new CountDownLatch(1); + var client = OpenAiClient.realtimeClient(); + + try (AudioInputChannel input = + client.speechToSpeech( + (byteChunk, isLast) -> { + outputBuffer.writeBytes(byteChunk); + if (isLast) { + monitor.countDown(); + } + }, + RealtimeParamTurnDetection.EACH_CALL_IS_A_TURN)) { + input.inputAudio(QUESTION_FIXTURE_PCM); + monitor.await(); + } catch (Exception e) { + if (!(e instanceof InterruptedException)) { + fail(e); + } + // do nothing, test has either been interrupted by user (intended) or by jupiter if timeout is + // reached + return; + } + + assertThat(monitor.getCount()).isEqualTo(0); + assertThat(outputBuffer.size()).isGreaterThan(0); + + var metrics = pcm16AudioMetrics(outputBuffer.toByteArray()); + + // root mean squire (measures the deviation from an average value) + assertThat(metrics.rms).isGreaterThan(500d); + // asserts that variety of deviation is sufficient (not a trivial repeating pattern) + assertThat(metrics.entropy).isGreaterThan(4); + } + + private record AudioMetrics(double rms, double entropy) {} + ; + + /** + * Audio quality metrics for signed 16-bit little-endian PCM, computed on the mean-subtracted + * (DC-offset-removed) signal so they reflect only the varying, audible part of the audio. + * + * @param pcm - The raw PCM audio buffer. + * @return The RMS amplitude and Shannon entropy (bits) of the zero-mean samples. + */ + AudioMetrics pcm16AudioMetrics(byte[] pcm) { + if (pcm.length < 1) { + return new AudioMetrics(0, 0); + } + var sum = 0L; + for (var i = 0; i < pcm.length; i += 2) { + sum += (pcm[i]) | ((pcm[i + 1]) << 8); + } + var mean = sum / (pcm.length / 2); + + var sumOfSquares = 0d; + var counts = new char[Character.MAX_VALUE]; + for (var i = 0; i < pcm.length; i += 2) { + var residual = ((pcm[i]) | ((pcm[i + 1]) << 8)) - mean; + sumOfSquares += residual * residual; + var key = (int) (residual + Short.MAX_VALUE + 1); + counts[key]++; + } + + var rms = Math.sqrt(sumOfSquares / (double) (pcm.length / 2)); + var entropy = 0d; + for (var i = 0; i < counts.length; i++) { + var p = counts[i] / counts.length; + entropy -= p * log2(p); + } + + return new AudioMetrics(rms, entropy); + } + + private static double log2(double logNumber) { + return Math.log(logNumber) / LOG_2; + } +} diff --git a/sample-code/spring-app/src/test/resources/fixtures/question.pcm b/sample-code/spring-app/src/test/resources/fixtures/question.pcm new file mode 100644 index 000000000..d0707295b Binary files /dev/null and b/sample-code/spring-app/src/test/resources/fixtures/question.pcm differ