Refact onFinishReason method to utility class

Test AdvisedResponseStreamUtils

Add java docs

Signed-off-by: ghdcksgml1 <ghdcksgml2@naver.com>
This commit is contained in:
ghdcksgml1
2025-03-15 11:38:59 +09:00
committed by Ilayaperumal Gopinathan
parent d25d37ab12
commit e4357ba2c1
4 changed files with 116 additions and 34 deletions

View File

@@ -19,19 +19,13 @@ package org.springframework.ai.chat.client.advisor;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.function.Predicate;
import java.util.stream.Collectors;
import org.springframework.ai.chat.client.advisor.api.*;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
import reactor.core.scheduler.Schedulers;
import org.springframework.ai.chat.client.advisor.api.AdvisedRequest;
import org.springframework.ai.chat.client.advisor.api.AdvisedResponse;
import org.springframework.ai.chat.client.advisor.api.CallAroundAdvisor;
import org.springframework.ai.chat.client.advisor.api.CallAroundAdvisorChain;
import org.springframework.ai.chat.client.advisor.api.StreamAroundAdvisor;
import org.springframework.ai.chat.client.advisor.api.StreamAroundAdvisorChain;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.document.Document;
@@ -201,7 +195,7 @@ public class QuestionAnswerAdvisor implements CallAroundAdvisor, StreamAroundAdv
// @formatter:on
return advisedResponses.map(ar -> {
if (onFinishReason().test(ar)) {
if (AdvisedResponseStreamUtils.onFinishReason().test(ar)) {
ar = after(ar);
}
return ar;
@@ -260,16 +254,6 @@ public class QuestionAnswerAdvisor implements CallAroundAdvisor, StreamAroundAdv
}
private Predicate<AdvisedResponse> onFinishReason() {
return advisedResponse -> advisedResponse.response()
.getResults()
.stream()
.filter(result -> result != null && result.getMetadata() != null
&& StringUtils.hasText(result.getMetadata().getFinishReason()))
.findFirst()
.isPresent();
}
public static final class Builder {
private final VectorStore vectorStore;

View File

@@ -0,0 +1,31 @@
package org.springframework.ai.chat.client.advisor.api;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.util.StringUtils;
import java.util.function.Predicate;
/**
* A stream utility class to provide support methods handling {@link AdvisedResponse}.
*/
public final class AdvisedResponseStreamUtils {
/**
* Returns a predicate that checks whether the provided {@link AdvisedResponse}
* contains a {@link ChatResponse} with at least one result having a non-empty finish
* reason in its metadata.
* @return a {@link Predicate} that evaluates whether the finish reason exists within
* the response metadata.
*/
public static Predicate<AdvisedResponse> onFinishReason() {
return advisedResponse -> {
ChatResponse chatResponse = advisedResponse.response();
return chatResponse != null && chatResponse.getResults() != null
&& chatResponse.getResults()
.stream()
.anyMatch(result -> result != null && result.getMetadata() != null
&& StringUtils.hasText(result.getMetadata().getFinishReason()));
};
}
}

View File

@@ -16,16 +16,12 @@
package org.springframework.ai.chat.client.advisor.api;
import java.util.function.Predicate;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
import reactor.core.scheduler.Scheduler;
import reactor.core.scheduler.Schedulers;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
/**
* Base advisor that implements common aspects of the {@link CallAroundAdvisor} and
@@ -65,24 +61,13 @@ public interface BaseAdvisor extends CallAroundAdvisor, StreamAroundAdvisor {
.flatMapMany(chain::nextAroundStream);
return advisedResponses.map(ar -> {
if (onFinishReason().test(ar)) {
if (AdvisedResponseStreamUtils.onFinishReason().test(ar)) {
ar = after(ar);
}
return ar;
}).onErrorResume(error -> Flux.error(new IllegalStateException("Stream processing failed", error)));
}
private Predicate<AdvisedResponse> onFinishReason() {
return advisedResponse -> {
ChatResponse chatResponse = advisedResponse.response();
return chatResponse != null && chatResponse.getResults() != null
&& chatResponse.getResults()
.stream()
.anyMatch(result -> result != null && result.getMetadata() != null
&& StringUtils.hasText(result.getMetadata().getFinishReason()));
};
}
@Override
default String getName() {
return this.getClass().getSimpleName();

View File

@@ -0,0 +1,82 @@
package org.springframework.ai.chat.client.advisor.api;
import org.junit.jupiter.api.Nested;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import java.util.List;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.BDDMockito.given;
import static org.mockito.Mockito.mock;
/**
* Unit tests for {@link AdvisedResponseStreamUtils}.
*
* @author ghdcksgml1
*/
class AdvisedResponseStreamUtilsTest {
@Nested
class OnFinishReason {
@Test
void whenChatResponseIsNullThenReturnFalse() {
AdvisedResponse response = mock(AdvisedResponse.class);
given(response.response()).willReturn(null);
boolean result = AdvisedResponseStreamUtils.onFinishReason().test(response);
assertFalse(result);
}
@Test
void whenChatResponseResultsIsNullThenReturnFalse() {
AdvisedResponse response = mock(AdvisedResponse.class);
ChatResponse chatResponse = mock(ChatResponse.class);
given(chatResponse.getResults()).willReturn(null);
given(response.response()).willReturn(chatResponse);
boolean result = AdvisedResponseStreamUtils.onFinishReason().test(response);
assertFalse(result);
}
@Test
void whenChatIsRunningThenReturnFalse() {
AdvisedResponse response = mock(AdvisedResponse.class);
ChatResponse chatResponse = mock(ChatResponse.class);
Generation generation = new Generation(new AssistantMessage("running.."), ChatGenerationMetadata.NULL);
given(chatResponse.getResults()).willReturn(List.of(generation));
given(response.response()).willReturn(chatResponse);
boolean result = AdvisedResponseStreamUtils.onFinishReason().test(response);
assertFalse(result);
}
@Test
void whenChatIsStopThenReturnTrue() {
AdvisedResponse response = mock(AdvisedResponse.class);
ChatResponse chatResponse = mock(ChatResponse.class);
Generation generation = new Generation(new AssistantMessage("finish."),
ChatGenerationMetadata.builder().finishReason("STOP").build());
given(chatResponse.getResults()).willReturn(List.of(generation));
given(response.response()).willReturn(chatResponse);
boolean result = AdvisedResponseStreamUtils.onFinishReason().test(response);
assertTrue(result);
}
}
}