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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@ void transcribe() {
AudioTranscriptionResponse response = transcriptionModel
.call(new AudioTranscriptionPrompt(new ClassPathResource("/speech.flac")));
assertThat(response.getResults()).hasSize(1);
assertThat(response.getResult().getOutput()).isNotBlank();
assertThat(response.getResult().getOutput().text()).isNotBlank();
});
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.io.InputStream;
import java.time.Duration;
import java.util.ArrayList;
import java.util.List;
import java.util.Objects;
Expand All @@ -28,7 +29,13 @@
import com.openai.core.MultipartField;
import com.openai.models.audio.transcriptions.TranscriptionCreateParams;
import com.openai.models.audio.transcriptions.TranscriptionCreateResponse;
import com.openai.models.audio.transcriptions.TranscriptionDiarized;
import com.openai.models.audio.transcriptions.TranscriptionDiarizedSegment;
import com.openai.models.audio.transcriptions.TranscriptionSegment;
import com.openai.models.audio.transcriptions.TranscriptionStreamEvent;
import com.openai.models.audio.transcriptions.TranscriptionTextSegmentEvent;
import com.openai.models.audio.transcriptions.TranscriptionVerbose;
import com.openai.models.audio.transcriptions.TranscriptionWord;
import io.micrometer.observation.ObservationRegistry;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
Expand All @@ -39,6 +46,9 @@
import org.springframework.ai.audio.transcription.AudioTranscriptionPrompt;
import org.springframework.ai.audio.transcription.AudioTranscriptionResponse;
import org.springframework.ai.audio.transcription.AudioTranscriptionResponseMetadata;
import org.springframework.ai.audio.transcription.AudioTranscriptionResult;
import org.springframework.ai.audio.transcription.AudioTranscriptionSegment;
import org.springframework.ai.audio.transcription.AudioTranscriptionWord;
import org.springframework.ai.audio.transcription.TranscriptionModel;
import org.springframework.ai.openai.http.okhttp.OpenAiHttpClientBuilderCustomizer;
import org.springframework.ai.openai.setup.OpenAiSetup;
Expand Down Expand Up @@ -129,8 +139,7 @@ public AudioTranscriptionResponse call(AudioTranscriptionPrompt transcriptionPro
}

TranscriptionCreateResponse response = this.openAiClient.audio().transcriptions().create(params);
String text = extractText(response);
AudioTranscription transcript = new AudioTranscription(text);
AudioTranscription transcript = new AudioTranscription(toTranscriptionResult(response));
return new AudioTranscriptionResponse(transcript, new AudioTranscriptionResponseMetadata());
}

Expand Down Expand Up @@ -165,9 +174,8 @@ public Flux<AudioTranscriptionResponse> stream(AudioTranscriptionPrompt transcri
}
}));

