Copy metadata in ChatResponse Builder's from() method

- Update ChatResponse.Builder to copy all metadata fields when using from()
 - Expand test case to verify correct metadata copying in QA advisor

 Resolves #1537
This commit is contained in:
Christian Tzolov
2024-10-14 22:57:27 +02:00
parent a39aadc327
commit 3d252ca173
2 changed files with 78 additions and 3 deletions

View File

@@ -125,6 +125,11 @@ public class ChatResponse implements ModelResponse<Generation> {
public Builder from(ChatResponse other) {
this.generations = other.generations;
this.chatResponseMetadataBuilder.withModel(other.chatResponseMetadata.getModel());
this.chatResponseMetadataBuilder.withId(other.chatResponseMetadata.getId());
this.chatResponseMetadataBuilder.withRateLimit(other.chatResponseMetadata.getRateLimit());
this.chatResponseMetadataBuilder.withUsage(other.chatResponseMetadata.getUsage());
this.chatResponseMetadataBuilder.withPromptMetadata(other.chatResponseMetadata.getPromptMetadata());
Set<Map.Entry<String, Object>> entries = other.chatResponseMetadata.entrySet();
for (Map.Entry<String, Object> entry : entries) {
this.chatResponseMetadataBuilder.withKeyValue(entry.getKey(), entry.getValue());

View File

@@ -19,7 +19,9 @@ package org.springframework.ai.chat.client.advisor;
import static org.assertj.core.api.Assertions.assertThat;
import static org.mockito.Mockito.when;
import java.time.Duration;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
@@ -28,8 +30,12 @@ import org.mockito.Captor;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.MessageType;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
@@ -60,8 +66,50 @@ public class QuestionAnswerAdvisorTests {
@Test
public void qaAdvisorWithDynamicFilterExpressions() {
// @formatter:off
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("Your answer is ZXY"))));
.thenReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("Your answer is ZXY"))),
ChatResponseMetadata.builder()
.withId("678")
.withModel("model1")
.withKeyValue("key6", "value6")
.withMetadata(Map.of("key1","value1" ))
.withPromptMetadata(null)
.withRateLimit(new RateLimit() {
@Override
public Long getRequestsLimit() {
return 5l;
}
@Override
public Long getRequestsRemaining() {
return 6l;
}
@Override
public Duration getRequestsReset() {
return Duration.ofSeconds(7);
}
@Override
public Long getTokensLimit() {
return 8l;
}
@Override
public Long getTokensRemaining() {
return 8l;
}
@Override
public Duration getTokensReset() {
return Duration.ofSeconds(9);
}
})
.withUsage(new DefaultUsage(6l, 7l))
.build()));
// @formatter:on
when(vectorStore.similaritySearch(vectorSearchCaptor.capture()))
.thenReturn(List.of(new Document("doc1"), new Document("doc2")));
@@ -75,13 +123,33 @@ public class QuestionAnswerAdvisorTests {
.build();
// @formatter:off
var content = chatClient.prompt()
var response = chatClient.prompt()
.user("Please answer my question XYZ")
.advisors(a -> a.param(QuestionAnswerAdvisor.FILTER_EXPRESSION, "type == 'Spring'"))
.call()
.content();
.chatResponse();
//formatter:on
// Ensure the metadata is correctly copied over
assertThat(response.getMetadata().getModel()).isEqualTo("model1");
assertThat(response.getMetadata().getId()).isEqualTo("678");
assertThat(response.getMetadata().getRateLimit().getRequestsLimit()).isEqualTo(5l);
assertThat(response.getMetadata().getRateLimit().getRequestsRemaining()).isEqualTo(6l);
assertThat(response.getMetadata().getRateLimit().getRequestsReset()).isEqualTo(Duration.ofSeconds(7));
assertThat(response.getMetadata().getRateLimit().getTokensLimit()).isEqualTo(8l);
assertThat(response.getMetadata().getRateLimit().getTokensRemaining()).isEqualTo(8l);
assertThat(response.getMetadata().getRateLimit().getTokensReset()).isEqualTo(Duration.ofSeconds(9));
assertThat(response.getMetadata().getUsage().getPromptTokens()).isEqualTo(6l);
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isEqualTo(7l);
assertThat(response.getMetadata().getUsage().getTotalTokens()).isEqualTo(6l + 7l);
assertThat(response.getMetadata().get("key6").toString()).isEqualTo("value6");
assertThat(response.getMetadata().get("key1").toString()).isEqualTo("value1");
String content = response.getResult().getOutput().getContent();
assertThat(content).isEqualTo("Your answer is ZXY");
Message systemMessage = promptCaptor.getValue().getInstructions().get(0);
@@ -112,5 +180,7 @@ public class QuestionAnswerAdvisorTests {
assertThat(vectorSearchCaptor.getValue().getFilterExpression()).isEqualTo(new FilterExpressionBuilder().eq("type", "Spring").build());
assertThat(vectorSearchCaptor.getValue().getSimilarityThreshold()).isEqualTo(0.99d);
assertThat(vectorSearchCaptor.getValue().getTopK()).isEqualTo(6);
}
}