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 <christian.tzolov@broadcom.com>
This commit is contained in:
Christian Tzolov
2025-05-13 12:07:40 +02:00
committed by Ilayaperumal Gopinathan
parent a03e7cba2f
commit 54e5c07428
4 changed files with 66 additions and 24 deletions

View File

@@ -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<ChatClientResponse> aggregateChatClientResponse(Flux<ChatClientResponse> chatClientResponses,
Consumer<ChatClientResponse> aggregationHandler) {
AtomicReference<Map<String, Object>> 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());
}
}

View File

@@ -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)));
}

View File

@@ -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<ChatClientResponse> chatClientResponses = streamAdvisorChain.nextStream(chatClientRequest);
return new MessageAggregator().aggregateChatClientResponse(chatClientResponses, this::logResponse);
return new ChatClientMessageAggregator().aggregateChatClientResponse(chatClientResponses, this::logResponse);
}
private void logRequest(ChatClientRequest request) {

View File

@@ -1,198 +0,0 @@
/*
* 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.model;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.concurrent.atomic.AtomicReference;
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;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.EmptyRateLimit;
import org.springframework.ai.chat.metadata.PromptMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.StringUtils;
/**
* 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 MessageAggregator {
private static final Logger logger = LoggerFactory.getLogger(MessageAggregator.class);
public Flux<ChatClientResponse> aggregateChatClientResponse(Flux<ChatClientResponse> chatClientResponses,
Consumer<ChatClientResponse> aggregationHandler) {
AtomicReference<Map<String, Object>> 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<ChatResponse> aggregate(Flux<ChatResponse> fluxChatResponse,
Consumer<ChatResponse> onAggregationComplete) {
// Assistant Message
AtomicReference<StringBuilder> messageTextContentRef = new AtomicReference<>(new StringBuilder());
AtomicReference<Map<String, Object>> messageMetadataMapRef = new AtomicReference<>();
// ChatGeneration Metadata
AtomicReference<ChatGenerationMetadata> generationMetadataRef = new AtomicReference<>(
ChatGenerationMetadata.NULL);
// Usage
AtomicReference<Integer> metadataUsagePromptTokensRef = new AtomicReference<Integer>(0);
AtomicReference<Integer> metadataUsageGenerationTokensRef = new AtomicReference<Integer>(0);
AtomicReference<Integer> metadataUsageTotalTokensRef = new AtomicReference<Integer>(0);
AtomicReference<PromptMetadata> metadataPromptMetadataRef = new AtomicReference<>(PromptMetadata.empty());
AtomicReference<RateLimit> metadataRateLimitRef = new AtomicReference<>(new EmptyRateLimit());
AtomicReference<String> metadataIdRef = new AtomicReference<>("");
AtomicReference<String> metadataModelRef = new AtomicReference<>("");
return fluxChatResponse.doOnSubscribe(subscription -> {
messageTextContentRef.set(new StringBuilder());
messageMetadataMapRef.set(new HashMap<>());
metadataIdRef.set("");
metadataModelRef.set("");
metadataUsagePromptTokensRef.set(0);
metadataUsageGenerationTokensRef.set(0);
metadataUsageTotalTokensRef.set(0);
metadataPromptMetadataRef.set(PromptMetadata.empty());
metadataRateLimitRef.set(new EmptyRateLimit());
}).doOnNext(chatResponse -> {
if (chatResponse.getResult() != null) {
if (chatResponse.getResult().getMetadata() != null
&& chatResponse.getResult().getMetadata() != ChatGenerationMetadata.NULL) {
generationMetadataRef.set(chatResponse.getResult().getMetadata());
}
if (chatResponse.getResult().getOutput().getText() != null) {
messageTextContentRef.get().append(chatResponse.getResult().getOutput().getText());
}
if (chatResponse.getResult().getOutput().getMetadata() != null) {
messageMetadataMapRef.get().putAll(chatResponse.getResult().getOutput().getMetadata());
}
}
if (chatResponse.getMetadata() != null) {
if (chatResponse.getMetadata().getUsage() != null) {
Usage usage = chatResponse.getMetadata().getUsage();
metadataUsagePromptTokensRef.set(
usage.getPromptTokens() > 0 ? usage.getPromptTokens() : metadataUsagePromptTokensRef.get());
metadataUsageGenerationTokensRef.set(usage.getCompletionTokens() > 0 ? usage.getCompletionTokens()
: metadataUsageGenerationTokensRef.get());
metadataUsageTotalTokensRef
.set(usage.getTotalTokens() > 0 ? usage.getTotalTokens() : metadataUsageTotalTokensRef.get());
}
if (chatResponse.getMetadata().getPromptMetadata() != null
&& chatResponse.getMetadata().getPromptMetadata().iterator().hasNext()) {
metadataPromptMetadataRef.set(chatResponse.getMetadata().getPromptMetadata());
}
if (chatResponse.getMetadata().getRateLimit() != null
&& !(metadataRateLimitRef.get() instanceof EmptyRateLimit)) {
metadataRateLimitRef.set(chatResponse.getMetadata().getRateLimit());
}
if (StringUtils.hasText(chatResponse.getMetadata().getId())) {
metadataIdRef.set(chatResponse.getMetadata().getId());
}
if (StringUtils.hasText(chatResponse.getMetadata().getModel())) {
metadataModelRef.set(chatResponse.getMetadata().getModel());
}
}
}).doOnComplete(() -> {
var usage = new DefaultUsage(metadataUsagePromptTokensRef.get(), metadataUsageGenerationTokensRef.get(),
metadataUsageTotalTokensRef.get());
var chatResponseMetadata = ChatResponseMetadata.builder()
.id(metadataIdRef.get())
.model(metadataModelRef.get())
.rateLimit(metadataRateLimitRef.get())
.usage(usage)
.promptMetadata(metadataPromptMetadataRef.get())
.build();
onAggregationComplete.accept(new ChatResponse(List.of(new Generation(
new AssistantMessage(messageTextContentRef.get().toString(), messageMetadataMapRef.get()),
generationMetadataRef.get())), chatResponseMetadata));
messageTextContentRef.set(new StringBuilder());
messageMetadataMapRef.set(new HashMap<>());
metadataIdRef.set("");
metadataModelRef.set("");
metadataUsagePromptTokensRef.set(0);
metadataUsageGenerationTokensRef.set(0);
metadataUsageTotalTokensRef.set(0);
metadataPromptMetadataRef.set(PromptMetadata.empty());
metadataRateLimitRef.set(new EmptyRateLimit());
}).doOnError(e -> logger.error("Aggregation Error", e));
}
public record DefaultUsage(Integer promptTokens, Integer completionTokens, Integer totalTokens) implements Usage {
@Override
public Integer getPromptTokens() {
return promptTokens();
}
@Override
public Integer getCompletionTokens() {
return completionTokens();
}
@Override
public Integer getTotalTokens() {
return totalTokens();
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", promptTokens());
usage.put("completionTokens", completionTokens());
usage.put("totalTokens", totalTokens());
return usage;
}
}
}