Add Supplier and Consumer function callback support in function calling

Add support for no-argument Supplier and single-argument Consumer function
callbacks in the Spring AI core module. This enhancement allows:
- Registration of Supplier<O> callbacks with no input (Void) type
- Registration of Consumer<I> callbacks with no output (Void) type
- Support for Kotlin Function0 (equivalent to Java Supplier)
- Handle empty properties for Void input types in schema generation
- Enhance FunctionCallback builder to support Supplier/Consumer patterns

Additional changes:
- Add test coverage for both Supplier and Consumer callbacks in various scenarios
- Enhance TypeResolverHelper to support Consumer input type resolution
- Support lambda-style function declarations for improved ergonomics
- Add test cases for void input/output handling in OpenAI chat model
- Include examples of function calls without return values
- Add support for parameterless functions through Supplier interface

Add comprehensive documentation for the FunctionCallback API:
- Overview of the interface and its key methods
- Builder pattern usage with function and method invocation approaches
- Examples for different function types (Function, BiFunction, Supplier, Consumer)
- Best practices and common pitfalls
- Schema generation and customization options

Resolves #1718 , #1277 , #1118, #860
This commit is contained in:
Christian Tzolov
2024-11-16 19:03:42 +01:00
committed by Mark Pollack
parent f9a9c02e2b
commit 432954dad7
14 changed files with 804 additions and 187 deletions

View File

