11package dev .langchain4j .service .guardrail ;
22
3- import static org .assertj .core .api .Assertions .assertThat ;
4- import static org .assertj .core .api .Assertions .assertThatExceptionOfType ;
5-
63import dev .langchain4j .data .image .Image ;
74import dev .langchain4j .data .message .AiMessage ;
85import dev .langchain4j .data .message .ChatMessageType ;
1613import dev .langchain4j .guardrail .OutputGuardrail ;
1714import dev .langchain4j .guardrail .OutputGuardrailResult ;
1815import dev .langchain4j .model .chat .ChatModel ;
16+ import dev .langchain4j .model .chat .mock .ChatModelMock ;
1917import dev .langchain4j .model .chat .request .ChatRequest ;
2018import dev .langchain4j .model .chat .response .ChatResponse ;
2119import dev .langchain4j .rag .AugmentationRequest ;
2220import dev .langchain4j .rag .AugmentationResult ;
2321import dev .langchain4j .rag .RetrievalAugmentor ;
2422import dev .langchain4j .service .AiServices ;
25- import java .util .concurrent .atomic .AtomicReference ;
26- import java .util .stream .Stream ;
2723import org .junit .jupiter .api .Test ;
2824import org .junit .jupiter .params .ParameterizedTest ;
2925import org .junit .jupiter .params .provider .Arguments ;
3026import 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+
3234class 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