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:
Ilayaperumal Gopinathan
2025-04-25 19:46:15 +01:00
parent 49e1dd10a9
commit 9692ebd7b4
12 changed files with 334 additions and 244 deletions

View File

@@ -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>

View File

@@ -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;
}
}

View File

@@ -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()

View File

@@ -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()))

View File

@@ -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()

View File

@@ -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())

View File

@@ -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;
}

View File

@@ -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;
}

View File

@@ -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");
}
}

View File

@@ -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

View File

@@ -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()))

View File

@@ -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);
}
}