From 5bad0c42161dae2cf40783305aa7cecbb0cb62b4 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Sun, 11 May 2025 06:58:39 -0400 Subject: [PATCH] fix failing test --- .../advisor/PromptChatMemoryAdvisor.java | 49 +++++++++---------- 1 file changed, 24 insertions(+), 25 deletions(-) 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 71943dc73..705ab6adf 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 @@ -35,7 +35,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; /** @@ -111,15 +110,36 @@ public class PromptChatMemoryAdvisor extends AbstractConversationHistoryAdvisor streamAdvisorChain, this::before); // Ensure memory is updated after each streamed response - return chatClientResponses.doOnNext(this::after) - .transform(responses -> new MessageAggregator().aggregateChatClientResponse(responses, null)); + return chatClientResponses.doOnNext(this::after); } @Override protected ChatClientRequest before(ChatClientRequest chatClientRequest) { String conversationId = this.doGetConversationId(chatClientRequest.context()); - // 1. Add all user messages from the current prompt to memory + // 1. Retrieve the chat memory for the current conversation. + List memoryMessages = this.getChatMemoryStore().get(conversationId); + logger.debug("[PromptChatMemoryAdvisor.before] Memory before processing for conversationId={}: {}", + conversationId, memoryMessages); + + // 2. Process memory messages as a string. + String memory = memoryMessages.stream() + .filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT) + .map(m -> m.getMessageType() + ":" + m.getText()) + .collect(Collectors.joining(System.lineSeparator())); + + // 3. Augment the system message. + SystemMessage systemMessage = chatClientRequest.prompt().getSystemMessage(); + String augmentedSystemText = this.systemPromptTemplate + .render(Map.of("instructions", systemMessage.getText(), "memory", memory)); + + // 4. Create a new request with the augmented system message. + ChatClientRequest processedChatClientRequest = chatClientRequest.mutate() + .prompt(chatClientRequest.prompt().augmentSystemMessage(augmentedSystemText)) + .build(); + + // 5. Add all user messages from the current prompt to memory (after system + // message is generated) List userMessages = chatClientRequest.prompt().getUserMessages(); for (UserMessage userMessage : userMessages) { this.getChatMemoryStore().add(conversationId, userMessage); @@ -127,27 +147,6 @@ public class PromptChatMemoryAdvisor extends AbstractConversationHistoryAdvisor conversationId, userMessage.getText()); } - // 2. Retrieve the chat memory for the current conversation. - List memoryMessages = this.getChatMemoryStore().get(conversationId); - logger.debug("[PromptChatMemoryAdvisor.before] Memory after USER add for conversationId={}: {}", conversationId, - memoryMessages); - - // 3. Process memory messages as a string. - String memory = memoryMessages.stream() - .filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT) - .map(m -> m.getMessageType() + ":" + m.getText()) - .collect(Collectors.joining(System.lineSeparator())); - - // 4. Augment the system message. - SystemMessage systemMessage = chatClientRequest.prompt().getSystemMessage(); - String augmentedSystemText = this.systemPromptTemplate - .render(Map.of("instructions", systemMessage.getText(), "memory", memory)); - - // 5. Create a new request with the augmented system message. - ChatClientRequest processedChatClientRequest = chatClientRequest.mutate() - .prompt(chatClientRequest.prompt().augmentSystemMessage(augmentedSystemText)) - .build(); - return processedChatClientRequest; }