Skip to content

Commit a639f2e

Browse files
committed
fix: materialize content arguments before input guardrails (langchain4j#4679)
1 parent 676832e commit a639f2e

2 files changed

Lines changed: 20 additions & 38 deletions

File tree

langchain4j/src/main/java/dev/langchain4j/service/DefaultAiServices.java

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -205,6 +205,8 @@ public Object invoke(Method method, Object[] args, InvocationContext invocationC
205205
userMessageForAugmentation = (UserMessage) augmentationResult.chatMessage();
206206
}
207207

208+
UserMessage userMessage = addContentsToUserMessage(method, args, userMessageForAugmentation);
209+
208210
var commonGuardrailParam = GuardrailRequestParams.builder()
209211
.chatMemory(chatMemory)
210212
.augmentationResult(augmentationResult)
@@ -214,7 +216,6 @@ public Object invoke(Method method, Object[] args, InvocationContext invocationC
214216
.variables(variables)
215217
.build();
216218

217-
UserMessage userMessage = addContentsToUserMessage(method, args, userMessageForAugmentation);
218219
userMessage = invokeInputGuardrails(
219220
context.guardrailService(), method, userMessage, commonGuardrailParam);
220221

langchain4j/src/test/java/dev/langchain4j/service/guardrail/AiServiceGuardrailTests.java

Lines changed: 18 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,5 @@
11
package dev.langchain4j.service.guardrail;
22

3-
import static org.assertj.core.api.Assertions.assertThat;
4-
import static org.assertj.core.api.Assertions.assertThatExceptionOfType;
5-
63
import dev.langchain4j.data.image.Image;
74
import dev.langchain4j.data.message.AiMessage;
85
import dev.langchain4j.data.message.ChatMessageType;
@@ -16,19 +13,24 @@
1613
import dev.langchain4j.guardrail.OutputGuardrail;
1714
import dev.langchain4j.guardrail.OutputGuardrailResult;
1815
import dev.langchain4j.model.chat.ChatModel;
16+
import dev.langchain4j.model.chat.mock.ChatModelMock;
1917
import dev.langchain4j.model.chat.request.ChatRequest;
2018
import dev.langchain4j.model.chat.response.ChatResponse;
2119
import dev.langchain4j.rag.AugmentationRequest;
2220
import dev.langchain4j.rag.AugmentationResult;
2321
import dev.langchain4j.rag.RetrievalAugmentor;
2422
import dev.langchain4j.service.AiServices;
25-
import java.util.concurrent.atomic.AtomicReference;
26-
import java.util.stream.Stream;
2723
import org.junit.jupiter.api.Test;
2824
import org.junit.jupiter.params.ParameterizedTest;
2925
import org.junit.jupiter.params.provider.Arguments;
3026
import org.junit.jupiter.params.provider.MethodSource;
3127

28+
import java.util.concurrent.atomic.AtomicReference;
29+
import java.util.stream.Stream;
30+
31+
import static org.assertj.core.api.Assertions.assertThat;
32+
import static org.assertj.core.api.Assertions.assertThatExceptionOfType;
33+
3234
class AiServiceGuardrailTests {
3335
private static final ImageContent IMAGE_CONTENT = ImageContent.from(
3436
Image.builder().url("https://example.com/image.png").build());
@@ -110,11 +112,10 @@ void classAndMethodLevelAssistant() {
110112

111113
@Test
112114
void input_guardrail_should_receive_materialized_multimodal_user_message() {
115+
ChatModelMock chatModelMock = ChatModelMock.thatAlwaysResponds("does not matter");
113116
RecordingInputGuardrail inputGuardrail = new RecordingInputGuardrail();
114-
AtomicReference<UserMessage> userMessageSeenByChatModel = new AtomicReference<>();
115-
116117
VisionAssistant assistant = AiServices.builder(VisionAssistant.class)
117-
.chatModel(new RecordingChatModel(userMessageSeenByChatModel))
118+
.chatModel(chatModelMock)
118119
.inputGuardrails(inputGuardrail)
119120
.build();
120121

@@ -124,22 +125,21 @@ void input_guardrail_should_receive_materialized_multimodal_user_message() {
124125
assertThat(inputGuardrail.observedUserMessage().contents())
125126
.containsExactly(TextContent.from("Describe this image"), IMAGE_CONTENT);
126127
assertThat(inputGuardrail.observedUserMessage().hasSingleText()).isFalse();
127-
assertThat(userMessageSeenByChatModel.get()).isEqualTo(inputGuardrail.observedUserMessage());
128+
assertThat(chatModelMock.request().messages().get(0)).isEqualTo(inputGuardrail.observedUserMessage());
128129
}
129130

130131
@Test
131132
void input_guardrail_should_observe_augmented_user_message_after_rag_and_before_chat_request() {
133+
ChatModelMock chatModelMock = ChatModelMock.thatAlwaysResponds("does not matter");
132134
RecordingInputGuardrail inputGuardrail = new RecordingInputGuardrail();
133135
AtomicReference<UserMessage> userMessageSeenByAugmentor = new AtomicReference<>();
134-
AtomicReference<UserMessage> userMessageSeenByChatModel = new AtomicReference<>();
135-
136136
RetrievalAugmentor retrievalAugmentor = (AugmentationRequest request) -> {
137137
userMessageSeenByAugmentor.set((UserMessage) request.chatMessage());
138138
return new AugmentationResult(UserMessage.from("Augmented prompt"), null);
139139
};
140140

141141
VisionAssistant assistant = AiServices.builder(VisionAssistant.class)
142-
.chatModel(new RecordingChatModel(userMessageSeenByChatModel))
142+
.chatModel(chatModelMock)
143143
.inputGuardrails(inputGuardrail)
144144
.retrievalAugmentor(retrievalAugmentor)
145145
.build();
@@ -150,22 +150,22 @@ void input_guardrail_should_observe_augmented_user_message_after_rag_and_before_
150150
.containsExactly(TextContent.from("Describe this image"));
151151
assertThat(inputGuardrail.observedUserMessage().contents())
152152
.containsExactly(TextContent.from("Augmented prompt"), IMAGE_CONTENT);
153-
assertThat(userMessageSeenByChatModel.get()).isEqualTo(inputGuardrail.observedUserMessage());
153+
assertThat(chatModelMock.request().messages().get(0)).isEqualTo(inputGuardrail.observedUserMessage());
154154
}
155155

156156
@Test
157157
void input_guardrail_rewrite_should_still_work_for_plain_text_requests() {
158-
AtomicReference<UserMessage> userMessageSeenByChatModel = new AtomicReference<>();
159-
158+
ChatModelMock chatModelMock = ChatModelMock.thatAlwaysResponds("does not matter");
160159
PlainTextAssistant assistant = AiServices.builder(PlainTextAssistant.class)
161-
.chatModel(new RecordingChatModel(userMessageSeenByChatModel))
160+
.chatModel(chatModelMock)
162161
.inputGuardrails(new RewritingInputGuardrail())
163162
.build();
164163

165164
assistant.chat("Original prompt");
166165

167-
assertThat(userMessageSeenByChatModel.get().contents()).containsExactly(TextContent.from("Rewritten prompt"));
168-
assertThat(userMessageSeenByChatModel.get().hasSingleText()).isTrue();
166+
UserMessage userMessage = (UserMessage) chatModelMock.request().messages().get(0);
167+
assertThat(userMessage.contents()).containsExactly(TextContent.from("Rewritten prompt"));
168+
assertThat(userMessage.hasSingleText()).isTrue();
169169
}
170170

171171
static Stream<Arguments> classLevelAssistants() {
@@ -379,23 +379,4 @@ public ChatResponse doChat(ChatRequest chatRequest) {
379379
.build();
380380
}
381381
}
382-
383-
static class RecordingChatModel implements ChatModel {
384-
385-
private final AtomicReference<UserMessage> observedUserMessage;
386-
387-
RecordingChatModel(AtomicReference<UserMessage> observedUserMessage) {
388-
this.observedUserMessage = observedUserMessage;
389-
}
390-
391-
@Override
392-
public ChatResponse doChat(ChatRequest chatRequest) {
393-
observedUserMessage.set(chatRequest.messages().stream()
394-
.filter(message -> message.type() == ChatMessageType.USER)
395-
.map(UserMessage.class::cast)
396-
.findFirst()
397-
.orElseThrow(() -> new IllegalStateException("No user message found")));
398-
return ChatResponse.builder().aiMessage(AiMessage.from("ok")).build();
399-
}
400-
}
401382
}

0 commit comments

Comments
 (0)