diff --git a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/ChatClientMessageAggregator.java b/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/ChatClientMessageAggregator.java new file mode 100644 index 000000000..582d77b48 --- /dev/null +++ b/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/ChatClientMessageAggregator.java @@ -0,0 +1,60 @@ +/* + * Copyright 2023-2025 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.ai.chat.client; + +import java.util.HashMap; +import java.util.Map; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; + +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import reactor.core.publisher.Flux; + +import org.springframework.ai.chat.model.MessageAggregator; + +/** + * Helper that for streaming chat responses, aggregate the chat response messages into a + * single AssistantMessage. Job is performed in parallel to the chat response processing. + * + * @author Christian Tzolov + * @author Alexandros Pappas + * @author Thomas Vitale + * @since 1.0.0 + */ +public class ChatClientMessageAggregator { + + private static final Logger logger = LoggerFactory.getLogger(ChatClientMessageAggregator.class); + + public Flux aggregateChatClientResponse(Flux chatClientResponses, + Consumer aggregationHandler) { + + AtomicReference> context = new AtomicReference<>(new HashMap<>()); + + return new MessageAggregator().aggregate(chatClientResponses.mapNotNull(chatClientResponse -> { + context.get().putAll(chatClientResponse.context()); + return chatClientResponse.chatResponse(); + }), aggregatedChatResponse -> { + ChatClientResponse aggregatedChatClientResponse = ChatClientResponse.builder() + .chatResponse(aggregatedChatResponse) + .context(context.get()) + .build(); + aggregationHandler.accept(aggregatedChatClientResponse); + }).map(chatResponse -> ChatClientResponse.builder().chatResponse(chatResponse).context(context.get()).build()); + } + +} diff --git a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/PromptChatMemoryAdvisor.java b/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/PromptChatMemoryAdvisor.java index 7c1344895..b30c8b197 100644 --- a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/PromptChatMemoryAdvisor.java +++ b/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/PromptChatMemoryAdvisor.java @@ -28,6 +28,7 @@ import reactor.core.publisher.Mono; import reactor.core.scheduler.Scheduler; import reactor.core.scheduler.Schedulers; +import org.springframework.ai.chat.client.ChatClientMessageAggregator; import org.springframework.ai.chat.client.ChatClientRequest; import org.springframework.ai.chat.client.ChatClientResponse; import org.springframework.ai.chat.client.advisor.api.Advisor; @@ -40,7 +41,6 @@ import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; -import org.springframework.ai.chat.model.MessageAggregator; import org.springframework.ai.chat.prompt.PromptTemplate; /** @@ -172,7 +172,7 @@ public class PromptChatMemoryAdvisor implements BaseChatMemoryAdvisor { .publishOn(scheduler) .map(request -> this.before(request, streamAdvisorChain)) .flatMapMany(streamAdvisorChain::nextStream) - .transform(flux -> new MessageAggregator().aggregateChatClientResponse(flux, + .transform(flux -> new ChatClientMessageAggregator().aggregateChatClientResponse(flux, response -> this.after(response, streamAdvisorChain))); } diff --git a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/SimpleLoggerAdvisor.java b/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/SimpleLoggerAdvisor.java index e5000a73a..0160e3155 100644 --- a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/SimpleLoggerAdvisor.java +++ b/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/advisor/SimpleLoggerAdvisor.java @@ -18,10 +18,11 @@ package org.springframework.ai.chat.client.advisor; import java.util.function.Function; -import reactor.core.publisher.Flux; - import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import reactor.core.publisher.Flux; + +import org.springframework.ai.chat.client.ChatClientMessageAggregator; import org.springframework.ai.chat.client.ChatClientRequest; import org.springframework.ai.chat.client.ChatClientResponse; import org.springframework.ai.chat.client.advisor.api.CallAdvisor; @@ -29,7 +30,6 @@ import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain; import org.springframework.ai.chat.client.advisor.api.StreamAdvisor; import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain; import org.springframework.ai.chat.model.ChatResponse; -import org.springframework.ai.chat.model.MessageAggregator; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.lang.Nullable; @@ -85,7 +85,7 @@ public class SimpleLoggerAdvisor implements CallAdvisor, StreamAdvisor { Flux chatClientResponses = streamAdvisorChain.nextStream(chatClientRequest); - return new MessageAggregator().aggregateChatClientResponse(chatClientResponses, this::logResponse); + return new ChatClientMessageAggregator().aggregateChatClientResponse(chatClientResponses, this::logResponse); } private void logRequest(ChatClientRequest request) { diff --git a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java b/spring-ai-model/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java similarity index 88% rename from spring-ai-client-chat/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java rename to spring-ai-model/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java index 433050cb3..839d99e23 100644 --- a/spring-ai-client-chat/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java +++ b/spring-ai-model/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java @@ -24,7 +24,6 @@ import java.util.function.Consumer; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.client.ChatClientResponse; import reactor.core.publisher.Flux; import org.springframework.ai.chat.messages.AssistantMessage; @@ -49,23 +48,6 @@ public class MessageAggregator { private static final Logger logger = LoggerFactory.getLogger(MessageAggregator.class); - public Flux aggregateChatClientResponse(Flux chatClientResponses, - Consumer aggregationHandler) { - - AtomicReference> context = new AtomicReference<>(new HashMap<>()); - - return new MessageAggregator().aggregate(chatClientResponses.mapNotNull(chatClientResponse -> { - context.get().putAll(chatClientResponse.context()); - return chatClientResponse.chatResponse(); - }), aggregatedChatResponse -> { - ChatClientResponse aggregatedChatClientResponse = ChatClientResponse.builder() - .chatResponse(aggregatedChatResponse) - .context(context.get()) - .build(); - aggregationHandler.accept(aggregatedChatClientResponse); - }).map(chatResponse -> ChatClientResponse.builder().chatResponse(chatResponse).context(context.get()).build()); - } - public Flux aggregate(Flux fluxChatResponse, Consumer onAggregationComplete) {