Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -498,6 +498,25 @@ interface ChatClientRequestSpec {
*/
interface Builder {

/**
* Configure the default factory used to create converters for type-based
* {@link CallResponseSpec#entity(Class)} and
* {@link CallResponseSpec#responseEntity(Class)} calls, including their
* {@link ParameterizedTypeReference} variants.
* <p>
* Explicit converter overloads bypass this factory.
* @param factory the structured output converter factory
* @return this builder
* @throws UnsupportedOperationException if this builder implementation does not
* support structured output converter factory configuration
* @since 2.0.1
*/
default Builder defaultStructuredOutputConverterFactory(StructuredOutputConverterFactory factory) {
Assert.notNull(factory, "factory cannot be null");
throw new UnsupportedOperationException(
"Structured output converter factory configuration is not supported by " + getClass().getName());
}

Builder defaultAdvisors(Advisor... advisors);

Builder defaultAdvisors(Consumer<AdvisorSpec> advisorSpecConsumer);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,13 @@ public class DefaultChatClient implements ChatClient {

private static final ChatClientMessageAggregator CHAT_CLIENT_MESSAGE_AGGREGATOR = new ChatClientMessageAggregator();

private static final StructuredOutputConverterFactory DEFAULT_STRUCTURED_OUTPUT_CONVERTER_FACTORY = new StructuredOutputConverterFactory() {
@Override
public <T> StructuredOutputConverter<T> create(ParameterizedTypeReference<T> targetType) {
return new BeanOutputConverter<>(targetType);
}
};

private final DefaultChatClientRequestSpec defaultChatClientRequest;

public DefaultChatClient(DefaultChatClientRequestSpec defaultChatClientRequest) {
Expand Down Expand Up @@ -438,47 +445,58 @@ public static class DefaultCallResponseSpec implements CallResponseSpec {

private final ChatClientObservationConvention observationConvention;

private final StructuredOutputConverterFactory structuredOutputConverterFactory;

public DefaultCallResponseSpec(ChatClientRequest chatClientRequest, BaseAdvisorChain advisorChain,
ObservationRegistry observationRegistry, ChatClientObservationConvention observationConvention) {
this(chatClientRequest, advisorChain, observationRegistry, observationConvention,
DEFAULT_STRUCTURED_OUTPUT_CONVERTER_FACTORY);
}

DefaultCallResponseSpec(ChatClientRequest chatClientRequest, BaseAdvisorChain advisorChain,
ObservationRegistry observationRegistry, ChatClientObservationConvention observationConvention,
StructuredOutputConverterFactory structuredOutputConverterFactory) {
Assert.notNull(chatClientRequest, "chatClientRequest cannot be null");
Assert.notNull(advisorChain, "advisorChain cannot be null");
Assert.notNull(observationRegistry, "observationRegistry cannot be null");
Assert.notNull(observationConvention, "observationConvention cannot be null");
Assert.notNull(structuredOutputConverterFactory, "structuredOutputConverterFactory cannot be null");

this.request = chatClientRequest;
this.advisorChain = advisorChain;
this.observationRegistry = observationRegistry;
this.observationConvention = observationConvention;
this.structuredOutputConverterFactory = structuredOutputConverterFactory;
}

@Override
public <T> ResponseEntity<ChatResponse, T> responseEntity(Class<T> type,
Consumer<EntityParamSpec> entityParamSpecConsumer) {
Assert.notNull(type, "type cannot be null");
Assert.notNull(entityParamSpecConsumer, "entityParamSpecConsumer cannot be null");
var converter = new BeanOutputConverter<>(type);
StructuredOutputConverter<T> converter = createConverter(type);
return doResponseEntity(converter, resolveAdvisorChain(entityParamSpecConsumer, converter));
}

@Override
public <T> ResponseEntity<ChatResponse, T> responseEntity(Class<T> type) {
Assert.notNull(type, "type cannot be null");
return doResponseEntity(new BeanOutputConverter<>(type));
return doResponseEntity(createConverter(type));
}

@Override
public <T> ResponseEntity<ChatResponse, T> responseEntity(ParameterizedTypeReference<T> type,
Consumer<EntityParamSpec> entityParamSpecConsumer) {
Assert.notNull(type, "type cannot be null");
Assert.notNull(entityParamSpecConsumer, "entityParamSpecConsumer cannot be null");
var converter = new BeanOutputConverter<>(type);
StructuredOutputConverter<T> converter = createConverter(type);
return doResponseEntity(converter, resolveAdvisorChain(entityParamSpecConsumer, converter));
}

@Override
public <T> ResponseEntity<ChatResponse, T> responseEntity(ParameterizedTypeReference<T> type) {
Assert.notNull(type, "type cannot be null");
return doResponseEntity(new BeanOutputConverter<>(type));
return doResponseEntity(createConverter(type));
}

@Override
Expand Down Expand Up @@ -530,43 +548,53 @@ protected <T> ResponseEntity<ChatResponse, T> doResponseEntity(StructuredOutputC
Consumer<EntityParamSpec> entitySpecConsumer) {
Assert.notNull(type, "type cannot be null");
Assert.notNull(entitySpecConsumer, "entitySpecConsumer cannot be null");
var converter = new BeanOutputConverter<>(type);
return doSingleWithBeanOutputConverter(converter, resolveAdvisorChain(entitySpecConsumer, converter));
StructuredOutputConverter<T> converter = createConverter(type);
return doSingleWithStructuredOutputConverter(converter, resolveAdvisorChain(entitySpecConsumer, converter));
}

@Override
public <T> @Nullable T entity(Class<T> type, Consumer<EntityParamSpec> entitySpecConsumer) {
Assert.notNull(type, "type cannot be null");
Assert.notNull(entitySpecConsumer, "entitySpecConsumer cannot be null");
var converter = new BeanOutputConverter<>(type);
return doSingleWithBeanOutputConverter(converter, resolveAdvisorChain(entitySpecConsumer, converter));
StructuredOutputConverter<T> converter = createConverter(type);
return doSingleWithStructuredOutputConverter(converter, resolveAdvisorChain(entitySpecConsumer, converter));
}

@Override
public <T> @Nullable T entity(ParameterizedTypeReference<T> type) {
Assert.notNull(type, "type cannot be null");
return doSingleWithBeanOutputConverter(new BeanOutputConverter<>(type));
return doSingleWithStructuredOutputConverter(createConverter(type));
}

@Override
public <T> @Nullable T entity(StructuredOutputConverter<T> structuredOutputConverter,
Consumer<EntityParamSpec> entitySpecConsumer) {
Assert.notNull(structuredOutputConverter, "structuredOutputConverter cannot be null");
Assert.notNull(entitySpecConsumer, "entitySpecConsumer cannot be null");
return doSingleWithBeanOutputConverter(structuredOutputConverter,
return doSingleWithStructuredOutputConverter(structuredOutputConverter,
resolveAdvisorChain(entitySpecConsumer, structuredOutputConverter));
}

@Override
public <T> @Nullable T entity(StructuredOutputConverter<T> structuredOutputConverter) {
Assert.notNull(structuredOutputConverter, "structuredOutputConverter cannot be null");
return doSingleWithBeanOutputConverter(structuredOutputConverter);
return doSingleWithStructuredOutputConverter(structuredOutputConverter);
}

@Override
public <T> @Nullable T entity(Class<T> type) {
Assert.notNull(type, "type cannot be null");
return doSingleWithBeanOutputConverter(new BeanOutputConverter<>(type));
return doSingleWithStructuredOutputConverter(createConverter(type));
}

private <T> StructuredOutputConverter<T> createConverter(Class<T> targetType) {
return createConverter(ParameterizedTypeReference.forType(targetType));
}

private <T> StructuredOutputConverter<T> createConverter(ParameterizedTypeReference<T> targetType) {
StructuredOutputConverter<T> converter = this.structuredOutputConverterFactory.create(targetType);
Assert.notNull(converter, "structuredOutputConverterFactory must not return null");
return converter;
}

private BaseAdvisorChain resolveAdvisorChain(Consumer<EntityParamSpec> consumer,
Expand All @@ -585,11 +613,11 @@ private BaseAdvisorChain resolveAdvisorChain(Consumer<EntityParamSpec> consumer,
return this.advisorChain;
}

private <T> @Nullable T doSingleWithBeanOutputConverter(StructuredOutputConverter<T> outputConverter) {
return doSingleWithBeanOutputConverter(outputConverter, this.advisorChain);
private <T> @Nullable T doSingleWithStructuredOutputConverter(StructuredOutputConverter<T> outputConverter) {
return doSingleWithStructuredOutputConverter(outputConverter, this.advisorChain);
}

private <T> @Nullable T doSingleWithBeanOutputConverter(StructuredOutputConverter<T> outputConverter,
private <T> @Nullable T doSingleWithStructuredOutputConverter(StructuredOutputConverter<T> outputConverter,
BaseAdvisorChain advisorChain) {

if (StringUtils.hasText(outputConverter.getFormat())) {
Expand Down Expand Up @@ -800,13 +828,16 @@ public static class DefaultChatClientRequestSpec implements ChatClientRequestSpe

private final ToolCallingAdvisor.Builder<?> toolCallingAdvisorBuilder;

private StructuredOutputConverterFactory structuredOutputConverterFactory = DEFAULT_STRUCTURED_OUTPUT_CONVERTER_FACTORY;

/* copy constructor */
DefaultChatClientRequestSpec(DefaultChatClientRequestSpec ccr) {
this(ccr.chatModel, ccr.userText, ccr.userParams, ccr.userMetadata, ccr.systemText, ccr.systemParams,
ccr.systemMetadata, ccr.toolCallbacks, ccr.toolCallbackProviders, ccr.messages, ccr.media,
ccr.optionsCustomizer, ccr.advisors, ccr.advisorParams, ccr.observationRegistry,
ccr.chatClientObservationConvention, ccr.toolContext, ccr.templateRenderer,
ccr.advisorObservationConvention, ccr.toolCallingAdvisorBuilder);
this.structuredOutputConverterFactory = ccr.structuredOutputConverterFactory;
}

public DefaultChatClientRequestSpec(ChatModel chatModel, @Nullable String userText,
Expand Down Expand Up @@ -917,6 +948,11 @@ public TemplateRenderer getTemplateRenderer() {
return this.templateRenderer;
}

void structuredOutputConverterFactory(StructuredOutputConverterFactory factory) {
Assert.notNull(factory, "factory cannot be null");
this.structuredOutputConverterFactory = factory;
}

/* package */ ChatModel getChatModel() {
return this.chatModel;
}
Expand All @@ -937,7 +973,8 @@ public Builder mutate() {
.defaultTemplateRenderer(this.templateRenderer)
.defaultTools(this.toolCallbacks.toArray(new ToolCallback[0]))
.defaultTools((Object[]) this.toolCallbackProviders.toArray(new ToolCallbackProvider[0]))
.defaultToolContext(this.toolContext);
.defaultToolContext(this.toolContext)
.defaultStructuredOutputConverterFactory(this.structuredOutputConverterFactory);

if (!CollectionUtils.isEmpty(this.advisors)) {
builder.defaultAdvisors(a -> a.advisors(this.advisors).params(this.advisorParams));
Expand Down Expand Up @@ -1176,7 +1213,8 @@ public ChatClientRequestSpec templateRenderer(TemplateRenderer templateRenderer)
public CallResponseSpec call() {
BaseAdvisorChain advisorChain = buildAdvisorChain();
return new DefaultCallResponseSpec(DefaultChatClientUtils.toChatClientRequest(this), advisorChain,
this.observationRegistry, this.chatClientObservationConvention);
this.observationRegistry, this.chatClientObservationConvention,
this.structuredOutputConverterFactory);
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,13 @@ public Builder clone() {
return this.defaultRequest.mutate();
}

@Override
public Builder defaultStructuredOutputConverterFactory(StructuredOutputConverterFactory factory) {
Assert.notNull(factory, "factory cannot be null");
this.defaultRequest.structuredOutputConverterFactory(factory);
return this;
}

public Builder defaultAdvisors(Advisor... advisors) {
this.defaultRequest.advisors(advisors);
return this;
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
/*
* Copyright 2023-present the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

package org.springframework.ai.chat.client;

import tools.jackson.databind.json.JsonMapper;

import org.springframework.ai.converter.BeanOutputConverter;
import org.springframework.ai.converter.ResponseTextCleaner;
import org.springframework.ai.converter.StructuredOutputConverter;
import org.springframework.core.ParameterizedTypeReference;
import org.springframework.util.Assert;

/**
* Factory for creating {@link StructuredOutputConverter} instances for type-based
* {@link ChatClient.CallResponseSpec#entity(Class)} and
* {@link ChatClient.CallResponseSpec#responseEntity(Class)} calls.
* <p>
* Implementations must return a non-null converter that supports the supplied target
* type. A factory is shared by the {@link ChatClient} instances created or derived from a
* configured builder and must therefore be safe to invoke concurrently. Any converter
* instance shared between invocations must also be safe for concurrent use.
* <p>
* Explicit converter overloads bypass this factory.
*
* @author Jaehyeon Park
* @since 2.0.1
*/
public interface StructuredOutputConverterFactory {

/**
* Create a converter for the given target type.
* @param targetType the target type
* @return a non-null structured output converter
* @since 2.0.1
*/
<T> StructuredOutputConverter<T> create(ParameterizedTypeReference<T> targetType);

/**
* Create a factory that uses {@link BeanOutputConverter} with the given JSON mapper.
* @param jsonMapper the JSON mapper
* @return a structured output converter factory
* @since 2.0.1
*/
static StructuredOutputConverterFactory beanOutputConverter(JsonMapper jsonMapper) {
Assert.notNull(jsonMapper, "jsonMapper cannot be null");
return new StructuredOutputConverterFactory() {
@Override
public <T> StructuredOutputConverter<T> create(ParameterizedTypeReference<T> targetType) {
return new BeanOutputConverter<>(targetType, jsonMapper);
}
};
}

/**
* Create a factory that uses {@link BeanOutputConverter} with the given JSON mapper
* and response text cleaner.
* @param jsonMapper the JSON mapper
* @param textCleaner the response text cleaner
* @return a structured output converter factory
* @since 2.0.1
*/
static StructuredOutputConverterFactory beanOutputConverter(JsonMapper jsonMapper,
ResponseTextCleaner textCleaner) {
Assert.notNull(jsonMapper, "jsonMapper cannot be null");
Assert.notNull(textCleaner, "textCleaner cannot be null");
return new StructuredOutputConverterFactory() {
@Override
public <T> StructuredOutputConverter<T> create(ParameterizedTypeReference<T> targetType) {
return new BeanOutputConverter<>(targetType, jsonMapper, textCleaner);
}
};
}

}
Loading