From 54e5c07428909ceec248e3bbd71e2df4b0812e49 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Tue, 13 May 2025 12:07:40 +0200 Subject: [PATCH] refactor: Move MessageAggregator to spring-ai-model module Improve separation of concerns by keeping model-related functionality in the model module while maintaining client-specific functionality in the client module. - Move MessageAggregator class from spring-ai-client-chat to spring-ai-model module - Create new ChatClientMessageAggregator in spring-ai-client-chat module to handle client-specific aggregation - Extract client-specific aggregation logic from MessageAggregator to ChatClientMessageAggregator - Update references in advisor classes to use the new ChatClientMessageAggregator Signed-off-by: Christian Tzolov --- .../client/ChatClientMessageAggregator.java | 60 +++++++++++++++++++ .../advisor/PromptChatMemoryAdvisor.java | 4 +- .../client/advisor/SimpleLoggerAdvisor.java | 8 +-- .../ai/chat/model/MessageAggregator.java | 18 ------ 4 files changed, 66 insertions(+), 24 deletions(-) create mode 100644 spring-ai-client-chat/src/main/java/org/springframework/ai/chat/client/ChatClientMessageAggregator.java rename {spring-ai-client-chat => spring-ai-model}/src/main/java/org/springframework/ai/chat/model/MessageAggregator.java (88%) 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) {