@@ -106,7 +106,8 @@ class OpenAiSpeechModelIT extends AbstractIT {
List<SpeechResponse> responses = responseFlux.collectList().block();
assertThat(responses).isNotNull();
responses.forEach(response -> {
System.out.println("Audio data chunk size: " + response.getResult().getOutput().length);
// System.out.println("Audio data chunk size: " +
// response.getResult().getOutput().length);
assertThat(response.getResult().getOutput()).isNotEmpty();
});
}

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.openai.chat;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.function.BiFunction;
import java.util.stream.Collectors;
@@ -28,6 +29,7 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import reactor.core.publisher.Flux;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
@@ -59,6 +61,25 @@ class OpenAiChatModelFunctionCallingIT {
@Autowired
ChatModel chatModel;
@Test
void functionCallSupplier() {
Map<String, Object> state = new ConcurrentHashMap<>();
// @formatter:off
String response = ChatClient.create(this.chatModel).prompt()
.user("Turn the light on in the living room")
.functions(FunctionCallback.builder()
.function("turnsLightOnInTheLivingRoom", () -> state.put("Light", "ON"))
.build())
.call()
.content();
// @formatter:on
logger.info("Response: {}", response);
assertThat(state).containsEntry("Light", "ON");
}
@Test
void functionCallTest() {
functionCallTest(OpenAiChatOptions.builder()

View File

@@ -340,7 +340,7 @@ public abstract class ModelOptionsUtils {
* @return the generated JSON Schema as a String.
* @deprecated use {@link #getJsonSchema(Type, boolean)} instead.
*/
@Deprecated
@Deprecated(since = "1.0 M4")
public static String getJsonSchema(Class<?> clazz, boolean toUpperCaseTypeValues) {
if (SCHEMA_GENERATOR_CACHE.get() == null) {
@@ -395,6 +395,11 @@ public abstract class ModelOptionsUtils {
}
ObjectNode node = SCHEMA_GENERATOR_CACHE.get().generateSchema(inputType);
if ((inputType == Void.class) && !node.has("properties")) {
node.putObject("properties");
}
if (toUpperCaseTypeValues) { // Required for OpenAPI 3.0 (at least Vertex AI
// version of it).
toUpperCaseTypeValues(node);

View File

@@ -1,24 +1,27 @@
/*
* Copyright 2024 - 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
* Copyright 2023-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.model.function;
import java.lang.reflect.Type;
import java.util.Arrays;
import java.util.function.BiFunction;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.DeserializationFeature;
@@ -43,7 +46,7 @@ import org.springframework.util.StringUtils;
/**
* Default implementation of the {@link FunctionCallback.Builder}.
*
*
* @author Christian Tzolov
* @since 1.0.0
*/
@@ -137,6 +140,20 @@ public class DefaultFunctionCallbackBuilder implements FunctionCallback.Builder
return new DefaultFunctionInvokingSpec<>(name, biFunction);
}
@Override
public <O> FunctionInvokingSpec<Void, O> function(String name, Supplier<O> supplier) {
Function<Void, O> function = (input) -> supplier.get();
return new DefaultFunctionInvokingSpec<>(name, function).inputType(Void.class);
}
public <I> FunctionInvokingSpec<I, Void> function(String name, Consumer<I> consumer) {
Function<I, Void> function = (I input) -> {
consumer.accept(input);
return null;
};
return new DefaultFunctionInvokingSpec<>(name, function);
}
@Override
public MethodInvokingSpec method(String methodName, Class<?>... argumentTypes) {
return new DefaultMethodInvokingSpec(methodName, argumentTypes);

View File

@@ -17,7 +17,9 @@
package org.springframework.ai.model.function;
import java.util.function.BiFunction;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
import com.fasterxml.jackson.databind.ObjectMapper;
@@ -141,6 +143,16 @@ public interface FunctionCallback {
*/
<I, O> FunctionInvokingSpec<I, O> function(String name, BiFunction<I, ToolContext, O> biFunction);
/**
* Builds a {@link Supplier} invoking {@link FunctionCallback} instance.
*/
<O> FunctionInvokingSpec<Void, O> function(String name, Supplier<O> supplier);
/**
* Builds a {@link Consumer} invoking {@link FunctionCallback} instance.
*/
<I> FunctionInvokingSpec<I, Void> function(String name, Consumer<I> consumer);
/**
* Builds a Method invoking {@link FunctionCallback} instance.
*/
@@ -189,14 +201,14 @@ public interface FunctionCallback {
MethodInvokingSpec name(String name);
/**
* For non static objects the target object is used to invoke the method.
* For non-static objects the target object is used to invoke the method.
* @param methodObject target object where the method is defined.
*/
MethodInvokingSpec targetObject(Object methodObject);
/**
* Target class where the method is defined. Used for static methods. For non
* static methods the target object is used.
* Target class where the method is defined. Used for static methods. For
* non-static methods the target object is used.
* @param targetClass method target class.
*/
MethodInvokingSpec targetClass(Class<?> targetClass);

View File

@@ -17,9 +17,12 @@
package org.springframework.ai.model.function;
import java.util.function.BiFunction;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
import com.fasterxml.jackson.annotation.JsonClassDescription;
import kotlin.jvm.functions.Function0;
import kotlin.jvm.functions.Function1;
import kotlin.jvm.functions.Function2;
@@ -30,6 +33,7 @@ import org.springframework.context.ApplicationContextAware;
import org.springframework.context.annotation.Description;
import org.springframework.context.support.GenericApplicationContext;
import org.springframework.core.KotlinDetector;
import org.springframework.core.ParameterizedTypeReference;
import org.springframework.core.ResolvableType;
import org.springframework.lang.NonNull;
import org.springframework.lang.Nullable;
@@ -38,9 +42,9 @@ import org.springframework.util.StringUtils;
/**
* A Spring {@link ApplicationContextAware} implementation that provides a way to retrieve
* a {@link Function} from the Spring context and wrap it into a {@link FunctionCallback}.
*
* <p>
* The name of the function is determined by the bean name.
*
* <p>
* The description of the function is determined by the following rules:
* <ul>
* <li>Provided as a default description</li>
@@ -69,24 +73,28 @@ public class FunctionCallbackContext implements ApplicationContextAware {
@SuppressWarnings({ "unchecked" })
public FunctionCallback getFunctionCallback(@NonNull String beanName, @Nullable String defaultDescription) {
ResolvableType functionType = TypeResolverHelper.resolveBeanType(this.applicationContext, beanName);
ResolvableType functionInputType = TypeResolverHelper.getFunctionArgumentType(functionType, 0);
ResolvableType functionInputType = (ResolvableType.forType(Supplier.class).isAssignableFrom(functionType))
? ResolvableType.forType(Void.class) : TypeResolverHelper.getFunctionArgumentType(functionType, 0);
Class<?> functionInputClass = functionInputType.toClass();
String functionDescription = resolveFunctionDescription(beanName, defaultDescription,
functionInputType.toClass());
Object bean = this.applicationContext.getBean(beanName);
return buildFunctionCallback(beanName, functionType, functionInputType, functionDescription, bean);
}
private String resolveFunctionDescription(String beanName, String defaultDescription, Class<?> functionInputClass) {
String functionDescription = defaultDescription;
if (!StringUtils.hasText(functionDescription)) {
// Look for a Description annotation on the bean
Description descriptionAnnotation = this.applicationContext.findAnnotationOnBean(beanName,
Description.class);
if (descriptionAnnotation != null) {
functionDescription = descriptionAnnotation.value();
}
if (!StringUtils.hasText(functionDescription)) {
// Look for a JsonClassDescription annotation on the input class
JsonClassDescription jsonClassDescriptionAnnotation = functionInputClass
.getAnnotation(JsonClassDescription.class);
if (jsonClassDescriptionAnnotation != null) {
@@ -95,13 +103,17 @@ public class FunctionCallbackContext implements ApplicationContextAware {
}
if (!StringUtils.hasText(functionDescription)) {
throw new IllegalStateException("Could not determine function description."
throw new IllegalStateException("Could not determine function description. "
+ "Please provide a description either as a default parameter, via @Description annotation on the bean "
+ "or @JsonClassDescription annotation on the input class.");
}
}
Object bean = this.applicationContext.getBean(beanName);
return functionDescription;
}
private FunctionCallback buildFunctionCallback(String beanName, ResolvableType functionType,
ResolvableType functionInputType, String functionDescription, Object bean) {
if (KotlinDetector.isKotlinPresent()) {
if (KotlinDelegate.isKotlinFunction(functionType.toClass())) {
@@ -109,37 +121,61 @@ public class FunctionCallbackContext implements ApplicationContextAware {
.schemaType(this.schemaType)
.description(functionDescription)
.function(beanName, KotlinDelegate.wrapKotlinFunction(bean))
.inputType(functionInputClass)
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
else if (KotlinDelegate.isKotlinBiFunction(functionType.toClass())) {
if (KotlinDelegate.isKotlinBiFunction(functionType.toClass())) {
return FunctionCallback.builder()
.description(functionDescription)
.schemaType(this.schemaType)
.function(beanName, KotlinDelegate.wrapKotlinBiFunction(bean))
.inputType(functionInputClass)
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
if (KotlinDelegate.isKotlinSupplier(functionType.toClass())) {
return FunctionCallback.builder()
.description(functionDescription)
.schemaType(this.schemaType)
.function(beanName, KotlinDelegate.wrapKotlinSupplier(bean))
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
}
if (bean instanceof Function<?, ?> function) {
return FunctionCallback.builder()
.schemaType(this.schemaType)
.description(functionDescription)
.function(beanName, function)
.inputType(functionInputClass)
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
else if (bean instanceof BiFunction<?, ?, ?>) {
if (bean instanceof BiFunction<?, ?, ?>) {
return FunctionCallback.builder()
.description(functionDescription)
.schemaType(this.schemaType)
.function(beanName, (BiFunction<?, ToolContext, ?>) bean)
.inputType(functionInputClass)
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
else {
throw new IllegalStateException();
if (bean instanceof Supplier<?> supplier) {
return FunctionCallback.builder()
.description(functionDescription)
.schemaType(this.schemaType)
.function(beanName, supplier)
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
if (bean instanceof Consumer<?> consumer) {
return FunctionCallback.builder()
.description(functionDescription)
.schemaType(this.schemaType)
.function(beanName, consumer)
.inputType(ParameterizedTypeReference.forType(functionInputType.getType()))
.build();
}
throw new IllegalStateException("Unsupported function type");
}
public enum SchemaType {
@@ -148,7 +184,16 @@ public class FunctionCallbackContext implements ApplicationContextAware {
}
private static class KotlinDelegate {
private static final class KotlinDelegate {
public static boolean isKotlinSupplier(Class<?> clazz) {
return Function0.class.isAssignableFrom(clazz);
}
@SuppressWarnings("unchecked")
public static Supplier<?> wrapKotlinSupplier(Object function) {
return () -> ((Function0<Object>) function).invoke();
}
public static boolean isKotlinFunction(Class<?> clazz) {
return Function1.class.isAssignableFrom(clazz);

View File

@@ -20,8 +20,11 @@ import java.lang.reflect.Method;
import java.lang.reflect.Modifier;
import java.util.Arrays;
import java.util.function.BiFunction;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
import kotlin.jvm.functions.Function0;
import kotlin.jvm.functions.Function1;
import kotlin.jvm.functions.Function2;
@@ -44,6 +47,16 @@ import org.springframework.util.ReflectionUtils;
*/
public abstract class TypeResolverHelper {
/**
* Returns the input class of a given Consumer class.
* @param consumerClass The consumer class.
* @return The input class of the consumer.
*/
public static Class<?> getConsumerInputClass(Class<? extends Consumer<?>> consumerClass) {
ResolvableType resolvableType = ResolvableType.forClass(consumerClass).as(Consumer.class);
return (resolvableType == ResolvableType.NONE ? Object.class : resolvableType.getGeneric(0).toClass());
}
/**
* Returns the input class of a given function class.
* @param biFunctionClass The function class.
@@ -108,51 +121,79 @@ public abstract class TypeResolverHelper {
* resolvable.
*/
public static ResolvableType resolveBeanType(GenericApplicationContext applicationContext, String beanName) {
BeanDefinition beanDefinition;
BeanDefinition beanDefinition = getBeanDefinition(applicationContext, beanName);
// Try to resolve directly
ResolvableType functionType = beanDefinition.getResolvableType();
if (functionType.resolve() != null) {
return functionType;
}
// Handle root bean definitions with factory methods
if (beanDefinition instanceof RootBeanDefinition rootBeanDefinition) {
return resolveRootBeanDefinitionType(applicationContext, rootBeanDefinition);
}
// Handle @Component beans
return resolveComponentBeanType(applicationContext, beanDefinition, beanName);
}
private static BeanDefinition getBeanDefinition(GenericApplicationContext applicationContext, String beanName) {
try {
beanDefinition = applicationContext.getBeanDefinition(beanName);
return applicationContext.getBeanDefinition(beanName);
}
catch (NoSuchBeanDefinitionException ex) {
throw new IllegalArgumentException(
"Functional bean with name " + beanName + " does not exist in the context.");
}
ResolvableType functionType = beanDefinition.getResolvableType();
Class<?> resolvableClass = functionType.resolve();
if (resolvableClass != null) {
return functionType;
}
if (beanDefinition instanceof RootBeanDefinition rootBeanDefinition) {
Class<?> factoryClass;
boolean isStatic;
if (rootBeanDefinition.getFactoryBeanName() != null) {
factoryClass = applicationContext.getBeanFactory().getType(rootBeanDefinition.getFactoryBeanName());
isStatic = false;
}
else {
factoryClass = rootBeanDefinition.getBeanClass();
isStatic = true;
}
Assert.state(factoryClass != null, "Unresolvable factory class");
factoryClass = ClassUtils.getUserClass(factoryClass);
}
Method[] candidates = getCandidateMethods(factoryClass, rootBeanDefinition);
Method uniqueCandidate = null;
for (Method candidate : candidates) {
if ((!isStatic || isStaticCandidate(candidate, factoryClass))
&& rootBeanDefinition.isFactoryMethod(candidate)) {
if (uniqueCandidate == null) {
uniqueCandidate = candidate;
}
else if (isParamMismatch(uniqueCandidate, candidate)) {
uniqueCandidate = null;
break;
}
private static ResolvableType resolveRootBeanDefinitionType(GenericApplicationContext applicationContext,
RootBeanDefinition rootBeanDefinition) {
Class<?> factoryClass;
boolean isStatic;
if (rootBeanDefinition.getFactoryBeanName() != null) {
factoryClass = applicationContext.getBeanFactory().getType(rootBeanDefinition.getFactoryBeanName());
isStatic = false;
}
else {
factoryClass = rootBeanDefinition.getBeanClass();
isStatic = true;
}
Assert.state(factoryClass != null, "Unresolvable factory class");
factoryClass = ClassUtils.getUserClass(factoryClass);
Method uniqueCandidate = findUniqueFactoryMethod(factoryClass, isStatic, rootBeanDefinition);
rootBeanDefinition.setResolvedFactoryMethod(uniqueCandidate);
return rootBeanDefinition.getResolvableType();
}
private static Method findUniqueFactoryMethod(Class<?> factoryClass, boolean isStatic,
RootBeanDefinition rootBeanDefinition) {
Method[] candidates = getCandidateMethods(factoryClass, rootBeanDefinition);
Method uniqueCandidate = null;
for (Method candidate : candidates) {
if ((!isStatic || isStaticCandidate(candidate, factoryClass))
&& rootBeanDefinition.isFactoryMethod(candidate)) {
if (uniqueCandidate == null) {
uniqueCandidate = candidate;
}
else if (isParamMismatch(uniqueCandidate, candidate)) {
uniqueCandidate = null;
break;
}
}
rootBeanDefinition.setResolvedFactoryMethod(uniqueCandidate);
return rootBeanDefinition.getResolvableType();
}
// Support for @Component
return uniqueCandidate;
}
private static ResolvableType resolveComponentBeanType(GenericApplicationContext applicationContext,
BeanDefinition beanDefinition, String beanName) {
if (beanDefinition.getFactoryMethodName() == null && beanDefinition.getBeanClassName() != null) {
try {
return ResolvableType.forClass(
@@ -199,6 +240,12 @@ public abstract class TypeResolverHelper {
else if (BiFunction.class.isAssignableFrom(resolvableClass)) {
functionArgumentResolvableType = functionType.as(BiFunction.class);
}
else if (Supplier.class.isAssignableFrom(resolvableClass)) {
functionArgumentResolvableType = functionType.as(Supplier.class);
}
else if (Consumer.class.isAssignableFrom(resolvableClass)) {
functionArgumentResolvableType = functionType.as(Consumer.class);
}
else if (KotlinDetector.isKotlinPresent()) {
if (KotlinDelegate.isKotlinFunction(resolvableClass)) {
functionArgumentResolvableType = KotlinDelegate.adaptToKotlinFunctionType(functionType);
@@ -206,6 +253,9 @@ public abstract class TypeResolverHelper {
else if (KotlinDelegate.isKotlinBiFunction(resolvableClass)) {
functionArgumentResolvableType = KotlinDelegate.adaptToKotlinBiFunctionType(functionType);
}
else if (KotlinDelegate.isKotlinSupplier(resolvableClass)) {
functionArgumentResolvableType = KotlinDelegate.adaptToKotlinSupplierType(functionType);
}
}
if (functionArgumentResolvableType == ResolvableType.NONE) {
@@ -216,7 +266,15 @@ public abstract class TypeResolverHelper {
return functionArgumentResolvableType.getGeneric(argumentIndex);
}
private static class KotlinDelegate {
private static final class KotlinDelegate {
public static boolean isKotlinSupplier(Class<?> clazz) {
return Function0.class.isAssignableFrom(clazz);
}
public static ResolvableType adaptToKotlinSupplierType(ResolvableType resolvableType) {
return resolvableType.as(Function0.class);
}
public static boolean isKotlinFunction(Class<?> clazz) {
return Function1.class.isAssignableFrom(clazz);

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.model.function;
import java.util.function.Consumer;
import java.util.function.Function;
import org.junit.jupiter.params.ParameterizedTest;
@@ -39,7 +40,7 @@ public class TypeResolverHelperIT {
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { "weatherClassDefinition", "weatherFunctionDefinition", "standaloneWeatherFunction",
"scannedStandaloneWeatherFunction", "componentWeatherFunction" })
"scannedStandaloneWeatherFunction", "componentWeatherFunction", "weatherConsumer" })
void beanInputTypeResolutionWithResolvableType(String beanName) {
assertThat(this.applicationContext).isNotNull();
ResolvableType functionType = TypeResolverHelper.resolveBeanType(this.applicationContext, beanName);
@@ -89,6 +90,13 @@ public class TypeResolverHelperIT {
return new StandaloneWeatherFunction();
}
@Bean
Consumer<WeatherRequest> weatherConsumer() {
return (weatherRequest) -> {
System.out.println(weatherRequest);
};
}
}
}

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.model.function;
import java.util.function.Consumer;
import java.util.function.Function;
import com.fasterxml.jackson.annotation.JsonClassDescription;
@@ -35,6 +36,12 @@ import static org.assertj.core.api.Assertions.assertThat;
*/
public class TypeResolverHelperTests {
@Test
public void testGetConsumerInputType() {
Class<?> inputType = TypeResolverHelper.getConsumerInputClass(MyConsumer.class);
assertThat(inputType).isEqualTo(Request.class);
}
@Test
public void testGetFunctionInputType() {
Class<?> inputType = TypeResolverHelper.getFunctionInputClass(MockWeatherService.class);
@@ -63,6 +70,14 @@ public class TypeResolverHelperTests {
}
public static class MyConsumer implements Consumer<Request> {
@Override
public void accept(Request request) {
}
}
public static class MockWeatherService implements Function<Request, Response> {
@Override

View File

@@ -97,6 +97,7 @@
* xref:api/prompt.adoc[]
* xref:api/structured-output-converter.adoc[Structured Output]
* xref:api/functions.adoc[Function Calling]
** xref:api/function-callback.adoc[FunctionCallback API]
* xref:api/multimodality.adoc[Multimodality]
* xref:api/etl-pipeline.adoc[]
* xref:api/testing.adoc[AI Model Evaluation]

View File

@@ -0,0 +1,240 @@
= FunctionCallback
== Overview
The `FunctionCallback` interface in Spring AI provides a standardized way to implement Large Language Model (LLM) function calling capabilities. It allows developers to register custom functions that can be called by AI models when specific conditions or intents are detected in the prompts.
== FunctionCallback Interface
The main interface defines several key methods:
* `getName()`: Returns the unique function name within the AI model context
* `getDescription()`: Provides a description that helps the model decide when to invoke the function
* `getInputTypeSchema()`: Defines the JSON schema for the function's input parameters
* `call(String functionInput)`: Handles the actual function execution
* `call(String functionInput, ToolContext toolContext)`: Extended version that supports additional context
== Builder Pattern
Spring AI provides a fluent builder API for creating `FunctionCallback` implementations.
This is particularly useful for defining function callbacks that you can register, pragmatically, on the fly, with your `ChatClient` or `ChatModel` model calls.
The builders helps with complex configurations, such as custom response handling, schema types (e.g. JSONSchema or OpenAPI), and object mapping.
=== Function-Invoking Approach
Converts any `java.util.function.Function`, `BiFunction`, `Supplier` or `Consumer` into a `FunctionCallback` that can be called by the AI model.
NOTE: You can use lambda expressions or method references to define the function logic but you must provide the input type of the function using the `inputType(TYPE)`.
==== Function<I, O>
[source,java]
----
FunctionCallback callback = FunctionCallback.builder()
.description("Process a new order")
.function("processOrder", (Order order) -> processOrderLogic(order))
.inputType(Order.class)
.build();
----
==== BiFunction<I, ToolContext, O> with ToolContext
[source,java]
----
FunctionCallback callback = FunctionCallback.builder()
.description("Process a new order with context")
.function("processOrder", (Order order, ToolContext context) ->
processOrderWithContext(order, context))
.inputType(Order.class)
.build();
----
==== Supplier<O>
Use `java.util.Supplier<O>` or `java.util.function.Function<Void, O>` to define functions that don't take any input:
[source,java]
----
FunctionCallback.builder()
.description("Turns light onn in the living room")
.function("turnsLight", () -> state.put("Light", "ON"))
.inputType(Void.class)
.build();
----
==== Consumer<I>
Use `java.util.Consumer<I>` or `java.util.function.Function<I, Void>` to define functions that don't produce output:
[source,java]
----
record LightInfo(String roomName, boolean isOn) {}
FunctionCallback.builder()
.description("Turns light on/off in a selected room")
.function("turnsLight", (LightInfo lightInfo) -> {
logger.info("Turning light to [" + lightInfo.isOn + "] in " + lightInfo.roomName());
})
.inputType(LightInfo.class)
.build();
----
==== Generics Input Type
Use the `ParameterizedTypeReference` to define functions with generic input types:
[source,java]
----
record TrainSearchRequest<T>(T data) {}
record TrainSearchSchedule(String from, String to, String date) {}
record TrainSearchScheduleResponse(String from, String to, String date, String trainNumber) {}
FunctionCallback.builder()
.description("Schedule a train reservation")
.function("trainSchedule", (TrainSearchRequest<TrainSearchSchedule> request) -> {
logger.info("Schedule: " + request.data().from() + " to " + request.data().to());
return new TrainSearchScheduleResponse(request.data().from(), request. data().to(), "", "123");
})
.inputType(new ParameterizedTypeReference<TrainSearchRequest<TrainSearchSchedule>>() {})
.build();
----
=== Method Invoking Approach
Enables method invocation through reflection while automatically handling JSON schema generation and parameter conversion. Its particularly useful for integrating Java methods as callable functions within AI model interactions.
The method invoking implements the `FunctionCallback` interface and provides:
- Automatic JSON schema generation for method parameters
- Support for both static and instance methods
- Any number of parameters (including none) and return values (including void)
- Any parameter/return types (primitives, objects, collections)
- Special handling for `ToolContext` parameters
==== Static Method Invocation
You can refer to a static method in a class by providing the method name, parameter types, and the target class.
[source,java]
----
public class WeatherService {
public static String getWeather(String city, TemperatureUnit unit) {
return "Temperature in " + city + ": 20" + unit;
}
}
FunctionCallback callback = FunctionCallback.builder()
.description("Get weather information for a city")
.method("getWeather", String.class, TemperatureUnit.class)
.targetClass(WeatherService.class)
.build();
----
==== Object instance Method Invocation
You can refer to an instance method in a class by providing the method name, parameter types, and the target object instance.
[source,java]
----
public class DeviceController {
public void setDeviceState(String deviceId, boolean state, ToolContext context) {
Map<String, Object> contextData = context.getContext();
// Implementation using context data
}
}
DeviceController controller = new DeviceController();
String response = ChatClient.create(chatModel).prompt()
.user("Turn on the living room lights")
.functions(FunctionCallback.builder()
.description("Control device state")
.method("setDeviceState", String.class,boolean.class,ToolContext.class)
.targetObject(controller)
.build())
.toolContext(Map.of("location", "home"))
.call()
.content();
----
TIP: Optionally, using the `.name()`, you can set a custom function name different from the method name.
== Schema Type Support
The framework supports different schema types for function parameter validation:
* JSON Schema (default)
* OpenAPI Schema (for Vertex AI compatibility)
[source,java]
----
FunctionCallback.builder()
.schemaType(SchemaType.OPEN_API_SCHEMA)
// ... other configuration
.build();
----
=== Custom Response Handling
[source,java]
----
FunctionCallback.builder()
.responseConverter(response ->
customResponseFormatter.format(response))
// ... other configuration
.build();
----
=== Custom Object Mapping
[source,java]
----
FunctionCallback.builder()
.objectMapper(customObjectMapper)
// ... other configuration
.build();
----
== Best Practices
=== Descriptive Names and Descriptions
* Provide unique function names
* Write comprehensive descriptions to help the model understand when to invoke the function
=== Input Type & Schema
* For the function invoking approach, define input types explicitly and use `ParameterizedTypeReference` for generic types.
* Consider using custom schema when auto-generated ones don't meet requirements.
=== Error Handling
* Implement proper error handling in function implementations and return the error message in the response
* You can use the ToolContext to provide additional error context when needed
=== Tool Context Usage
* Use ToolContext when additional state or context is required that is provided from the User and not part of the function input generated by the AI model.
* Use `BiFunction<I, ToolContext, O>` to access the ToolContext in the function invocation approach and add `ToolContext` parameter in the method invoking approach.
== Notes on Schema Generation
* The framework automatically generates JSON schemas from Java types
* For function invoking, the schema is generated based on the input type for the function that needs to be set using `inputType(TYPE)`. Use `ParameterizedTypeReference` for generic types.
* Generated schemas respect Jackson annotations on model classes
* You can bypass the automatic generation by providing custom schemas using `inputTypeSchema()`
== Common Pitfalls to Avoid
=== Lack of Description
* Always provide explicit descriptions instead of relying on auto-generated ones
* Clear descriptions improve model's function selection accuracy
=== Schema Mismatches
* Ensure input types match the Function's input parameter types.
* Use `ParameterizedTypeReference` for generic types.

View File

@@ -66,7 +66,8 @@ public class PaymentStatusPromptIT {
var promptOptions = MistralAiChatOptions.builder()
.withFunctionCallbacks(List.of(FunctionCallback.builder()
.description("Get payment status of a transaction")
.function("retrievePaymentStatus", transaction -> new Status(DATA.get(transaction).status()))
.function("retrievePaymentStatus",
(Transaction transaction) -> new Status(DATA.get(transaction).status()))
.inputType(Transaction.class)
.build()))
.build();

View File

@@ -16,6 +16,8 @@
package org.springframework.ai.autoconfigure.openai.tool;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.stream.Collectors;
import org.junit.jupiter.api.Test;
@@ -72,6 +74,36 @@ public class FunctionCallbackInPrompt2IT {
});
}
@Test
void lambdaFunctionCallTest() {
Map<String, Object> state = new ConcurrentHashMap<>();
record LightInfo(String roomName, boolean isOn) {
}
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// @formatter:off
String content = ChatClient.builder(chatModel).build().prompt()
.user("Turn the light on in the kitchen and in the living room!")
.functions(FunctionCallback.builder()
.description("Turn light on or off in a room")
.function("turnLight", (LightInfo lightInfo) -> {
logger.info("Turning light to [" + lightInfo.isOn + "] in " + lightInfo.roomName());
state.put(lightInfo.roomName(), lightInfo.isOn());
})
.inputType(LightInfo.class)
.build())
.call().content();
// @formatter:on
logger.info("Response: {}", content);
assertThat(state).containsEntry("kitchen", Boolean.TRUE);
assertThat(state).containsEntry("living room", Boolean.TRUE);
});
}
@Test
void functionCallTest2() {
this.contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName())

View File

@@ -18,10 +18,14 @@ package org.springframework.ai.autoconfigure.openai.tool;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.function.BiFunction;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
import java.util.stream.Collectors;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.slf4j.Logger;
@@ -52,179 +56,272 @@ import static org.assertj.core.api.Assertions.assertThat;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
class FunctionCallbackWithPlainFunctionBeanIT {
private final Logger logger = LoggerFactory.getLogger(FunctionCallbackWithPlainFunctionBeanIT.class);
private static final Logger logger = LoggerFactory.getLogger(FunctionCallbackWithPlainFunctionBeanIT.class);
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"),
"spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName())
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withUserConfiguration(Config.class);
private static Map<String, Object> feedback = new ConcurrentHashMap<>();
@BeforeEach
void setUp() {
feedback.clear();
}
@Test
void functionCallingVoidInput() {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage("Turn the light on in the living room");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("turnLivingRoomLightOn").build()));
logger.info("Response: {}", response);
assertThat(feedback).hasSize(1);
assertThat(feedback.get("turnLivingRoomLightOn")).isEqualTo(Boolean.valueOf(true));
});
}
@Test
void functionCallingSupplier() {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage("Turn the light on in the living room");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("turnLivingRoomLightOnSupplier").build()));
logger.info("Response: {}", response);
assertThat(feedback).hasSize(1);
assertThat(feedback.get("turnLivingRoomLightOnSupplier")).isEqualTo(Boolean.valueOf(true));
});
}
@Test
void functionCallingVoidOutput() {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage("Turn the light on in the kitchen and in the living room");
ChatResponse response = chatModel
.call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("turnLight").build()));
logger.info("Response: {}", response);
assertThat(feedback).hasSize(2);
assertThat(feedback.get("kitchen")).isEqualTo(Boolean.valueOf(true));
assertThat(feedback.get("living room")).isEqualTo(Boolean.valueOf(true));
});
}
@Test
void functionCallingConsumer() {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage("Turn the light on in the kitchen and in the living room");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("turnLightConsumer").build()));
logger.info("Response: {}", response);
assertThat(feedback).hasSize(2);
assertThat(feedback.get("kitchen")).isEqualTo(Boolean.valueOf(true));
assertThat(feedback.get("living room")).isEqualTo(Boolean.valueOf(true));
});
}
@Test
void trainScheduler() {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"Please schedule a train from San Francisco to Los Angeles on 2023-12-25");
PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder()
.withFunction("trainReservation")
.build();
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions));
logger.info("Response: {}", response.getResult().getOutput().getContent());
});
}
@Test
void functionCallWithDirectBiFunction() {
this.contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName())
.run(context -> {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
ChatClient chatClient = ChatClient.builder(chatModel).build();
ChatClient chatClient = ChatClient.builder(chatModel).build();
String content = chatClient.prompt("What's the weather like in San Francisco, Tokyo, and Paris?")
.functions("weatherFunctionWithContext")
.toolContext(Map.of("sessionId", "123"))
.call()
.content();
logger.info(content);
String content = chatClient.prompt("What's the weather like in San Francisco, Tokyo, and Paris?")
.functions("weatherFunctionWithContext")
.toolContext(Map.of("sessionId", "123"))
.call()
.content();
logger.info(content);
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder()
.withFunction("weatherFunctionWithContext")
.withToolContext(Map.of("sessionId", "123"))
.build()));
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder()
.withFunction("weatherFunctionWithContext")
.withToolContext(Map.of("sessionId", "123"))
.build()));
logger.info("Response: {}", response);
logger.info("Response: {}", response);
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
});
});
}
@Test
void functionCallWithBiFunctionClass() {
this.contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName())
.run(context -> {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
ChatClient chatClient = ChatClient.builder(chatModel).build();
ChatClient chatClient = ChatClient.builder(chatModel).build();
String content = chatClient.prompt("What's the weather like in San Francisco, Tokyo, and Paris?")
.functions("weatherFunctionWithClassBiFunction")
.toolContext(Map.of("sessionId", "123"))
.call()
.content();
logger.info(content);
String content = chatClient.prompt("What's the weather like in San Francisco, Tokyo, and Paris?")
.functions("weatherFunctionWithClassBiFunction")
.toolContext(Map.of("sessionId", "123"))
.call()
.content();
logger.info(content);
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder()
.withFunction("weatherFunctionWithClassBiFunction")
.withToolContext(Map.of("sessionId", "123"))
.build()));
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder()
.withFunction("weatherFunctionWithClassBiFunction")
.withToolContext(Map.of("sessionId", "123"))
.build()));
logger.info("Response: {}", response);
logger.info("Response: {}", response);
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
});
});
}
@Test
void functionCallTest() {
this.contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName())
.run(context -> {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunction").build()));
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunction").build()));
logger.info("Response: {}", response);
logger.info("Response: {}", response);
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
// Test weatherFunctionTwo
response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunctionTwo").build()));
// Test weatherFunctionTwo
response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunctionTwo").build()));
logger.info("Response: {}", response);
logger.info("Response: {}", response);
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
});
});
}
@Test
void functionCallWithPortableFunctionCallingOptions() {
this.contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName(),
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris?");
// Test weatherFunction
UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?");
PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder()
.withFunction("weatherFunction")
.build();
PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder()
.withFunction("weatherFunction")
.build();
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions));
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions));
logger.info("Response: {}", response.getResult().getOutput().getContent());
logger.info("Response: {}", response.getResult().getOutput().getContent());
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
});
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
});
}
@Test
void streamFunctionCallTest() {
this.contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=" + ChatModel.GPT_4_O_MINI.getName(),
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
this.contextRunner.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
// Test weatherFunction
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'");
Flux<ChatResponse> response = chatModel.stream(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunction").build()));
Flux<ChatResponse> response = chatModel.stream(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunction").build()));
String content = response.collectList()
.block()
.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
logger.info("Response: {}", content);
String content = response.collectList()
.block()
.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
logger.info("Response: {}", content);
assertThat(content).contains("30", "10", "15");
assertThat(content).contains("30", "10", "15");
// Test weatherFunctionTwo
response = chatModel.stream(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunctionTwo").build()));
// Test weatherFunctionTwo
response = chatModel.stream(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("weatherFunctionTwo").build()));
content = response.collectList()
.block()
.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
logger.info("Response: {}", content);
content = response.collectList()
.block()
.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
logger.info("Response: {}", content);
assertThat(content).isNotEmpty().withFailMessage("Content returned from OpenAI model is empty");
assertThat(content).contains("30", "10", "15");
assertThat(content).isNotEmpty().withFailMessage("Content returned from OpenAI model is empty");
assertThat(content).contains("30", "10", "15");
});
});
}
@Configuration
@@ -256,6 +353,70 @@ class FunctionCallbackWithPlainFunctionBeanIT {
return (weatherService::apply);
}
record LightInfo(String roomName, boolean isOn) {
}
@Bean
@Description("Turn light on or off in a room")
public Function<LightInfo, Void> turnLight() {
return (LightInfo lightInfo) -> {
logger.info("Turning light to [" + lightInfo.isOn + "] in " + lightInfo.roomName());
feedback.put(lightInfo.roomName(), lightInfo.isOn());
return null;
};
}
@Bean
@Description("Turn light on or off in a room")
public Consumer<LightInfo> turnLightConsumer() {
return (LightInfo lightInfo) -> {
logger.info("Turning light to [" + lightInfo.isOn + "] in " + lightInfo.roomName());
feedback.put(lightInfo.roomName(), lightInfo.isOn());
};
}
@Bean
@Description("Turns light on in the living room")
public Function<Void, String> turnLivingRoomLightOn() {
return (Void v) -> {
logger.info("Turning light on in the living room");
feedback.put("turnLivingRoomLightOn", Boolean.TRUE);
return "Done";
};
}
@Bean
@Description("Turns light on in the living room")
public Supplier<String> turnLivingRoomLightOnSupplier() {
return () -> {
logger.info("Turning light on in the living room");
feedback.put("turnLivingRoomLightOnSupplier", Boolean.TRUE);
return "Done";
};
}
record TrainSearchSchedule(String from, String to, String date) {
}
record TrainSearchScheduleResponse(String from, String to, String date, String trainNumber) {
}
record TrainSearchRequest<T>(T data) {
}
record TrainSearchResponse<T>(T data) {
}
@Bean
@Description("Schedule a train reservation")
public Function<TrainSearchRequest<TrainSearchSchedule>, TrainSearchResponse<TrainSearchScheduleResponse>> trainReservation() {
return (TrainSearchRequest<TrainSearchSchedule> request) -> {
logger.info("Turning light to [" + request.data().from() + "] in " + request.data().to());
return new TrainSearchResponse<>(
new TrainSearchScheduleResponse(request.data().from(), request.data().to(), "", "123"));
};
}
}
public static class MyBiFunction