Fix VectorStoreChatMemoryAdvisor streaming bug
- Override adviseStream method in VectorStoreChatMemoryAdvisor to properly handle streaming responses - Add tests to verify the fix works with both normal and problematic streaming scenarios Fixes #3152 Signed-off-by: Mark Pollack <mark.pollack@broadcom.com>
This commit is contained in:
@@ -21,14 +21,18 @@ import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import reactor.core.publisher.Flux;
|
||||
import reactor.core.publisher.Mono;
|
||||
import reactor.core.scheduler.Scheduler;
|
||||
|
||||
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;
|
||||
import org.springframework.ai.chat.client.advisor.api.AdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.BaseAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.BaseChatMemoryAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.ai.chat.memory.ChatMemory;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
@@ -167,6 +171,20 @@ public final class VectorStoreChatMemoryAdvisor implements BaseChatMemoryAdvisor
|
||||
return chatClientResponse;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest chatClientRequest,
|
||||
StreamAdvisorChain streamAdvisorChain) {
|
||||
// Get the scheduler from BaseAdvisor
|
||||
Scheduler scheduler = this.getScheduler();
|
||||
// Process the request with the before method
|
||||
return Mono.just(chatClientRequest)
|
||||
.publishOn(scheduler)
|
||||
.map(request -> this.before(request, streamAdvisorChain))
|
||||
.flatMapMany(streamAdvisorChain::nextStream)
|
||||
.transform(flux -> new ChatClientMessageAggregator().aggregateChatClientResponse(flux,
|
||||
response -> this.after(response, streamAdvisorChain)));
|
||||
}
|
||||
|
||||
private List<Document> toDocuments(List<Message> messages, String conversationId) {
|
||||
List<Document> docs = messages.stream()
|
||||
.filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT)
|
||||
|
||||
@@ -47,6 +47,7 @@ import org.springframework.ai.vectorstore.SearchRequest;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
import static org.assertj.core.api.Assertions.fail;
|
||||
import static org.mockito.ArgumentMatchers.any;
|
||||
import static org.mockito.BDDMockito.given;
|
||||
import static org.mockito.Mockito.mock;
|
||||
|
||||
Reference in New Issue
Block a user