Configure TemplateRenderer in ChatClient
- Extend the ChatClient with a new templateRenderer() method to pass a custom TemplateRenderer object used to render user and system templates.
- Evolve the QuestionAnswerAdvisor to accept a PromptTemplate for customising the RAG prompt and templating logic while maintaining backward compatibility.
- Introduce integration tests for the QuestionAnswerAdvisor.
- Document the TemplateRenderer API and how to use it to build PromptTemplate with custom templating logic.
- Document how to customise the templating logic used internally by the ChatClient via the TemplateRendererAPI.
Add validation tests and improve PromptTemplate resource handling
Enhance robustness and reliability of the PromptTemplate class with better
resource handling and comprehensive input validation:
- Add dedicated validation tests for builder methods with null/invalid inputs
- Improve renderResource method to gracefully handle edge cases:
- Null resources return empty string
- ByteArrayResource handling with proper charset (UTF-8)
- Empty resources check with proper existence test
- Better error handling with logging instead of exception propagation
- Add input validation assertions to all Builder methods
- Fix typo in deprecated annotation comment ("fahvor" → "favor")
Update documentation to clarify template rendering in different contexts:
- Add clear notes about TemplateRenderer usage in ChatClient vs Advisors
- Document how advisor template customization differs from ChatClient template rendering
- Add comprehensive API upgrade notes for template-related deprecations
- Include detailed migration examples for PromptTemplate and QuestionAnswerAdvisor
Fixes gh-355, gh-1687, gh-2448, gh-1849, gh-1428
Signed-off-by: Thomas Vitale <ThomasVitale@users.noreply.github.com>
This commit is contained in:
committed by
Mark Pollack
parent
b0d671944a
commit
5527d037f2
@@ -34,6 +34,7 @@ import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.content.Media;
|
||||
import org.springframework.ai.converter.StructuredOutputConverter;
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.ToolCallbackProvider;
|
||||
import org.springframework.core.ParameterizedTypeReference;
|
||||
@@ -247,6 +248,8 @@ public interface ChatClient {
|
||||
|
||||
ChatClientRequestSpec user(Consumer<PromptUserSpec> consumer);
|
||||
|
||||
ChatClientRequestSpec templateRenderer(TemplateRenderer templateRenderer);
|
||||
|
||||
CallResponseSpec call();
|
||||
|
||||
StreamResponseSpec stream();
|
||||
@@ -282,6 +285,8 @@ public interface ChatClient {
|
||||
|
||||
Builder defaultSystem(Consumer<PromptSystemSpec> systemSpecConsumer);
|
||||
|
||||
Builder defaultTemplateRenderer(TemplateRenderer templateRenderer);
|
||||
|
||||
Builder defaultTools(String... toolNames);
|
||||
|
||||
Builder defaultTools(ToolCallback... toolCallbacks);
|
||||
|
||||
@@ -57,6 +57,8 @@ import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.content.Media;
|
||||
import org.springframework.ai.converter.BeanOutputConverter;
|
||||
import org.springframework.ai.converter.StructuredOutputConverter;
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.template.st.StTemplateRenderer;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.ToolCallbackProvider;
|
||||
import org.springframework.ai.tool.ToolCallbacks;
|
||||
@@ -86,6 +88,8 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
private static final ChatClientObservationConvention DEFAULT_CHAT_CLIENT_OBSERVATION_CONVENTION = new DefaultChatClientObservationConvention();
|
||||
|
||||
private static final TemplateRenderer DEFAULT_TEMPLATE_RENDERER = StTemplateRenderer.builder().build();
|
||||
|
||||
private final DefaultChatClientRequestSpec defaultChatClientRequest;
|
||||
|
||||
public DefaultChatClient(DefaultChatClientRequestSpec defaultChatClientRequest) {
|
||||
@@ -136,7 +140,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
advisedRequest.toolCallbacks(), advisedRequest.messages(), advisedRequest.toolNames(),
|
||||
advisedRequest.media(), advisedRequest.chatOptions(), advisedRequest.advisors(),
|
||||
advisedRequest.advisorParams(), observationRegistry, customObservationConvention,
|
||||
advisedRequest.toolContext());
|
||||
advisedRequest.toolContext(), null);
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -638,6 +642,8 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
private final Map<String, Object> toolContext = new HashMap<>();
|
||||
|
||||
private TemplateRenderer templateRenderer;
|
||||
|
||||
@Nullable
|
||||
private String userText;
|
||||
|
||||
@@ -651,7 +657,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
DefaultChatClientRequestSpec(DefaultChatClientRequestSpec ccr) {
|
||||
this(ccr.chatModel, ccr.userText, ccr.userParams, ccr.systemText, ccr.systemParams, ccr.toolCallbacks,
|
||||
ccr.messages, ccr.toolNames, ccr.media, ccr.chatOptions, ccr.advisors, ccr.advisorParams,
|
||||
ccr.observationRegistry, ccr.observationConvention, ccr.toolContext);
|
||||
ccr.observationRegistry, ccr.observationConvention, ccr.toolContext, ccr.templateRenderer);
|
||||
}
|
||||
|
||||
public DefaultChatClientRequestSpec(ChatModel chatModel, @Nullable String userText,
|
||||
@@ -659,7 +665,8 @@ public class DefaultChatClient implements ChatClient {
|
||||
List<ToolCallback> toolCallbacks, List<Message> messages, List<String> toolNames, List<Media> media,
|
||||
@Nullable ChatOptions chatOptions, List<Advisor> advisors, Map<String, Object> advisorParams,
|
||||
ObservationRegistry observationRegistry,
|
||||
@Nullable ChatClientObservationConvention observationConvention, Map<String, Object> toolContext) {
|
||||
@Nullable ChatClientObservationConvention observationConvention, Map<String, Object> toolContext,
|
||||
@Nullable TemplateRenderer templateRenderer) {
|
||||
|
||||
Assert.notNull(chatModel, "chatModel cannot be null");
|
||||
Assert.notNull(userParams, "userParams cannot be null");
|
||||
@@ -692,6 +699,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
this.observationConvention = observationConvention != null ? observationConvention
|
||||
: DEFAULT_CHAT_CLIENT_OBSERVATION_CONVENTION;
|
||||
this.toolContext.putAll(toolContext);
|
||||
this.templateRenderer = templateRenderer != null ? templateRenderer : DEFAULT_TEMPLATE_RENDERER;
|
||||
}
|
||||
|
||||
private ObservationRegistry getObservationRegistry() {
|
||||
@@ -945,16 +953,22 @@ public class DefaultChatClient implements ChatClient {
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec templateRenderer(TemplateRenderer templateRenderer) {
|
||||
Assert.notNull(templateRenderer, "templateRenderer cannot be null");
|
||||
this.templateRenderer = templateRenderer;
|
||||
return this;
|
||||
}
|
||||
|
||||
public CallResponseSpec call() {
|
||||
BaseAdvisorChain advisorChain = buildAdvisorChain();
|
||||
return new DefaultCallResponseSpec(toAdvisedRequest(this).toChatClientRequest(), advisorChain,
|
||||
observationRegistry, observationConvention);
|
||||
return new DefaultCallResponseSpec(toAdvisedRequest(this).toChatClientRequest(this.templateRenderer),
|
||||
advisorChain, observationRegistry, observationConvention);
|
||||
}
|
||||
|
||||
public StreamResponseSpec stream() {
|
||||
BaseAdvisorChain advisorChain = buildAdvisorChain();
|
||||
return new DefaultStreamResponseSpec(toAdvisedRequest(this).toChatClientRequest(), advisorChain,
|
||||
observationRegistry, observationConvention);
|
||||
return new DefaultStreamResponseSpec(toAdvisedRequest(this).toChatClientRequest(this.templateRenderer),
|
||||
advisorChain, observationRegistry, observationConvention);
|
||||
}
|
||||
|
||||
private BaseAdvisorChain buildAdvisorChain() {
|
||||
@@ -963,7 +977,10 @@ public class DefaultChatClient implements ChatClient {
|
||||
this.advisors.add(ChatModelCallAdvisor.builder().chatModel(this.chatModel).build());
|
||||
this.advisors.add(ChatModelStreamAdvisor.builder().chatModel(this.chatModel).build());
|
||||
|
||||
return DefaultAroundAdvisorChain.builder(this.observationRegistry).pushAll(this.advisors).build();
|
||||
return DefaultAroundAdvisorChain.builder(this.observationRegistry)
|
||||
.pushAll(this.advisors)
|
||||
.templateRenderer(this.templateRenderer)
|
||||
.build();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -33,6 +33,7 @@ import org.springframework.ai.chat.client.observation.ChatClientObservationConve
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.ToolCallbackProvider;
|
||||
import org.springframework.ai.tool.function.FunctionToolCallback;
|
||||
@@ -66,7 +67,7 @@ public class DefaultChatClientBuilder implements Builder {
|
||||
Assert.notNull(observationRegistry, "the " + ObservationRegistry.class.getName() + " must be non-null");
|
||||
this.defaultRequest = new DefaultChatClientRequestSpec(chatModel, null, Map.of(), null, Map.of(), List.of(),
|
||||
List.of(), List.of(), List.of(), null, List.of(), Map.of(), observationRegistry,
|
||||
customObservationConvention, Map.of());
|
||||
customObservationConvention, Map.of(), null);
|
||||
}
|
||||
|
||||
public ChatClient build() {
|
||||
@@ -190,6 +191,12 @@ public class DefaultChatClientBuilder implements Builder {
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder defaultTemplateRenderer(TemplateRenderer templateRenderer) {
|
||||
Assert.notNull(templateRenderer, "templateRenderer cannot be null");
|
||||
this.defaultRequest.templateRenderer(templateRenderer);
|
||||
return this;
|
||||
}
|
||||
|
||||
void addMessages(List<Message> messages) {
|
||||
this.defaultRequest.messages(messages);
|
||||
}
|
||||
|
||||
@@ -33,6 +33,9 @@ import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAroundAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAroundAdvisor;
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.template.st.StTemplateRenderer;
|
||||
import org.springframework.lang.Nullable;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
import org.springframework.ai.chat.client.advisor.observation.AdvisorObservationContext;
|
||||
@@ -57,6 +60,8 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
|
||||
public static final AdvisorObservationConvention DEFAULT_OBSERVATION_CONVENTION = new DefaultAdvisorObservationConvention();
|
||||
|
||||
private static final TemplateRenderer DEFAULT_TEMPLATE_RENDERER = StTemplateRenderer.builder().build();
|
||||
|
||||
private final List<CallAroundAdvisor> originalCallAdvisors;
|
||||
|
||||
private final List<StreamAroundAdvisor> originalStreamAdvisors;
|
||||
@@ -67,14 +72,17 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
|
||||
private final ObservationRegistry observationRegistry;
|
||||
|
||||
DefaultAroundAdvisorChain(ObservationRegistry observationRegistry, Deque<CallAroundAdvisor> callAroundAdvisors,
|
||||
Deque<StreamAroundAdvisor> streamAroundAdvisors) {
|
||||
private final TemplateRenderer templateRenderer;
|
||||
|
||||
DefaultAroundAdvisorChain(ObservationRegistry observationRegistry, @Nullable TemplateRenderer templateRenderer,
|
||||
Deque<CallAroundAdvisor> callAroundAdvisors, Deque<StreamAroundAdvisor> streamAroundAdvisors) {
|
||||
|
||||
Assert.notNull(observationRegistry, "the observationRegistry must be non-null");
|
||||
Assert.notNull(callAroundAdvisors, "the callAroundAdvisors must be non-null");
|
||||
Assert.notNull(streamAroundAdvisors, "the streamAroundAdvisors must be non-null");
|
||||
|
||||
this.observationRegistry = observationRegistry;
|
||||
this.templateRenderer = templateRenderer != null ? templateRenderer : DEFAULT_TEMPLATE_RENDERER;
|
||||
this.callAroundAdvisors = callAroundAdvisors;
|
||||
this.streamAroundAdvisors = streamAroundAdvisors;
|
||||
this.originalCallAdvisors = List.copyOf(callAroundAdvisors);
|
||||
@@ -85,6 +93,11 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
return new Builder(observationRegistry);
|
||||
}
|
||||
|
||||
@Override
|
||||
public TemplateRenderer getTemplateRenderer() {
|
||||
return this.templateRenderer;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse nextCall(ChatClientRequest chatClientRequest) {
|
||||
Assert.notNull(chatClientRequest, "the chatClientRequest cannot be null");
|
||||
@@ -131,7 +144,7 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
|
||||
var observationContext = AdvisorObservationContext.builder()
|
||||
.advisorName(advisor.getName())
|
||||
.chatClientRequest(advisedRequest.toChatClientRequest())
|
||||
.chatClientRequest(advisedRequest.toChatClientRequest(templateRenderer))
|
||||
.order(advisor.getOrder())
|
||||
.build();
|
||||
|
||||
@@ -140,8 +153,8 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
.observe(() -> {
|
||||
// Supports both deprecated and new API.
|
||||
if (advisor instanceof CallAdvisor callAdvisor) {
|
||||
ChatClientResponse chatClientResponse = callAdvisor.adviseCall(advisedRequest.toChatClientRequest(),
|
||||
this);
|
||||
ChatClientResponse chatClientResponse = callAdvisor
|
||||
.adviseCall(advisedRequest.toChatClientRequest(templateRenderer), this);
|
||||
return AdvisedResponse.from(chatClientResponse);
|
||||
}
|
||||
AdvisedResponse advisedResponse = advisor.aroundCall(advisedRequest, this);
|
||||
@@ -209,7 +222,7 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
|
||||
AdvisorObservationContext observationContext = AdvisorObservationContext.builder()
|
||||
.advisorName(advisor.getName())
|
||||
.chatClientRequest(advisedRequest.toChatClientRequest())
|
||||
.chatClientRequest(advisedRequest.toChatClientRequest(templateRenderer))
|
||||
.order(advisor.getOrder())
|
||||
.build();
|
||||
|
||||
@@ -222,7 +235,7 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
return Flux.defer(() -> {
|
||||
// Supports both deprecated and new API.
|
||||
if (advisor instanceof StreamAdvisor streamAdvisor) {
|
||||
return streamAdvisor.adviseStream(advisedRequest.toChatClientRequest(), this)
|
||||
return streamAdvisor.adviseStream(advisedRequest.toChatClientRequest(templateRenderer), this)
|
||||
.doOnError(observation::error)
|
||||
.doFinally(s -> observation.stop())
|
||||
.contextWrite(ctx -> ctx.put(ObservationThreadLocalAccessor.KEY, observation))
|
||||
@@ -261,12 +274,19 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
|
||||
private final Deque<StreamAroundAdvisor> streamAroundAdvisors;
|
||||
|
||||
private TemplateRenderer templateRenderer;
|
||||
|
||||
public Builder(ObservationRegistry observationRegistry) {
|
||||
this.observationRegistry = observationRegistry;
|
||||
this.callAroundAdvisors = new ConcurrentLinkedDeque<>();
|
||||
this.streamAroundAdvisors = new ConcurrentLinkedDeque<>();
|
||||
}
|
||||
|
||||
public Builder templateRenderer(TemplateRenderer templateRenderer) {
|
||||
this.templateRenderer = templateRenderer;
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder push(Advisor advisor) {
|
||||
Assert.notNull(advisor, "the advisor must be non-null");
|
||||
return this.pushAll(List.of(advisor));
|
||||
@@ -315,8 +335,8 @@ public class DefaultAroundAdvisorChain implements BaseAdvisorChain {
|
||||
}
|
||||
|
||||
public DefaultAroundAdvisorChain build() {
|
||||
return new DefaultAroundAdvisorChain(this.observationRegistry, this.callAroundAdvisors,
|
||||
this.streamAroundAdvisors);
|
||||
return new DefaultAroundAdvisorChain(this.observationRegistry, this.templateRenderer,
|
||||
this.callAroundAdvisors, this.streamAroundAdvisors);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -37,6 +37,8 @@ import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.chat.prompt.PromptTemplate;
|
||||
import org.springframework.ai.content.Media;
|
||||
import org.springframework.ai.model.tool.ToolCallingChatOptions;
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.template.st.StTemplateRenderer;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.Assert;
|
||||
@@ -199,8 +201,12 @@ public record AdvisedRequest(
|
||||
}
|
||||
|
||||
public ChatClientRequest toChatClientRequest() {
|
||||
return toChatClientRequest(StTemplateRenderer.builder().build());
|
||||
}
|
||||
|
||||
public ChatClientRequest toChatClientRequest(TemplateRenderer templateRenderer) {
|
||||
return ChatClientRequest.builder()
|
||||
.prompt(toPrompt())
|
||||
.prompt(toPrompt(templateRenderer))
|
||||
.context(this.adviseContext)
|
||||
.context(ChatClientAttributes.ADVISORS.getKey(), this.advisors)
|
||||
.context(ChatClientAttributes.CHAT_MODEL.getKey(), this.chatModel)
|
||||
@@ -210,12 +216,21 @@ public record AdvisedRequest(
|
||||
}
|
||||
|
||||
public Prompt toPrompt() {
|
||||
return toPrompt(StTemplateRenderer.builder().build());
|
||||
}
|
||||
|
||||
public Prompt toPrompt(TemplateRenderer templateRenderer) {
|
||||
var messages = new ArrayList<>(this.messages());
|
||||
|
||||
String processedSystemText = this.systemText();
|
||||
if (StringUtils.hasText(processedSystemText)) {
|
||||
if (!CollectionUtils.isEmpty(this.systemParams())) {
|
||||
processedSystemText = new PromptTemplate(processedSystemText, this.systemParams()).render();
|
||||
processedSystemText = PromptTemplate.builder()
|
||||
.template(processedSystemText)
|
||||
.variables(this.systemParams())
|
||||
.renderer(templateRenderer)
|
||||
.build()
|
||||
.render();
|
||||
}
|
||||
messages.add(new SystemMessage(processedSystemText));
|
||||
}
|
||||
@@ -224,7 +239,12 @@ public record AdvisedRequest(
|
||||
Map<String, Object> userParams = new HashMap<>(this.userParams());
|
||||
String processedUserText = this.userText();
|
||||
if (!CollectionUtils.isEmpty(userParams)) {
|
||||
processedUserText = new PromptTemplate(processedUserText, userParams).render();
|
||||
processedUserText = PromptTemplate.builder()
|
||||
.template(processedUserText)
|
||||
.variables(userParams)
|
||||
.renderer(templateRenderer)
|
||||
.build()
|
||||
.render();
|
||||
}
|
||||
messages.add(new UserMessage(processedUserText, this.media()));
|
||||
}
|
||||
|
||||
@@ -16,6 +16,9 @@
|
||||
|
||||
package org.springframework.ai.chat.client.advisor.api;
|
||||
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.template.st.StTemplateRenderer;
|
||||
|
||||
/**
|
||||
* A base interface for advisor chains that can be used to chain multiple advisors
|
||||
* together, both for call and stream advisors.
|
||||
@@ -25,4 +28,8 @@ package org.springframework.ai.chat.client.advisor.api;
|
||||
*/
|
||||
public interface BaseAdvisorChain extends CallAdvisorChain, StreamAdvisorChain {
|
||||
|
||||
default TemplateRenderer getTemplateRenderer() {
|
||||
return StTemplateRenderer.builder().build();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -95,4 +95,11 @@ class DefaultChatClientBuilderTests {
|
||||
.hasMessage("charset cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenTemplateRendererIsNullThenThrows() {
|
||||
DefaultChatClientBuilder builder = new DefaultChatClientBuilder(mock(ChatModel.class));
|
||||
assertThatThrownBy(() -> builder.defaultTemplateRenderer(null)).isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("templateRenderer cannot be null");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -1302,7 +1302,7 @@ class DefaultChatClientTests {
|
||||
ChatModel chatModel = mock(ChatModel.class);
|
||||
DefaultChatClient.DefaultChatClientRequestSpec spec = new DefaultChatClient.DefaultChatClientRequestSpec(
|
||||
chatModel, null, Map.of(), null, Map.of(), List.of(), List.of(), List.of(), List.of(), null, List.of(),
|
||||
Map.of(), ObservationRegistry.NOOP, null, Map.of());
|
||||
Map.of(), ObservationRegistry.NOOP, null, Map.of(), null);
|
||||
assertThat(spec).isNotNull();
|
||||
}
|
||||
|
||||
@@ -1310,7 +1310,7 @@ class DefaultChatClientTests {
|
||||
void whenChatModelIsNullThenThrow() {
|
||||
assertThatThrownBy(() -> new DefaultChatClient.DefaultChatClientRequestSpec(null, null, Map.of(), null,
|
||||
Map.of(), List.of(), List.of(), List.of(), List.of(), null, List.of(), Map.of(),
|
||||
ObservationRegistry.NOOP, null, Map.of()))
|
||||
ObservationRegistry.NOOP, null, Map.of(), null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("chatModel cannot be null");
|
||||
}
|
||||
@@ -1319,7 +1319,7 @@ class DefaultChatClientTests {
|
||||
void whenObservationRegistryIsNullThenThrow() {
|
||||
assertThatThrownBy(() -> new DefaultChatClient.DefaultChatClientRequestSpec(mock(ChatModel.class), null,
|
||||
Map.of(), null, Map.of(), List.of(), List.of(), List.of(), List.of(), null, List.of(), Map.of(), null,
|
||||
null, Map.of()))
|
||||
null, Map.of(), null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("observationRegistry cannot be null");
|
||||
}
|
||||
|
||||
@@ -30,6 +30,8 @@ import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.content.Media;
|
||||
import org.springframework.ai.model.tool.ToolCallingChatOptions;
|
||||
import org.springframework.ai.template.TemplateRenderer;
|
||||
import org.springframework.ai.template.st.StTemplateRenderer;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
@@ -157,12 +159,12 @@ class AdvisedRequestTests {
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenConvertToAndFromChatClientRequest() {
|
||||
void whenConvertToAndFromChatClientRequestWithDefaultTemplateRenderer() {
|
||||
ChatModel chatModel = mock(ChatModel.class);
|
||||
ChatOptions chatOptions = ToolCallingChatOptions.builder().build();
|
||||
List<Message> messages = List.of(mock(UserMessage.class));
|
||||
SystemMessage systemMessage = new SystemMessage("Instructions {key}");
|
||||
UserMessage userMessage = new UserMessage("Question {key}", mock(Media.class));
|
||||
UserMessage userMessage = UserMessage.builder().text("Question {key}").media(mock(Media.class)).build();
|
||||
Map<String, Object> systemParams = Map.of("key", "value");
|
||||
Map<String, Object> userParams = Map.of("key", "value");
|
||||
List<String> toolNames = List.of("tool1", "tool2");
|
||||
@@ -208,6 +210,70 @@ class AdvisedRequestTests {
|
||||
AdvisedRequest convertedAdvisedRequest = AdvisedRequest.from(chatClientRequest);
|
||||
assertThat(convertedAdvisedRequest.toPrompt()).isEqualTo(chatClientRequest.prompt());
|
||||
assertThat(convertedAdvisedRequest.adviseContext()).containsAllEntriesOf(chatClientRequest.context());
|
||||
assertThat(chatClientRequest.context().get(ChatClientAttributes.USER_PARAMS.getKey())).isEqualTo(userParams);
|
||||
assertThat(chatClientRequest.context().get(ChatClientAttributes.SYSTEM_PARAMS.getKey()))
|
||||
.isEqualTo(systemParams);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenConvertToAndFromChatClientRequestWithCustomTemplateRenderer() {
|
||||
ChatModel chatModel = mock(ChatModel.class);
|
||||
ChatOptions chatOptions = ToolCallingChatOptions.builder().build();
|
||||
SystemMessage systemMessage = new SystemMessage("Instructions <name>");
|
||||
UserMessage userMessage = UserMessage.builder().text("Question <name>").media(mock(Media.class)).build();
|
||||
Map<String, Object> systemParams = Map.of("name", "Spring AI");
|
||||
Map<String, Object> userParams = Map.of("name", "Spring AI");
|
||||
|
||||
AdvisedRequest advisedRequest = AdvisedRequest.builder()
|
||||
.chatModel(chatModel)
|
||||
.chatOptions(chatOptions)
|
||||
.systemText(systemMessage.getText())
|
||||
.systemParams(systemParams)
|
||||
.userText(userMessage.getText())
|
||||
.userParams(userParams)
|
||||
.media(userMessage.getMedia())
|
||||
.build();
|
||||
|
||||
TemplateRenderer customRenderer = StTemplateRenderer.builder()
|
||||
.startDelimiterToken('<')
|
||||
.endDelimiterToken('>')
|
||||
.build();
|
||||
ChatClientRequest chatClientRequest = advisedRequest.toChatClientRequest(customRenderer);
|
||||
|
||||
assertThat(chatClientRequest.prompt().getInstructions()).hasSize(2);
|
||||
assertThat(chatClientRequest.prompt().getInstructions().get(0)).isInstanceOf(SystemMessage.class);
|
||||
assertThat(chatClientRequest.prompt().getInstructions().get(1)).isInstanceOf(UserMessage.class);
|
||||
assertThat(chatClientRequest.context().get(ChatClientAttributes.USER_PARAMS.getKey())).isEqualTo(userParams);
|
||||
assertThat(chatClientRequest.context().get(ChatClientAttributes.SYSTEM_PARAMS.getKey()))
|
||||
.isEqualTo(systemParams);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUsingToPromptWithCustomTemplateRenderer() {
|
||||
ChatModel chatModel = mock(ChatModel.class);
|
||||
SystemMessage systemMessage = new SystemMessage("Instructions <name>");
|
||||
UserMessage userMessage = UserMessage.builder().text("Question <name>").media(mock(Media.class)).build();
|
||||
Map<String, Object> systemParams = Map.of("name", "Spring AI");
|
||||
Map<String, Object> userParams = Map.of("name", "Spring AI");
|
||||
|
||||
AdvisedRequest advisedRequest = AdvisedRequest.builder()
|
||||
.chatModel(chatModel)
|
||||
.systemText(systemMessage.getText())
|
||||
.systemParams(systemParams)
|
||||
.userText(userMessage.getText())
|
||||
.userParams(userParams)
|
||||
.media(userMessage.getMedia())
|
||||
.build();
|
||||
|
||||
TemplateRenderer customRenderer = StTemplateRenderer.builder()
|
||||
.startDelimiterToken('<')
|
||||
.endDelimiterToken('>')
|
||||
.build();
|
||||
var prompt = advisedRequest.toPrompt(customRenderer);
|
||||
|
||||
assertThat(prompt.getInstructions()).hasSize(2);
|
||||
assertThat(prompt.getInstructions().get(0).getText()).isEqualTo("Instructions Spring AI");
|
||||
assertThat(prompt.getInstructions().get(1).getText()).isEqualTo("Question Spring AI");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user