Refactor toolcalling support for ZhipuAI
- Update ZhipuAI Chat Model to use ToolCalling Manager and ToolExecutionEligibilityPredicate - Update ZhipuAI ChatOptions to implement ToolCallingChatOptions - Update Autoconfiguration for ZhipuAI model to use ToolCallingAutoconfiguration - Update tests Signed-off-by: Ilayaperumal Gopinathan <ilayaperumal.gopinathan@broadcom.com>
This commit is contained in:
@@ -35,6 +35,12 @@
|
||||
|
||||
<!-- Spring AI auto configurations -->
|
||||
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-autoconfigure-model-tool</artifactId>
|
||||
<version>${project.parent.version}</version>
|
||||
</dependency>
|
||||
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-autoconfigure-retry</artifactId>
|
||||
|
||||
@@ -26,11 +26,16 @@ import org.springframework.ai.model.SpringAIModels;
|
||||
import org.springframework.ai.model.function.DefaultFunctionCallbackResolver;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.model.function.FunctionCallbackResolver;
|
||||
import org.springframework.ai.model.tool.DefaultToolExecutionEligibilityPredicate;
|
||||
import org.springframework.ai.model.tool.ToolCallingManager;
|
||||
import org.springframework.ai.model.tool.ToolExecutionEligibilityPredicate;
|
||||
import org.springframework.ai.model.tool.autoconfigure.ToolCallingAutoConfiguration;
|
||||
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatModel;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi;
|
||||
import org.springframework.beans.factory.ObjectProvider;
|
||||
import org.springframework.boot.autoconfigure.AutoConfiguration;
|
||||
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
|
||||
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
|
||||
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
|
||||
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
|
||||
@@ -50,28 +55,32 @@ import org.springframework.web.client.RestClient;
|
||||
* @author Geng Rong
|
||||
* @author Ilayaperumal Gopinathan
|
||||
*/
|
||||
@AutoConfiguration(after = { RestClientAutoConfiguration.class, SpringAiRetryAutoConfiguration.class })
|
||||
@AutoConfiguration(after = { RestClientAutoConfiguration.class, SpringAiRetryAutoConfiguration.class,
|
||||
ToolCallingAutoConfiguration.class })
|
||||
@ConditionalOnClass(ZhiPuAiApi.class)
|
||||
@ConditionalOnProperty(name = SpringAIModelProperties.CHAT_MODEL, havingValue = SpringAIModels.ZHIPUAI,
|
||||
matchIfMissing = true)
|
||||
@EnableConfigurationProperties({ ZhiPuAiConnectionProperties.class, ZhiPuAiChatProperties.class })
|
||||
@ImportAutoConfiguration(classes = { SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
|
||||
ToolCallingAutoConfiguration.class })
|
||||
public class ZhiPuAiChatAutoConfiguration {
|
||||
|
||||
@Bean
|
||||
@ConditionalOnMissingBean
|
||||
public ZhiPuAiChatModel zhiPuAiChatModel(ZhiPuAiConnectionProperties commonProperties,
|
||||
ZhiPuAiChatProperties chatProperties, ObjectProvider<RestClient.Builder> restClientBuilderProvider,
|
||||
List<FunctionCallback> toolFunctionCallbacks, FunctionCallbackResolver functionCallbackResolver,
|
||||
RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler,
|
||||
ObjectProvider<ObservationRegistry> observationRegistry,
|
||||
ObjectProvider<ChatModelObservationConvention> observationConvention) {
|
||||
ObjectProvider<ChatModelObservationConvention> observationConvention, ToolCallingManager toolCallingManager,
|
||||
ObjectProvider<ToolExecutionEligibilityPredicate> toolExecutionEligibilityPredicate) {
|
||||
|
||||
var zhiPuAiApi = zhiPuAiApi(chatProperties.getBaseUrl(), commonProperties.getBaseUrl(),
|
||||
chatProperties.getApiKey(), commonProperties.getApiKey(),
|
||||
restClientBuilderProvider.getIfAvailable(RestClient::builder), responseErrorHandler);
|
||||
|
||||
var chatModel = new ZhiPuAiChatModel(zhiPuAiApi, chatProperties.getOptions(), functionCallbackResolver,
|
||||
toolFunctionCallbacks, retryTemplate, observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP));
|
||||
var chatModel = new ZhiPuAiChatModel(zhiPuAiApi, chatProperties.getOptions(), toolCallingManager, retryTemplate,
|
||||
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
|
||||
toolExecutionEligibilityPredicate.getIfUnique(DefaultToolExecutionEligibilityPredicate::new));
|
||||
|
||||
observationConvention.ifAvailable(chatModel::setObservationConvention);
|
||||
|
||||
@@ -90,12 +99,4 @@ public class ZhiPuAiChatAutoConfiguration {
|
||||
return new ZhiPuAiApi(resolvedBaseUrl, resolvedApiKey, restClientBuilder, responseErrorHandler);
|
||||
}
|
||||
|
||||
@Bean
|
||||
@ConditionalOnMissingBean
|
||||
public FunctionCallbackResolver springAiFunctionManager(ApplicationContext context) {
|
||||
DefaultFunctionCallbackResolver manager = new DefaultFunctionCallbackResolver();
|
||||
manager.setApplicationContext(context);
|
||||
return manager;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -27,6 +27,7 @@ import reactor.core.publisher.Flux;
|
||||
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.embedding.EmbeddingResponse;
|
||||
import org.springframework.ai.image.ImagePrompt;
|
||||
@@ -58,8 +59,8 @@ public class ZhiPuAiAutoConfigurationIT {
|
||||
void generate() {
|
||||
this.contextRunner.withConfiguration(AutoConfigurations.of(ZhiPuAiChatAutoConfiguration.class)).run(context -> {
|
||||
ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class);
|
||||
String response = chatModel.call("Hello");
|
||||
assertThat(response).isNotEmpty();
|
||||
ChatResponse response = chatModel.call(new Prompt("Hello", ChatOptions.builder().build()));
|
||||
assertThat(response.getResult().getOutput().getText()).isNotEmpty();
|
||||
logger.info("Response: " + response);
|
||||
});
|
||||
}
|
||||
@@ -68,7 +69,8 @@ public class ZhiPuAiAutoConfigurationIT {
|
||||
void generateStreaming() {
|
||||
this.contextRunner.withConfiguration(AutoConfigurations.of(ZhiPuAiChatAutoConfiguration.class)).run(context -> {
|
||||
ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class);
|
||||
Flux<ChatResponse> responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello")));
|
||||
Flux<ChatResponse> responseFlux = chatModel
|
||||
.stream(new Prompt(new UserMessage("Hello"), ChatOptions.builder().build()));
|
||||
String response = responseFlux.collectList()
|
||||
.block()
|
||||
.stream()
|
||||
|
||||
@@ -33,6 +33,7 @@ import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.model.zhipuai.autoconfigure.ZhiPuAiChatAutoConfiguration;
|
||||
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
|
||||
import org.springframework.ai.tool.function.FunctionToolCallback;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatModel;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatOptions;
|
||||
import org.springframework.boot.autoconfigure.AutoConfigurations;
|
||||
@@ -64,8 +65,7 @@ public class FunctionCallbackInPromptIT {
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
var promptOptions = ZhiPuAiChatOptions.builder()
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function("CurrentWeatherService", new MockWeatherService())
|
||||
.toolCallbacks(List.of(FunctionToolCallback.builder("CurrentWeatherService", new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
// .responseConverter(response -> "" + response.temp() +
|
||||
@@ -92,8 +92,7 @@ public class FunctionCallbackInPromptIT {
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
var promptOptions = ZhiPuAiChatOptions.builder()
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function("CurrentWeatherService", new MockWeatherService())
|
||||
.toolCallbacks(List.of(FunctionToolCallback.builder("CurrentWeatherService", new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
.build()))
|
||||
|
||||
@@ -32,6 +32,7 @@ import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.function.FunctionCallingOptions;
|
||||
import org.springframework.ai.model.tool.ToolCallingChatOptions;
|
||||
import org.springframework.ai.model.zhipuai.autoconfigure.ZhiPuAiChatAutoConfiguration;
|
||||
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatModel;
|
||||
@@ -69,8 +70,8 @@ class FunctionCallbackWithPlainFunctionBeanIT {
|
||||
UserMessage userMessage = new UserMessage(
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
ChatResponse response = chatModel.call(
|
||||
new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().function("weatherFunction").build()));
|
||||
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
|
||||
ZhiPuAiChatOptions.builder().toolNames("weatherFunction").build()));
|
||||
|
||||
logger.info("Response: {}", response);
|
||||
|
||||
@@ -78,7 +79,7 @@ class FunctionCallbackWithPlainFunctionBeanIT {
|
||||
|
||||
// Test weatherFunctionTwo
|
||||
response = chatModel.call(new Prompt(List.of(userMessage),
|
||||
ZhiPuAiChatOptions.builder().function("weatherFunctionTwo").build()));
|
||||
ZhiPuAiChatOptions.builder().toolNames("weatherFunctionTwo").build()));
|
||||
|
||||
logger.info("Response: {}", response);
|
||||
|
||||
@@ -97,8 +98,8 @@ class FunctionCallbackWithPlainFunctionBeanIT {
|
||||
UserMessage userMessage = new UserMessage(
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
FunctionCallingOptions functionOptions = FunctionCallingOptions.builder()
|
||||
.function("weatherFunction")
|
||||
ToolCallingChatOptions functionOptions = ToolCallingChatOptions.builder()
|
||||
.toolNames("weatherFunction")
|
||||
.build();
|
||||
|
||||
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions));
|
||||
@@ -117,8 +118,8 @@ class FunctionCallbackWithPlainFunctionBeanIT {
|
||||
UserMessage userMessage = new UserMessage(
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
Flux<ChatResponse> response = chatModel.stream(
|
||||
new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().function("weatherFunction").build()));
|
||||
Flux<ChatResponse> response = chatModel.stream(new Prompt(List.of(userMessage),
|
||||
ZhiPuAiChatOptions.builder().toolNames("weatherFunction").build()));
|
||||
|
||||
String content = response.collectList()
|
||||
.block()
|
||||
@@ -136,7 +137,7 @@ class FunctionCallbackWithPlainFunctionBeanIT {
|
||||
|
||||
// Test weatherFunctionTwo
|
||||
response = chatModel.stream(new Prompt(List.of(userMessage),
|
||||
ZhiPuAiChatOptions.builder().function("weatherFunctionTwo").build()));
|
||||
ZhiPuAiChatOptions.builder().toolNames("weatherFunctionTwo").build()));
|
||||
|
||||
content = response.collectList()
|
||||
.block()
|
||||
|
||||
@@ -33,6 +33,8 @@ import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.model.zhipuai.autoconfigure.ZhiPuAiChatAutoConfiguration;
|
||||
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.function.FunctionToolCallback;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatModel;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatOptions;
|
||||
import org.springframework.boot.autoconfigure.AutoConfigurations;
|
||||
@@ -67,7 +69,7 @@ public class ZhipuAiFunctionCallbackIT {
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
ChatResponse response = chatModel
|
||||
.call(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().function("WeatherInfo").build()));
|
||||
.call(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().toolNames("WeatherInfo").build()));
|
||||
|
||||
logger.info("Response: {}", response);
|
||||
|
||||
@@ -85,8 +87,8 @@ public class ZhipuAiFunctionCallbackIT {
|
||||
UserMessage userMessage = new UserMessage(
|
||||
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
|
||||
|
||||
Flux<ChatResponse> response = chatModel
|
||||
.stream(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().function("WeatherInfo").build()));
|
||||
Flux<ChatResponse> response = chatModel.stream(
|
||||
new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().toolNames("WeatherInfo").build()));
|
||||
|
||||
String content = response.collectList()
|
||||
.block()
|
||||
@@ -109,10 +111,9 @@ public class ZhipuAiFunctionCallbackIT {
|
||||
static class Config {
|
||||
|
||||
@Bean
|
||||
public FunctionCallback weatherFunctionInfo() {
|
||||
public ToolCallback weatherFunctionInfo() {
|
||||
|
||||
return FunctionCallback.builder()
|
||||
.function("WeatherInfo", new MockWeatherService())
|
||||
return FunctionToolCallback.builder("WeatherInfo", new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
// .responseConverter(response -> "" + response.temp() + response.unit())
|
||||
|
||||
@@ -18,10 +18,8 @@ package org.springframework.ai.zhipuai;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Base64;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
import io.micrometer.observation.Observation;
|
||||
@@ -41,7 +39,6 @@ import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
|
||||
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
|
||||
import org.springframework.ai.chat.metadata.DefaultUsage;
|
||||
import org.springframework.ai.chat.metadata.EmptyUsage;
|
||||
import org.springframework.ai.chat.model.AbstractToolCallSupport;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
@@ -54,15 +51,17 @@ import org.springframework.ai.chat.observation.DefaultChatModelObservationConven
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.ModelOptionsUtils;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.model.function.FunctionCallbackResolver;
|
||||
import org.springframework.ai.model.function.FunctionCallingOptions;
|
||||
import org.springframework.ai.model.tool.DefaultToolExecutionEligibilityPredicate;
|
||||
import org.springframework.ai.model.tool.ToolCallingChatOptions;
|
||||
import org.springframework.ai.model.tool.ToolCallingManager;
|
||||
import org.springframework.ai.model.tool.ToolExecutionEligibilityPredicate;
|
||||
import org.springframework.ai.model.tool.ToolExecutionResult;
|
||||
import org.springframework.ai.retry.RetryUtils;
|
||||
import org.springframework.ai.tool.definition.ToolDefinition;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletion;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletion.Choice;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionChunk;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionFinishReason;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionMessage;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionMessage.ChatCompletionFunction;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionMessage.MediaContent;
|
||||
@@ -81,13 +80,14 @@ import org.springframework.util.MimeType;
|
||||
* backed by {@link ZhiPuAiApi}.
|
||||
*
|
||||
* @author Geng Rong
|
||||
* @author Alexandros Pappas
|
||||
* @author Ilayaperumal Gopinathan
|
||||
* @see ChatModel
|
||||
* @see StreamingChatModel
|
||||
* @see ZhiPuAiApi
|
||||
* @author Alexandros Pappas
|
||||
* @since 1.0.0 M1
|
||||
*/
|
||||
public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatModel, StreamingChatModel {
|
||||
public class ZhiPuAiChatModel implements ChatModel {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiChatModel.class);
|
||||
|
||||
@@ -113,11 +113,22 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
*/
|
||||
private final ObservationRegistry observationRegistry;
|
||||
|
||||
/**
|
||||
* The Tool calling manager.
|
||||
*/
|
||||
private final ToolCallingManager toolCallingManager;
|
||||
|
||||
/**
|
||||
* Conventions to use for generating observations.
|
||||
*/
|
||||
private ChatModelObservationConvention observationConvention = DEFAULT_OBSERVATION_CONVENTION;
|
||||
|
||||
/**
|
||||
* The tool execution eligibility predicate used to determine if a tool can be
|
||||
* executed.
|
||||
*/
|
||||
private final ToolExecutionEligibilityPredicate toolExecutionEligibilityPredicate;
|
||||
|
||||
/**
|
||||
* Creates an instance of the ZhiPuAiChatModel.
|
||||
* @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the
|
||||
@@ -135,7 +146,7 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
* @param options The ZhiPuAiChatOptions to configure the chat model.
|
||||
*/
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options) {
|
||||
this(zhiPuAiApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE);
|
||||
this(zhiPuAiApi, options, RetryUtils.DEFAULT_RETRY_TEMPLATE);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -143,12 +154,40 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
* @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the
|
||||
* ZhiPuAI Chat API.
|
||||
* @param options The ZhiPuAiChatOptions to configure the chat model.
|
||||
* @param functionCallbackResolver The function callback resolver.
|
||||
* @param retryTemplate The retry template.
|
||||
*/
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options,
|
||||
FunctionCallbackResolver functionCallbackResolver, RetryTemplate retryTemplate) {
|
||||
this(zhiPuAiApi, options, functionCallbackResolver, List.of(), retryTemplate, ObservationRegistry.NOOP);
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, RetryTemplate retryTemplate) {
|
||||
this(zhiPuAiApi, options, ToolCallingManager.builder().build(), retryTemplate, ObservationRegistry.NOOP,
|
||||
new DefaultToolExecutionEligibilityPredicate());
|
||||
}
|
||||
|
||||
/**
|
||||
* Initializes an instance of the ZhiPuAiChatModel.
|
||||
* @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the
|
||||
* ZhiPuAI Chat API.
|
||||
* @param options The ZhiPuAiChatOptions to configure the chat model.
|
||||
* @param retryTemplate The retry template.
|
||||
* @param observationRegistry The Observation Registry.
|
||||
*/
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, RetryTemplate retryTemplate,
|
||||
ObservationRegistry observationRegistry) {
|
||||
this(zhiPuAiApi, options, ToolCallingManager.builder().build(), retryTemplate, observationRegistry,
|
||||
new DefaultToolExecutionEligibilityPredicate());
|
||||
}
|
||||
|
||||
/**
|
||||
* Initializes an instance of the ZhiPuAiChatModel.
|
||||
* @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the
|
||||
* ZhiPuAI Chat API.
|
||||
* @param toolCallingManager The tool calling manager
|
||||
* @param options The ZhiPuAiChatOptions to configure the chat model.
|
||||
* @param retryTemplate The retry template.
|
||||
* @param observationRegistry The Observation Registry.
|
||||
*/
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, ToolCallingManager toolCallingManager,
|
||||
RetryTemplate retryTemplate, ObservationRegistry observationRegistry) {
|
||||
this(zhiPuAiApi, options, toolCallingManager, retryTemplate, observationRegistry,
|
||||
new DefaultToolExecutionEligibilityPredicate());
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -156,25 +195,26 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
* @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the
|
||||
* ZhiPuAI Chat API.
|
||||
* @param options The ZhiPuAiChatOptions to configure the chat model.
|
||||
* @param functionCallbackResolver The function callback resolver.
|
||||
* @param toolFunctionCallbacks The tool function callbacks.
|
||||
* @param toolCallingManager The tool calling manager
|
||||
* @param retryTemplate The retry template.
|
||||
* @param observationRegistry The ObservationRegistry used for instrumentation.
|
||||
* @param toolExecutionEligibilityPredicate The Tool execution eligibility predicate.
|
||||
*/
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options,
|
||||
FunctionCallbackResolver functionCallbackResolver, List<FunctionCallback> toolFunctionCallbacks,
|
||||
RetryTemplate retryTemplate, ObservationRegistry observationRegistry) {
|
||||
super(functionCallbackResolver, options, toolFunctionCallbacks);
|
||||
public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, ToolCallingManager toolCallingManager,
|
||||
RetryTemplate retryTemplate, ObservationRegistry observationRegistry,
|
||||
ToolExecutionEligibilityPredicate toolExecutionEligibilityPredicate) {
|
||||
Assert.notNull(zhiPuAiApi, "ZhiPuAiApi must not be null");
|
||||
Assert.notNull(options, "Options must not be null");
|
||||
Assert.notNull(retryTemplate, "RetryTemplate must not be null");
|
||||
Assert.isTrue(CollectionUtils.isEmpty(options.getFunctionCallbacks()),
|
||||
"The default function callbacks must be set via the toolFunctionCallbacks constructor parameter");
|
||||
Assert.notNull(observationRegistry, "ObservationRegistry must not be null");
|
||||
Assert.notNull(toolCallingManager, "toolCallingManager cannot be null");
|
||||
Assert.notNull(toolExecutionEligibilityPredicate, "toolExecutionEligibilityPredicate cannot be null");
|
||||
this.zhiPuAiApi = zhiPuAiApi;
|
||||
this.defaultOptions = options;
|
||||
this.retryTemplate = retryTemplate;
|
||||
this.observationRegistry = observationRegistry;
|
||||
this.toolCallingManager = toolCallingManager;
|
||||
this.toolExecutionEligibilityPredicate = toolExecutionEligibilityPredicate;
|
||||
}
|
||||
|
||||
private static Generation buildGeneration(Choice choice, Map<String, Object> metadata) {
|
||||
@@ -194,12 +234,15 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
|
||||
@Override
|
||||
public ChatResponse call(Prompt prompt) {
|
||||
ChatCompletionRequest request = createRequest(prompt, false);
|
||||
// Before moving any further, build the final request Prompt,
|
||||
// merging runtime and default options.
|
||||
Prompt requestPrompt = buildRequestPrompt(prompt);
|
||||
ChatCompletionRequest request = createRequest(requestPrompt, false);
|
||||
|
||||
ChatModelObservationContext observationContext = ChatModelObservationContext.builder()
|
||||
.prompt(prompt)
|
||||
.prompt(requestPrompt)
|
||||
.provider(ZhiPuApiConstants.PROVIDER_NAME)
|
||||
.requestOptions(buildRequestOptions(request))
|
||||
.requestOptions(prompt.getOptions())
|
||||
.build();
|
||||
|
||||
ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION
|
||||
@@ -236,14 +279,20 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
|
||||
return chatResponse;
|
||||
});
|
||||
if (!isProxyToolCalls(prompt, this.defaultOptions) && isToolCall(response,
|
||||
Set.of(ChatCompletionFinishReason.TOOL_CALLS.name(), ChatCompletionFinishReason.STOP.name()))) {
|
||||
var toolCallConversation = handleToolCalls(prompt, response);
|
||||
// Recursively call the call method with the tool call message
|
||||
// conversation that contains the call responses.
|
||||
return this.call(new Prompt(toolCallConversation, prompt.getOptions()));
|
||||
if (this.toolExecutionEligibilityPredicate.isToolExecutionRequired(requestPrompt.getOptions(), response)) {
|
||||
var toolExecutionResult = this.toolCallingManager.executeToolCalls(requestPrompt, response);
|
||||
if (toolExecutionResult.returnDirect()) {
|
||||
// Return tool execution result directly to the client.
|
||||
return ChatResponse.builder()
|
||||
.from(response)
|
||||
.generations(ToolExecutionResult.buildGenerations(toolExecutionResult))
|
||||
.build();
|
||||
}
|
||||
else {
|
||||
// Send the tool execution result back to the model.
|
||||
return this.call(new Prompt(toolExecutionResult.conversationHistory(), requestPrompt.getOptions()));
|
||||
}
|
||||
}
|
||||
|
||||
return response;
|
||||
}
|
||||
|
||||
@@ -255,7 +304,10 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
@Override
|
||||
public Flux<ChatResponse> stream(Prompt prompt) {
|
||||
return Flux.deferContextual(contextView -> {
|
||||
ChatCompletionRequest request = createRequest(prompt, true);
|
||||
// Before moving any further, build the final request Prompt,
|
||||
// merging runtime and default options.
|
||||
Prompt requestPrompt = buildRequestPrompt(prompt);
|
||||
ChatCompletionRequest request = createRequest(requestPrompt, true);
|
||||
|
||||
Flux<ChatCompletionChunk> completionChunks = this.retryTemplate
|
||||
.execute(ctx -> this.zhiPuAiApi.chatCompletionStream(request));
|
||||
@@ -265,7 +317,7 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
ConcurrentHashMap<String, String> roleMap = new ConcurrentHashMap<>();
|
||||
|
||||
final ChatModelObservationContext observationContext = ChatModelObservationContext.builder()
|
||||
.prompt(prompt)
|
||||
.prompt(requestPrompt)
|
||||
.provider(ZhiPuApiConstants.PROVIDER_NAME)
|
||||
.requestOptions(buildRequestOptions(request))
|
||||
.build();
|
||||
@@ -306,17 +358,24 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
|
||||
// @formatter:off
|
||||
Flux<ChatResponse> flux = chatResponse.flatMap(response -> {
|
||||
if (!isProxyToolCalls(prompt, this.defaultOptions) && isToolCall(response, Set.of(ChatCompletionFinishReason.TOOL_CALLS.name(), ChatCompletionFinishReason.STOP.name()))) {
|
||||
// FIXME: bounded elastic needs to be used since tool calling
|
||||
// is currently only synchronous
|
||||
return Flux.defer(() -> {
|
||||
var toolCallConversation = handleToolCalls(prompt, response);
|
||||
// Recursively call the stream method with the tool call message
|
||||
// conversation that contains the call responses.
|
||||
return this.stream(new Prompt(toolCallConversation, prompt.getOptions()));
|
||||
}).subscribeOn(Schedulers.boundedElastic());
|
||||
}
|
||||
return Flux.just(response);
|
||||
if (this.toolExecutionEligibilityPredicate.isToolExecutionRequired(requestPrompt.getOptions(), response)) {
|
||||
return Flux.defer(() -> {
|
||||
// FIXME: bounded elastic needs to be used since tool calling
|
||||
// is currently only synchronous
|
||||
var toolExecutionResult = this.toolCallingManager.executeToolCalls(requestPrompt, response);
|
||||
if (toolExecutionResult.returnDirect()) {
|
||||
// Return tool execution result directly to the client.
|
||||
return Flux.just(ChatResponse.builder().from(response)
|
||||
.generations(ToolExecutionResult.buildGenerations(toolExecutionResult))
|
||||
.build());
|
||||
}
|
||||
else {
|
||||
// Send the tool execution result back to the model.
|
||||
return this.stream(new Prompt(toolExecutionResult.conversationHistory(), requestPrompt.getOptions()));
|
||||
}
|
||||
}).subscribeOn(Schedulers.boundedElastic());
|
||||
}
|
||||
return Flux.just(response);
|
||||
})
|
||||
.doOnError(observation::error)
|
||||
.doFinally(s -> observation.stop())
|
||||
@@ -360,6 +419,57 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
"chat.completion", null);
|
||||
}
|
||||
|
||||
private List<ZhiPuAiApi.FunctionTool> getFunctionTools(List<ToolDefinition> toolDefinitions) {
|
||||
return toolDefinitions.stream().map(toolDefinition -> {
|
||||
var function = new ZhiPuAiApi.FunctionTool.Function(toolDefinition.description(), toolDefinition.name(),
|
||||
toolDefinition.inputSchema());
|
||||
return new ZhiPuAiApi.FunctionTool(function);
|
||||
}).toList();
|
||||
}
|
||||
|
||||
Prompt buildRequestPrompt(Prompt prompt) {
|
||||
// Process runtime options
|
||||
ZhiPuAiChatOptions runtimeOptions = null;
|
||||
if (prompt.getOptions() != null) {
|
||||
if (prompt.getOptions() instanceof ToolCallingChatOptions toolCallingChatOptions) {
|
||||
runtimeOptions = ModelOptionsUtils.copyToTarget(toolCallingChatOptions, ToolCallingChatOptions.class,
|
||||
ZhiPuAiChatOptions.class);
|
||||
}
|
||||
else {
|
||||
runtimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
|
||||
ZhiPuAiChatOptions.class);
|
||||
}
|
||||
}
|
||||
|
||||
// Define request options by merging runtime options and default options
|
||||
ZhiPuAiChatOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.defaultOptions,
|
||||
ZhiPuAiChatOptions.class);
|
||||
|
||||
// Merge @JsonIgnore-annotated options explicitly since they are ignored by
|
||||
// Jackson, used by ModelOptionsUtils.
|
||||
if (runtimeOptions != null) {
|
||||
requestOptions.setInternalToolExecutionEnabled(
|
||||
ModelOptionsUtils.mergeOption(runtimeOptions.getInternalToolExecutionEnabled(),
|
||||
this.defaultOptions.getInternalToolExecutionEnabled()));
|
||||
requestOptions.setToolNames(ToolCallingChatOptions.mergeToolNames(runtimeOptions.getToolNames(),
|
||||
this.defaultOptions.getToolNames()));
|
||||
requestOptions.setToolCallbacks(ToolCallingChatOptions.mergeToolCallbacks(runtimeOptions.getToolCallbacks(),
|
||||
this.defaultOptions.getToolCallbacks()));
|
||||
requestOptions.setToolContext(ToolCallingChatOptions.mergeToolContext(runtimeOptions.getToolContext(),
|
||||
this.defaultOptions.getToolContext()));
|
||||
}
|
||||
else {
|
||||
requestOptions.setInternalToolExecutionEnabled(this.defaultOptions.getInternalToolExecutionEnabled());
|
||||
requestOptions.setToolNames(this.defaultOptions.getToolNames());
|
||||
requestOptions.setToolCallbacks(this.defaultOptions.getToolCallbacks());
|
||||
requestOptions.setToolContext(this.defaultOptions.getToolContext());
|
||||
}
|
||||
|
||||
ToolCallingChatOptions.validateToolCallbacks(requestOptions.getToolCallbacks());
|
||||
|
||||
return new Prompt(prompt.getInstructions(), requestOptions);
|
||||
}
|
||||
|
||||
/**
|
||||
* Accessible for testing.
|
||||
*/
|
||||
@@ -416,37 +526,31 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
|
||||
ChatCompletionRequest request = new ChatCompletionRequest(chatCompletionMessages, stream);
|
||||
|
||||
Set<String> enabledToolsToUse = new HashSet<>();
|
||||
|
||||
if (prompt.getOptions() != null) {
|
||||
ZhiPuAiChatOptions updatedRuntimeOptions;
|
||||
if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) {
|
||||
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions,
|
||||
FunctionCallingOptions.class, ZhiPuAiChatOptions.class);
|
||||
if (prompt.getOptions() instanceof ToolCallingChatOptions toolCallingChatOptions) {
|
||||
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(toolCallingChatOptions,
|
||||
ToolCallingChatOptions.class, ZhiPuAiChatOptions.class);
|
||||
}
|
||||
else {
|
||||
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
|
||||
ZhiPuAiChatOptions.class);
|
||||
}
|
||||
|
||||
enabledToolsToUse.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions));
|
||||
|
||||
request = ModelOptionsUtils.merge(updatedRuntimeOptions, request, ChatCompletionRequest.class);
|
||||
}
|
||||
|
||||
if (!CollectionUtils.isEmpty(this.defaultOptions.getFunctions())) {
|
||||
enabledToolsToUse.addAll(this.defaultOptions.getFunctions());
|
||||
}
|
||||
|
||||
request = ModelOptionsUtils.merge(request, this.defaultOptions, ChatCompletionRequest.class);
|
||||
|
||||
if (!CollectionUtils.isEmpty(enabledToolsToUse)) {
|
||||
ZhiPuAiChatOptions requestOptions = (ZhiPuAiChatOptions) prompt.getOptions();
|
||||
|
||||
// Add the tool definitions to the request's tools parameter.
|
||||
List<ToolDefinition> toolDefinitions = this.toolCallingManager.resolveToolDefinitions(requestOptions);
|
||||
if (!CollectionUtils.isEmpty(toolDefinitions)) {
|
||||
request = ModelOptionsUtils.merge(
|
||||
ZhiPuAiChatOptions.builder().tools(this.getFunctionTools(enabledToolsToUse)).build(), request,
|
||||
ZhiPuAiChatOptions.builder().tools(this.getFunctionTools(toolDefinitions)).build(), request,
|
||||
ChatCompletionRequest.class);
|
||||
}
|
||||
|
||||
return request;
|
||||
}
|
||||
|
||||
@@ -476,14 +580,6 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
|
||||
.build();
|
||||
}
|
||||
|
||||
private List<ZhiPuAiApi.FunctionTool> getFunctionTools(Set<String> functionNames) {
|
||||
return this.resolveFunctionCallbacks(functionNames).stream().map(functionCallback -> {
|
||||
var function = new ZhiPuAiApi.FunctionTool.Function(functionCallback.getDescription(),
|
||||
functionCallback.getName(), functionCallback.getInputTypeSchema());
|
||||
return new ZhiPuAiApi.FunctionTool(function);
|
||||
}).toList();
|
||||
}
|
||||
|
||||
public void setObservationConvention(ChatModelObservationConvention observationConvention) {
|
||||
this.observationConvention = observationConvention;
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
package org.springframework.ai.zhipuai;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Arrays;
|
||||
import java.util.HashMap;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
@@ -29,9 +30,10 @@ import com.fasterxml.jackson.annotation.JsonInclude.Include;
|
||||
import com.fasterxml.jackson.annotation.JsonProperty;
|
||||
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.model.function.FunctionCallingOptions;
|
||||
import org.springframework.ai.model.tool.ToolCallingChatOptions;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.Assert;
|
||||
|
||||
/**
|
||||
@@ -43,7 +45,7 @@ import org.springframework.util.Assert;
|
||||
* @since 1.0.0 M1
|
||||
*/
|
||||
@JsonInclude(Include.NON_NULL)
|
||||
public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
public class ZhiPuAiChatOptions implements ToolCallingChatOptions {
|
||||
|
||||
// @formatter:off
|
||||
/**
|
||||
@@ -106,31 +108,25 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
private @JsonProperty("do_sample") Boolean doSample;
|
||||
|
||||
/**
|
||||
* ZhiPuAI Tool Function Callbacks to register with the ChatModel.
|
||||
* For Prompt Options the functionCallbacks are automatically enabled for the duration of the prompt execution.
|
||||
* For Default Options the functionCallbacks are registered but disabled by default. Use the enableFunctions to set the functions
|
||||
* from the registry to be used by the ChatModel chat completion requests.
|
||||
* Collection of {@link ToolCallback}s to be used for tool calling in the chat completion requests.
|
||||
*/
|
||||
@JsonIgnore
|
||||
private List<FunctionCallback> functionCallbacks = new ArrayList<>();
|
||||
private List<ToolCallback> toolCallbacks = new ArrayList<>();
|
||||
|
||||
/**
|
||||
* List of functions, identified by their names, to configure for function calling in
|
||||
* the chat completion requests.
|
||||
* Functions with those names must exist in the functionCallbacks registry.
|
||||
* The {@link #functionCallbacks} from the PromptOptions are automatically enabled for the duration of the prompt execution.
|
||||
*
|
||||
* Note that function enabled with the default options are enabled for all chat completion requests. This could impact the token count and the billing.
|
||||
* If the functions is set in a prompt options, then the enabled functions are only active for the duration of this prompt execution.
|
||||
* Collection of tool names to be resolved at runtime and used for tool calling in the chat completion requests.
|
||||
*/
|
||||
@JsonIgnore
|
||||
private Set<String> functions = new HashSet<>();
|
||||
private Set<String> toolNames = new HashSet<>();
|
||||
|
||||
/**
|
||||
* Whether to enable the tool execution lifecycle internally in ChatModel.
|
||||
*/
|
||||
@JsonIgnore
|
||||
private Boolean internalToolExecutionEnabled;
|
||||
|
||||
@JsonIgnore
|
||||
private Boolean proxyToolCalls;
|
||||
|
||||
@JsonIgnore
|
||||
private Map<String, Object> toolContext;
|
||||
private Map<String, Object> toolContext = new HashMap<>();
|
||||
// @formatter:on
|
||||
|
||||
public static Builder builder() {
|
||||
@@ -149,9 +145,9 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
.user(fromOptions.getUser())
|
||||
.requestId(fromOptions.getRequestId())
|
||||
.doSample(fromOptions.getDoSample())
|
||||
.functionCallbacks(fromOptions.getFunctionCallbacks())
|
||||
.functions(fromOptions.getFunctions())
|
||||
.proxyToolCalls(fromOptions.getProxyToolCalls())
|
||||
.toolCallbacks(fromOptions.getToolCallbacks())
|
||||
.toolNames(fromOptions.getToolNames())
|
||||
.internalToolExecutionEnabled(fromOptions.getInternalToolExecutionEnabled())
|
||||
.toolContext(fromOptions.getToolContext())
|
||||
.build();
|
||||
}
|
||||
@@ -251,25 +247,6 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
this.doSample = doSample;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<FunctionCallback> getFunctionCallbacks() {
|
||||
return this.functionCallbacks;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void setFunctionCallbacks(List<FunctionCallback> functionCallbacks) {
|
||||
this.functionCallbacks = functionCallbacks;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Set<String> getFunctions() {
|
||||
return this.functions;
|
||||
}
|
||||
|
||||
public void setFunctions(Set<String> functionNames) {
|
||||
this.functions = functionNames;
|
||||
}
|
||||
|
||||
@Override
|
||||
@JsonIgnore
|
||||
public Double getFrequencyPenalty() {
|
||||
@@ -289,12 +266,45 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
}
|
||||
|
||||
@Override
|
||||
public Boolean getProxyToolCalls() {
|
||||
return this.proxyToolCalls;
|
||||
@JsonIgnore
|
||||
public List<ToolCallback> getToolCallbacks() {
|
||||
return this.toolCallbacks;
|
||||
}
|
||||
|
||||
public void setProxyToolCalls(Boolean proxyToolCalls) {
|
||||
this.proxyToolCalls = proxyToolCalls;
|
||||
@Override
|
||||
@JsonIgnore
|
||||
public void setToolCallbacks(List<ToolCallback> toolCallbacks) {
|
||||
Assert.notNull(toolCallbacks, "toolCallbacks cannot be null");
|
||||
Assert.noNullElements(toolCallbacks, "toolCallbacks cannot contain null elements");
|
||||
this.toolCallbacks = toolCallbacks;
|
||||
}
|
||||
|
||||
@Override
|
||||
@JsonIgnore
|
||||
public Set<String> getToolNames() {
|
||||
return this.toolNames;
|
||||
}
|
||||
|
||||
@Override
|
||||
@JsonIgnore
|
||||
public void setToolNames(Set<String> toolNames) {
|
||||
Assert.notNull(toolNames, "toolNames cannot be null");
|
||||
Assert.noNullElements(toolNames, "toolNames cannot contain null elements");
|
||||
toolNames.forEach(tool -> Assert.hasText(tool, "toolNames cannot contain empty elements"));
|
||||
this.toolNames = toolNames;
|
||||
}
|
||||
|
||||
@Override
|
||||
@Nullable
|
||||
@JsonIgnore
|
||||
public Boolean getInternalToolExecutionEnabled() {
|
||||
return this.internalToolExecutionEnabled;
|
||||
}
|
||||
|
||||
@Override
|
||||
@JsonIgnore
|
||||
public void setInternalToolExecutionEnabled(@Nullable Boolean internalToolExecutionEnabled) {
|
||||
this.internalToolExecutionEnabled = internalToolExecutionEnabled;
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -319,7 +329,10 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
result = prime * result + ((this.tools == null) ? 0 : this.tools.hashCode());
|
||||
result = prime * result + ((this.toolChoice == null) ? 0 : this.toolChoice.hashCode());
|
||||
result = prime * result + ((this.user == null) ? 0 : this.user.hashCode());
|
||||
result = prime * result + ((this.proxyToolCalls == null) ? 0 : this.proxyToolCalls.hashCode());
|
||||
result = prime * result
|
||||
+ ((this.internalToolExecutionEnabled == null) ? 0 : this.internalToolExecutionEnabled.hashCode());
|
||||
result = prime * result + ((this.toolCallbacks == null) ? 0 : this.toolCallbacks.hashCode());
|
||||
result = prime * result + ((this.toolNames == null) ? 0 : this.toolNames.hashCode());
|
||||
result = prime * result + ((this.toolContext == null) ? 0 : this.toolContext.hashCode());
|
||||
return result;
|
||||
}
|
||||
@@ -416,12 +429,12 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
else if (!this.doSample.equals(other.doSample)) {
|
||||
return false;
|
||||
}
|
||||
if (this.proxyToolCalls == null) {
|
||||
if (other.proxyToolCalls != null) {
|
||||
if (this.internalToolExecutionEnabled == null) {
|
||||
if (other.internalToolExecutionEnabled != null) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
else if (!this.proxyToolCalls.equals(other.proxyToolCalls)) {
|
||||
else if (!this.internalToolExecutionEnabled.equals(other.internalToolExecutionEnabled)) {
|
||||
return false;
|
||||
}
|
||||
if (this.toolContext == null) {
|
||||
@@ -440,7 +453,7 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
return fromOptions(this);
|
||||
}
|
||||
|
||||
public FunctionCallingOptions merge(ChatOptions options) {
|
||||
public ToolCallingChatOptions merge(ChatOptions options) {
|
||||
ZhiPuAiChatOptions.Builder builder = ZhiPuAiChatOptions.builder();
|
||||
|
||||
// Merge chat-specific options
|
||||
@@ -450,44 +463,43 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
.temperature(options.getTemperature() != null ? options.getTemperature() : this.getTemperature())
|
||||
.topP(options.getTopP() != null ? options.getTopP() : this.getTopP());
|
||||
|
||||
// Try to get function-specific properties if options is a FunctionCallingOptions
|
||||
if (options instanceof FunctionCallingOptions functionOptions) {
|
||||
builder.proxyToolCalls(functionOptions.getProxyToolCalls() != null ? functionOptions.getProxyToolCalls()
|
||||
: this.proxyToolCalls);
|
||||
// Try to get tool-specific properties if options is a ToolCallingChatOptions
|
||||
if (options instanceof ToolCallingChatOptions toolCallingChatOptions) {
|
||||
builder.internalToolExecutionEnabled(toolCallingChatOptions.getInternalToolExecutionEnabled() != null
|
||||
? (toolCallingChatOptions).getInternalToolExecutionEnabled()
|
||||
: this.getInternalToolExecutionEnabled());
|
||||
|
||||
Set<String> functions = new HashSet<>();
|
||||
if (this.functions != null) {
|
||||
functions.addAll(this.functions);
|
||||
Set<String> toolNames = new HashSet<>();
|
||||
if (this.toolNames != null) {
|
||||
toolNames.addAll(this.toolNames);
|
||||
}
|
||||
if (functionOptions.getFunctions() != null) {
|
||||
functions.addAll(functionOptions.getFunctions());
|
||||
if (toolCallingChatOptions.getToolNames() != null) {
|
||||
toolNames.addAll(toolCallingChatOptions.getToolNames());
|
||||
}
|
||||
builder.functions(functions);
|
||||
builder.toolNames(toolNames);
|
||||
|
||||
List<FunctionCallback> functionCallbacks = new ArrayList<>();
|
||||
if (this.functionCallbacks != null) {
|
||||
functionCallbacks.addAll(this.functionCallbacks);
|
||||
List<ToolCallback> toolCallbacks = new ArrayList<>();
|
||||
if (this.toolCallbacks != null) {
|
||||
toolCallbacks.addAll(this.toolCallbacks);
|
||||
}
|
||||
if (functionOptions.getFunctionCallbacks() != null) {
|
||||
functionCallbacks.addAll(functionOptions.getFunctionCallbacks());
|
||||
if (toolCallingChatOptions.getToolCallbacks() != null) {
|
||||
toolCallbacks.addAll(toolCallingChatOptions.getToolCallbacks());
|
||||
}
|
||||
builder.functionCallbacks(functionCallbacks);
|
||||
builder.toolCallbacks(toolCallbacks);
|
||||
|
||||
Map<String, Object> context = new HashMap<>();
|
||||
if (this.toolContext != null) {
|
||||
context.putAll(this.toolContext);
|
||||
}
|
||||
if (functionOptions.getToolContext() != null) {
|
||||
context.putAll(functionOptions.getToolContext());
|
||||
if (toolCallingChatOptions.getToolContext() != null) {
|
||||
context.putAll(toolCallingChatOptions.getToolContext());
|
||||
}
|
||||
builder.toolContext(context);
|
||||
}
|
||||
else {
|
||||
// If not a FunctionCallingOptions, preserve current function-specific
|
||||
// properties
|
||||
builder.proxyToolCalls(this.proxyToolCalls);
|
||||
builder.functions(this.functions != null ? new HashSet<>(this.functions) : null);
|
||||
builder.functionCallbacks(this.functionCallbacks != null ? new ArrayList<>(this.functionCallbacks) : null);
|
||||
builder.internalToolExecutionEnabled(this.internalToolExecutionEnabled);
|
||||
builder.toolNames(this.toolNames != null ? new HashSet<>(this.toolNames) : null);
|
||||
builder.toolCallbacks(this.toolCallbacks != null ? new ArrayList<>(this.toolCallbacks) : null);
|
||||
builder.toolContext(this.toolContext != null ? new HashMap<>(this.toolContext) : null);
|
||||
}
|
||||
|
||||
@@ -563,25 +575,31 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions {
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder functionCallbacks(List<FunctionCallback> functionCallbacks) {
|
||||
this.options.functionCallbacks = functionCallbacks;
|
||||
public Builder toolCallbacks(List<ToolCallback> toolCallbacks) {
|
||||
this.options.setToolCallbacks(toolCallbacks);
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder functions(Set<String> functionNames) {
|
||||
Assert.notNull(functionNames, "Function names must not be null");
|
||||
this.options.functions = functionNames;
|
||||
public Builder toolCallbacks(ToolCallback... toolCallbacks) {
|
||||
Assert.notNull(toolCallbacks, "toolCallbacks cannot be null");
|
||||
this.options.toolCallbacks.addAll(Arrays.asList(toolCallbacks));
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder function(String functionName) {
|
||||
Assert.hasText(functionName, "Function name must not be empty");
|
||||
this.options.functions.add(functionName);
|
||||
public Builder toolNames(Set<String> toolNames) {
|
||||
Assert.notNull(toolNames, "toolNames cannot be null");
|
||||
this.options.setToolNames(toolNames);
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder proxyToolCalls(Boolean proxyToolCalls) {
|
||||
this.options.proxyToolCalls = proxyToolCalls;
|
||||
public Builder toolNames(String... toolNames) {
|
||||
Assert.notNull(toolNames, "toolNames cannot be null");
|
||||
this.options.toolNames.addAll(Set.of(toolNames));
|
||||
return this;
|
||||
}
|
||||
|
||||
public Builder internalToolExecutionEnabled(@Nullable Boolean internalToolExecutionEnabled) {
|
||||
this.options.setInternalToolExecutionEnabled(internalToolExecutionEnabled);
|
||||
return this;
|
||||
}
|
||||
|
||||
|
||||
@@ -20,8 +20,10 @@ import java.util.List;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.tool.function.FunctionToolCallback;
|
||||
import org.springframework.ai.zhipuai.api.MockWeatherService;
|
||||
import org.springframework.ai.zhipuai.api.ZhiPuAiApi;
|
||||
|
||||
@@ -38,7 +40,9 @@ public class ChatCompletionRequestTests {
|
||||
var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"),
|
||||
ZhiPuAiChatOptions.builder().model("DEFAULT_MODEL").temperature(66.6).build());
|
||||
|
||||
var request = client.createRequest(new Prompt("Test message content"), false);
|
||||
var prompt = client.buildRequestPrompt(new Prompt("Test message content"));
|
||||
|
||||
var request = client.createRequest(prompt, false);
|
||||
|
||||
assertThat(request.messages()).hasSize(1);
|
||||
assertThat(request.stream()).isFalse();
|
||||
@@ -67,17 +71,13 @@ public class ChatCompletionRequestTests {
|
||||
var request = client.createRequest(new Prompt("Test message content",
|
||||
ZhiPuAiChatOptions.builder()
|
||||
.model("PROMPT_MODEL")
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function(TOOL_FUNCTION_NAME, new MockWeatherService())
|
||||
.toolCallbacks(List.of(FunctionToolCallback.builder(TOOL_FUNCTION_NAME, new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
.build()))
|
||||
.build()),
|
||||
false);
|
||||
|
||||
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
|
||||
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
|
||||
|
||||
assertThat(request.messages()).hasSize(1);
|
||||
assertThat(request.stream()).isFalse();
|
||||
assertThat(request.model()).isEqualTo("PROMPT_MODEL");
|
||||
@@ -94,55 +94,19 @@ public class ChatCompletionRequestTests {
|
||||
var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"),
|
||||
ZhiPuAiChatOptions.builder()
|
||||
.model("DEFAULT_MODEL")
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function(TOOL_FUNCTION_NAME, new MockWeatherService())
|
||||
.toolCallbacks(List.of(FunctionToolCallback.builder(TOOL_FUNCTION_NAME, new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
.build()))
|
||||
.build());
|
||||
|
||||
var request = client.createRequest(new Prompt("Test message content"), false);
|
||||
var prompt = client.buildRequestPrompt(new Prompt("Test message content"));
|
||||
|
||||
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
|
||||
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
|
||||
assertThat(client.getFunctionCallbackRegister().get(TOOL_FUNCTION_NAME).getDescription())
|
||||
.isEqualTo("Get the weather in location");
|
||||
var request = client.createRequest(prompt, false);
|
||||
|
||||
assertThat(request.messages()).hasSize(1);
|
||||
assertThat(request.stream()).isFalse();
|
||||
assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
|
||||
|
||||
assertThat(request.tools()).as("Default Options callback functions are not automatically enabled!")
|
||||
.isNullOrEmpty();
|
||||
|
||||
// Explicitly enable the function
|
||||
request = client.createRequest(
|
||||
new Prompt("Test message content", ZhiPuAiChatOptions.builder().function(TOOL_FUNCTION_NAME).build()),
|
||||
false);
|
||||
|
||||
assertThat(request.tools()).hasSize(1);
|
||||
assertThat(request.tools().get(0).getFunction().getName()).as("Explicitly enabled function")
|
||||
.isEqualTo(TOOL_FUNCTION_NAME);
|
||||
|
||||
// Override the default options function with one from the prompt
|
||||
request = client.createRequest(new Prompt("Test message content",
|
||||
ZhiPuAiChatOptions.builder()
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function(TOOL_FUNCTION_NAME, new MockWeatherService())
|
||||
.description("Overridden function description")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
.build()))
|
||||
.build()),
|
||||
false);
|
||||
|
||||
assertThat(request.tools()).hasSize(1);
|
||||
assertThat(request.tools().get(0).getFunction().getName()).as("Explicitly enabled function")
|
||||
.isEqualTo(TOOL_FUNCTION_NAME);
|
||||
|
||||
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
|
||||
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
|
||||
assertThat(client.getFunctionCallbackRegister().get(TOOL_FUNCTION_NAME).getDescription())
|
||||
.isEqualTo("Overridden function description");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -26,6 +26,7 @@ import org.mockito.Mock;
|
||||
import org.mockito.junit.jupiter.MockitoExtension;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.document.MetadataMode;
|
||||
import org.springframework.ai.image.ImageMessage;
|
||||
@@ -88,7 +89,7 @@ public class ZhiPuAiRetryTests {
|
||||
this.retryListener = new TestRetryListener();
|
||||
this.retryTemplate.registerListener(this.retryListener);
|
||||
|
||||
this.chatModel = new ZhiPuAiChatModel(this.zhiPuAiApi, ZhiPuAiChatOptions.builder().build(), null,
|
||||
this.chatModel = new ZhiPuAiChatModel(this.zhiPuAiApi, ZhiPuAiChatOptions.builder().build(),
|
||||
this.retryTemplate);
|
||||
this.embeddingModel = new ZhiPuAiEmbeddingModel(this.zhiPuAiApi, MetadataMode.EMBED,
|
||||
ZhiPuAiEmbeddingOptions.builder().build(), this.retryTemplate);
|
||||
@@ -109,7 +110,7 @@ public class ZhiPuAiRetryTests {
|
||||
.willThrow(new TransientAiException("Transient Error 2"))
|
||||
.willReturn(ResponseEntity.of(Optional.of(expectedChatCompletion)));
|
||||
|
||||
var result = this.chatModel.call(new Prompt("text"));
|
||||
var result = this.chatModel.call(new Prompt("text", ChatOptions.builder().build()));
|
||||
|
||||
assertThat(result).isNotNull();
|
||||
assertThat(result.getResult().getOutput().getText()).isSameAs("Response");
|
||||
@@ -121,7 +122,8 @@ public class ZhiPuAiRetryTests {
|
||||
public void zhiPuAiChatNonTransientError() {
|
||||
given(this.zhiPuAiApi.chatCompletionEntity(isA(ChatCompletionRequest.class)))
|
||||
.willThrow(new RuntimeException("Non Transient Error"));
|
||||
assertThrows(RuntimeException.class, () -> this.chatModel.call(new Prompt("text")));
|
||||
assertThrows(RuntimeException.class,
|
||||
() -> this.chatModel.call(new Prompt("text", ChatOptions.builder().build())));
|
||||
}
|
||||
|
||||
@Test
|
||||
|
||||
@@ -40,6 +40,7 @@ import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
import org.springframework.ai.chat.model.StreamingChatModel;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.chat.prompt.PromptTemplate;
|
||||
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
|
||||
@@ -48,6 +49,7 @@ import org.springframework.ai.converter.BeanOutputConverter;
|
||||
import org.springframework.ai.converter.ListOutputConverter;
|
||||
import org.springframework.ai.converter.MapOutputConverter;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.ai.tool.function.FunctionToolCallback;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatOptions;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiTestConfiguration;
|
||||
import org.springframework.ai.zhipuai.api.MockWeatherService;
|
||||
@@ -86,7 +88,7 @@ class ZhiPuAiChatModelIT {
|
||||
"Tell me about 3 famous pirates from the Golden Age of Piracy and what they did.");
|
||||
SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(this.systemResource);
|
||||
Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate"));
|
||||
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
|
||||
Prompt prompt = new Prompt(List.of(userMessage, systemMessage), ChatOptions.builder().build());
|
||||
ChatResponse response = this.chatModel.call(prompt);
|
||||
assertThat(response.getResults()).hasSize(1);
|
||||
assertThat(response.getResults().get(0).getOutput().getText()).contains("Blackbeard");
|
||||
@@ -128,7 +130,7 @@ class ZhiPuAiChatModelIT {
|
||||
""";
|
||||
PromptTemplate promptTemplate = new PromptTemplate(template,
|
||||
Map.of("subject", "ice cream flavors", "format", format));
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage());
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage(), ChatOptions.builder().build());
|
||||
Generation generation = this.chatModel.call(prompt).getResult();
|
||||
|
||||
List<String> list = outputConverter.convert(generation.getOutput().getText());
|
||||
@@ -147,7 +149,7 @@ class ZhiPuAiChatModelIT {
|
||||
""";
|
||||
PromptTemplate promptTemplate = new PromptTemplate(template,
|
||||
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage());
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage(), ChatOptions.builder().build());
|
||||
Generation generation = this.chatModel.call(prompt).getResult();
|
||||
|
||||
Map<String, Object> result = outputConverter.convert(generation.getOutput().getText());
|
||||
@@ -166,7 +168,7 @@ class ZhiPuAiChatModelIT {
|
||||
{format}
|
||||
""";
|
||||
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage());
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage(), ChatOptions.builder().build());
|
||||
Generation generation = this.chatModel.call(prompt).getResult();
|
||||
|
||||
ActorsFilms actorsFilms = outputConverter.convert(generation.getOutput().getText());
|
||||
@@ -183,7 +185,7 @@ class ZhiPuAiChatModelIT {
|
||||
{format}
|
||||
""";
|
||||
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage());
|
||||
Prompt prompt = new Prompt(promptTemplate.createMessage(), ChatOptions.builder().build());
|
||||
Generation generation = this.chatModel.call(prompt).getResult();
|
||||
|
||||
ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getText());
|
||||
@@ -230,8 +232,7 @@ class ZhiPuAiChatModelIT {
|
||||
|
||||
var promptOptions = ZhiPuAiChatOptions.builder()
|
||||
.model(ZhiPuAiApi.ChatModel.GLM_4.getValue())
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function("getCurrentWeather", new MockWeatherService())
|
||||
.toolCallbacks(List.of(FunctionToolCallback.builder("getCurrentWeather", new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
.build()))
|
||||
@@ -256,8 +257,7 @@ class ZhiPuAiChatModelIT {
|
||||
|
||||
var promptOptions = ZhiPuAiChatOptions.builder()
|
||||
.model(ZhiPuAiApi.ChatModel.GLM_4.getValue())
|
||||
.functionCallbacks(List.of(FunctionCallback.builder()
|
||||
.function("getCurrentWeather", new MockWeatherService())
|
||||
.toolCallbacks(List.of(FunctionToolCallback.builder("getCurrentWeather", new MockWeatherService())
|
||||
.description("Get the weather in location")
|
||||
.inputType(MockWeatherService.Request.class)
|
||||
.build()))
|
||||
|
||||
@@ -31,6 +31,7 @@ import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.observation.DefaultChatModelObservationConvention;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.function.DefaultFunctionCallbackResolver;
|
||||
import org.springframework.ai.model.tool.ToolCallingManager;
|
||||
import org.springframework.ai.observation.conventions.AiOperationType;
|
||||
import org.springframework.ai.observation.conventions.AiProvider;
|
||||
import org.springframework.ai.zhipuai.ZhiPuAiChatModel;
|
||||
@@ -166,8 +167,7 @@ public class ZhiPuAiChatModelObservationIT {
|
||||
@Bean
|
||||
public ZhiPuAiChatModel zhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, TestObservationRegistry observationRegistry) {
|
||||
return new ZhiPuAiChatModel(zhiPuAiApi, ZhiPuAiChatOptions.builder().build(),
|
||||
new DefaultFunctionCallbackResolver(), List.of(), RetryTemplate.defaultInstance(),
|
||||
observationRegistry);
|
||||
RetryTemplate.defaultInstance(), observationRegistry);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user