return chunks.map(event -> {
String text = extractStreamEventText(event);
AudioTranscription transcript = new AudioTranscription(text);
return chunks.filter(event -> !event.isTranscriptTextDone()).map(event -> {
AudioTranscription transcript = new AudioTranscription(toTranscriptionResult(event));
return new AudioTranscriptionResponse(transcript, new AudioTranscriptionResponseMetadata());
});
}
Expand All @@ -178,19 +186,10 @@ private TranscriptionCreateParams buildParams(OpenAiAudioTranscriptionOptions op
.value(new ByteArrayInputStream(audioBytes))
.filename(filename)
.build();
String model;
if (options.getDeploymentName() != null) {
model = options.getDeploymentName();
}
else {
model = options.getModel();
}
Assert.notNull(model, "Model must not be null");
String model = options.getDeploymentName() != null ? options.getDeploymentName() : options.getModel();
TranscriptionCreateParams.Builder builder = TranscriptionCreateParams.builder().file(fileField).model(model);

if (options.getResponseFormat() != null) {
builder.responseFormat(options.getResponseFormat());
}
builder.responseFormat(options.getResponseFormat());
if (options.getLanguage() != null) {
builder.language(options.getLanguage());
}
Expand All @@ -206,33 +205,82 @@ private TranscriptionCreateParams buildParams(OpenAiAudioTranscriptionOptions op
return builder.build();
}

private static String extractText(TranscriptionCreateResponse response) {
private static AudioTranscriptionResult toTranscriptionResult(TranscriptionCreateResponse response) {
if (response.isTranscription()) {
return response.asTranscription().text();
return new AudioTranscriptionResult(response.asTranscription().text());
}
if (response.isVerbose()) {
return response.asVerbose().text();
return toTranscriptionResult(response.asVerbose());
}
if (response.isDiarized()) {
return response.asDiarized().text();
return toTranscriptionResult(response.asDiarized());
}
return "";
return new AudioTranscriptionResult("");
}

private static AudioTranscriptionResult toTranscriptionResult(TranscriptionVerbose transcription) {
List<AudioTranscriptionSegment> segments = transcription.segments()
.orElseGet(List::of)
.stream()
.map(OpenAiAudioTranscriptionModel::toTranscriptionSegment)
.toList();
List<AudioTranscriptionWord> words = transcription.words()
.orElseGet(List::of)
.stream()
.map(OpenAiAudioTranscriptionModel::toTranscriptionWord)
.toList();
return new AudioTranscriptionResult(transcription.text(), transcription.language(),
toDuration(transcription.duration()), segments, words);
}

private static AudioTranscriptionResult toTranscriptionResult(TranscriptionDiarized transcription) {
List<AudioTranscriptionSegment> segments = transcription.segments()
.stream()
.map(OpenAiAudioTranscriptionModel::toTranscriptionSegment)
.toList();
return new AudioTranscriptionResult(transcription.text(), null, toDuration(transcription.duration()), segments,
List.of());
}

private static AudioTranscriptionSegment toTranscriptionSegment(TranscriptionSegment segment) {
return new AudioTranscriptionSegment(Long.toString(segment.id()), null, toDuration(segment.start()),
toDuration(segment.end()), segment.text());
}

private static String extractStreamEventText(TranscriptionStreamEvent event) {
private static AudioTranscriptionSegment toTranscriptionSegment(TranscriptionDiarizedSegment segment) {
return new AudioTranscriptionSegment(segment.id(), segment.speaker(), toDuration(segment.start()),
toDuration(segment.end()), segment.text());
}

private static AudioTranscriptionWord toTranscriptionWord(TranscriptionWord word) {
return new AudioTranscriptionWord(word.word(), toDuration(word.start()), toDuration(word.end()));
}

private static Duration toDuration(double seconds) {
return Duration.ofNanos(Math.round(seconds * 1_000_000_000));
}

private static AudioTranscriptionResult toTranscriptionResult(TranscriptionStreamEvent event) {
if (event.isTranscriptTextDelta()) {
return event.asTranscriptTextDelta().delta();
return new AudioTranscriptionResult(event.asTranscriptTextDelta().delta());
}
if (event.isTranscriptTextSegment()) {
return event.asTranscriptTextSegment().text();
TranscriptionTextSegmentEvent segment = event.asTranscriptTextSegment();
return new AudioTranscriptionResult(segment.text(), null, null, List.of(toTranscriptionSegment(segment)),
List.of());
}
return "";
return new AudioTranscriptionResult("");
}

private static AudioTranscriptionSegment toTranscriptionSegment(TranscriptionTextSegmentEvent segment) {
return new AudioTranscriptionSegment(segment.id(), segment.speaker(), toDuration(segment.start()),
toDuration(segment.end()), segment.text());
}

private static byte[] toBytes(Resource resource) {
Assert.notNull(resource, "Resource must not be null");
try {
return resource.getInputStream().readAllBytes();
try (InputStream inputStream = resource.getInputStream()) {
return inputStream.readAllBytes();
}
catch (IOException e) {
throw new IllegalArgumentException("Failed to read resource: " + resource, e);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import org.springframework.ai.audio.transcription.AudioTranscription;
import org.springframework.ai.audio.transcription.AudioTranscriptionPrompt;
import org.springframework.ai.audio.transcription.AudioTranscriptionResponse;
import org.springframework.ai.audio.transcription.AudioTranscriptionResult;
import org.springframework.ai.openai.OpenAiAudioTranscriptionModel;
import org.springframework.core.io.Resource;

Expand Down Expand Up @@ -53,7 +54,7 @@ void transcrbeRequestReturnsResponseCorrectly() {

// Create a mock Transcript
AudioTranscription transcript = Mockito.mock(AudioTranscription.class);
given(transcript.getOutput()).willReturn(mockTranscription);
given(transcript.getOutput()).willReturn(new AudioTranscriptionResult(mockTranscription));

// Create a mock TranscriptionResponse with the mock Transcript
AudioTranscriptionResponse response = Mockito.mock(AudioTranscriptionResponse.class);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,8 @@ void transcriptionTest() {
transcriptionOptions);
AudioTranscriptionResponse response = this.transcriptionModel.call(transcriptionRequest);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getResults().get(0).getOutput().toLowerCase(Locale.ROOT).contains("fellow")).isTrue();
assertThat(response.getResults().get(0).getOutput().text().toLowerCase(Locale.ROOT).contains("fellow"))
.isTrue();
}

@Test
Expand All @@ -79,7 +80,8 @@ void transcriptionTestWithOptions() {
transcriptionOptions);
AudioTranscriptionResponse response = this.transcriptionModel.call(transcriptionRequest);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getResults().get(0).getOutput().toLowerCase(Locale.ROOT).contains("fellow")).isTrue();
assertThat(response.getResults().get(0).getOutput().text().toLowerCase(Locale.ROOT).contains("fellow"))
.isTrue();
}

@SpringBootConfiguration
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
import org.springframework.ai.audio.transcription.AudioTranscription;
import org.springframework.ai.audio.transcription.AudioTranscriptionPrompt;
import org.springframework.ai.audio.transcription.AudioTranscriptionResponse;
import org.springframework.ai.audio.transcription.AudioTranscriptionResult;
import org.springframework.ai.openai.OpenAiAudioTranscriptionModel;
import org.springframework.ai.openai.OpenAiAudioTranscriptionOptions;
import org.springframework.ai.openai.OpenAiTestConfiguration;
Expand Down Expand Up @@ -60,7 +61,7 @@ void callTest() {
AudioTranscriptionResponse response = this.transcriptionModel.call(prompt);

assertThat(response.getResults()).hasSize(1);
assertThat(response.getResult().getOutput()).isNotBlank();
assertThat(response.getResult().getOutput().text()).isNotBlank();
}

@Test
Expand All @@ -82,7 +83,7 @@ void transcribeWithOptionsTest() {
AudioTranscriptionResponse response = this.transcriptionModel.call(prompt);

assertThat(response.getResults()).hasSize(1);
assertThat(response.getResult().getOutput()).isNotBlank();
assertThat(response.getResult().getOutput().text()).isNotBlank();
}

@Test
Expand Down Expand Up @@ -123,7 +124,7 @@ void callTestWithVttFormat() {
AudioTranscriptionResponse response = this.transcriptionModel.call(prompt);

assertThat(response.getResults()).hasSize(1);
assertThat(response.getResult().getOutput()).isNotBlank();
assertThat(response.getResult().getOutput().text()).isNotBlank();
}

@Test
Expand All @@ -141,6 +142,7 @@ void streamTest() {
.map(AudioTranscriptionResponse::getResult)
.filter(Objects::nonNull)
.map(AudioTranscription::getOutput)
.map(AudioTranscriptionResult::text)
.collect(Collectors.joining());
assertThat(text).isNotBlank();
}
Expand Down
Loading