From b54bb8bc8a56a3a9e6fc5241b87c433fbadf25fc Mon Sep 17 00:00:00 2001 From: Jemin Huh Date: Sun, 17 Dec 2023 22:02:41 +0900 Subject: [PATCH] OpenAi clien improvments - OpenAiApi's ObjectMapper is configured to DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES = false - Use Java Predicate for the SSE_DONE condition used by the takeUntil and filter for the Flux stream. - Add retryTemplate to OpenAI Chat Client. Make OpenAiApi toleratant to uknown response fields. --- .../ai/openai/api/OpenAiApi.java | 14 +- .../ai/openai/client/OpenAiClient.java | 132 ++++++++++-------- 2 files changed, 81 insertions(+), 65 deletions(-) diff --git a/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java b/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java index 1cdd2e626..084858aa5 100644 --- a/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java +++ b/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java @@ -20,11 +20,13 @@ import java.io.IOException; import java.util.List; import java.util.Map; import java.util.function.Consumer; +import java.util.function.Predicate; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.core.type.TypeReference; +import com.fasterxml.jackson.databind.DeserializationFeature; import com.fasterxml.jackson.databind.ObjectMapper; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; @@ -51,7 +53,7 @@ public class OpenAiApi { private static final String DEFAULT_BASE_URL = "https://api.openai.com"; private static final String DEFAULT_EMBEDDING_MODEL = "text-embedding-ada-002"; - private static final String SSE_DONE = "[DONE]"; + private static final Predicate SSE_DONE_PREDICATE = "[DONE]"::equals; private final RestClient restClient; private final WebClient webClient; @@ -75,7 +77,7 @@ public class OpenAiApi { */ public OpenAiApi(String baseUrl, String openAiToken, RestClient.Builder restClientBuilder) { - this.objectMapper = new ObjectMapper(); + this.objectMapper = new ObjectMapper().configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); Consumer jsonContentHeaders = headers -> { headers.setBearerAuth(openAiToken); @@ -623,10 +625,10 @@ public class OpenAiApi { .body(Mono.just(chatRequest), ChatCompletionRequest.class) .retrieve() .bodyToFlux(String.class) - // cancels the flux stream after the SSE_DONE is received. - .takeUntil(content -> content.contains(SSE_DONE)) - // filters out the SSE_DONE message. - .filter(content -> !content.contains(SSE_DONE)) + // cancels the flux stream after the "[DONE]" is received. + .takeUntil(SSE_DONE_PREDICATE) + // filters out the "[DONE]" message. + .filter(SSE_DONE_PREDICATE.negate()) .map(content -> parseJson(content, ChatCompletionChunk.class)); } diff --git a/spring-ai-openai/src/main/java/org/springframework/ai/openai/client/OpenAiClient.java b/spring-ai-openai/src/main/java/org/springframework/ai/openai/client/OpenAiClient.java index 604e20ff2..858ac4776 100644 --- a/spring-ai-openai/src/main/java/org/springframework/ai/openai/client/OpenAiClient.java +++ b/spring-ai-openai/src/main/java/org/springframework/ai/openai/client/OpenAiClient.java @@ -15,6 +15,7 @@ */ package org.springframework.ai.openai.client; +import java.time.Duration; import java.util.List; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; @@ -32,11 +33,13 @@ import org.springframework.ai.metadata.RateLimit; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage; +import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiException; import org.springframework.ai.openai.metadata.OpenAiGenerationMetadata; import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor; import org.springframework.ai.prompt.Prompt; import org.springframework.ai.prompt.messages.Message; import org.springframework.http.ResponseEntity; +import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; /** @@ -60,6 +63,12 @@ public class OpenAiClient implements AiClient, AiStreamClient { private final Logger logger = LoggerFactory.getLogger(getClass()); + public final RetryTemplate retryTemplate = RetryTemplate.builder() + .maxAttempts(10) + .retryOn(OpenAiApiException.class) + .exponentialBackoff(Duration.ofMillis(2000), 5, Duration.ofMillis(3 * 60000)) + .build(); + private final OpenAiApi openAiApi; public OpenAiClient(OpenAiApi openAiApi) { @@ -86,81 +95,86 @@ public class OpenAiClient implements AiClient, AiStreamClient { @Override public AiResponse generate(Prompt prompt) { - List messages = prompt.getMessages(); + return this.retryTemplate.execute(ctx -> { + List messages = prompt.getMessages(); - List chatCompletionMessages = messages.stream() - .map(m -> new ChatCompletionMessage(m.getContent(), - ChatCompletionMessage.Role.valueOf(m.getMessageType().getValue()))) - .toList(); + List chatCompletionMessages = messages.stream() + .map(m -> new ChatCompletionMessage(m.getContent(), + ChatCompletionMessage.Role.valueOf(m.getMessageType().getValue()))) + .toList(); - ResponseEntity completionEntity = this.openAiApi.chatCompletionEntity( - new OpenAiApi.ChatCompletionRequest(chatCompletionMessages, this.model, this.temperature.floatValue())); + ResponseEntity completionEntity = this.openAiApi + .chatCompletionEntity(new OpenAiApi.ChatCompletionRequest(chatCompletionMessages, this.model, + this.temperature.floatValue())); - var chatCompletion = completionEntity.getBody(); - if (chatCompletion == null) { - logger.warn("No chat completion returned for request: {}", chatCompletionMessages); - return new AiResponse(List.of()); - } + var chatCompletion = completionEntity.getBody(); + if (chatCompletion == null) { + logger.warn("No chat completion returned for request: {}", chatCompletionMessages); + return new AiResponse(List.of()); + } - RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); + RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); - List generations = chatCompletion.choices().stream().map(choice -> { - return new Generation(choice.message().content(), Map.of("role", choice.message().role().name())) - .withChoiceMetadata(ChoiceMetadata.from(choice.finishReason().name(), null)); - }).toList(); + List generations = chatCompletion.choices().stream().map(choice -> { + return new Generation(choice.message().content(), Map.of("role", choice.message().role().name())) + .withChoiceMetadata(ChoiceMetadata.from(choice.finishReason().name(), null)); + }).toList(); - return new AiResponse(generations, - OpenAiGenerationMetadata.from(completionEntity.getBody()).withRateLimit(rateLimits)); + return new AiResponse(generations, + OpenAiGenerationMetadata.from(completionEntity.getBody()).withRateLimit(rateLimits)); + }); } @Override public Flux generateStream(Prompt prompt) { - List messages = prompt.getMessages(); + return this.retryTemplate.execute(ctx -> { + List messages = prompt.getMessages(); - List chatCompletionMessages = messages.stream() - .map(m -> new ChatCompletionMessage(m.getContent(), - ChatCompletionMessage.Role.valueOf(m.getMessageType().getValue()))) - .toList(); + List chatCompletionMessages = messages.stream() + .map(m -> new ChatCompletionMessage(m.getContent(), + ChatCompletionMessage.Role.valueOf(m.getMessageType().getValue()))) + .toList(); - Flux completionChunks = this.openAiApi - .chatCompletionStream(new OpenAiApi.ChatCompletionRequest(chatCompletionMessages, this.model, - this.temperature.floatValue(), true)); + Flux completionChunks = this.openAiApi + .chatCompletionStream(new OpenAiApi.ChatCompletionRequest(chatCompletionMessages, this.model, + this.temperature.floatValue(), true)); - // For chunked responses, only the first chunk contains the choice role. - // The rest of the chunks with same ID share the same role. - ConcurrentHashMap roleMap = new ConcurrentHashMap<>(); + // For chunked responses, only the first chunk contains the choice role. + // The rest of the chunks with same ID share the same role. + ConcurrentHashMap roleMap = new ConcurrentHashMap<>(); - // An alternative implementation that returns Flux instead of - // Flux. - // Flux generationFlux = completionChunks.map(chunk -> { - // String chunkId = chunk.id(); - // return chunk.choices().stream() - // .map(choice -> { - // if (choice.delta().role() != null) { - // roleMap.putIfAbsent(chunkId, choice.delta().role().name()); - // } - // return new Generation(choice.delta().content(), - // Map.of("role", roleMap.get(chunkId))); - // }) - // .toList(); - // }).flatMapIterable(generations -> generations); - // return generationFlux; + // An alternative implementation that returns Flux instead of + // Flux. + // Flux generationFlux = completionChunks.map(chunk -> { + // String chunkId = chunk.id(); + // return chunk.choices().stream() + // .map(choice -> { + // if (choice.delta().role() != null) { + // roleMap.putIfAbsent(chunkId, choice.delta().role().name()); + // } + // return new Generation(choice.delta().content(), + // Map.of("role", roleMap.get(chunkId))); + // }) + // .toList(); + // }).flatMapIterable(generations -> generations); + // return generationFlux; - return completionChunks.map(chunk -> { - String chunkId = chunk.id(); - List generations = chunk.choices().stream().map(choice -> { - if (choice.delta().role() != null) { - roleMap.putIfAbsent(chunkId, choice.delta().role().name()); - } - var generation = new Generation(choice.delta().content(), Map.of("role", roleMap.get(chunkId))); - if (choice.finishReason() != null) { - generation = generation.withChoiceMetadata(ChoiceMetadata.from(choice.finishReason().name(), null)); - } - return generation; - }).toList(); - return new AiResponse(generations); + return completionChunks.map(chunk -> { + String chunkId = chunk.id(); + List generations = chunk.choices().stream().map(choice -> { + if (choice.delta().role() != null) { + roleMap.putIfAbsent(chunkId, choice.delta().role().name()); + } + var generation = new Generation(choice.delta().content(), Map.of("role", roleMap.get(chunkId))); + if (choice.finishReason() != null) { + generation = generation + .withChoiceMetadata(ChoiceMetadata.from(choice.finishReason().name(), null)); + } + return generation; + }).toList(); + return new AiResponse(generations); + }); }); - } }