Spring AI Function Calling实现原理与自定义函数

Spring AI Function Calling实现原理与自定义函数

前置知识

  • 理解Java反射机制
  • 了解LLM的Function Calling协议
  • Spring AI ChatClient基础用法
  • JSON Schema规范

核心概念

Function Calling是大模型调用外部工具的标准协议,允许模型通过结构化输出决定何时、如何调用预定义函数。Spring AI通过注解和反射自动完成函数注册、参数序列化和结果回传。

工作原理

复制代码
用户输入 → LLM判断是否需要调用函数 → 生成函数调用请求 → 
Spring AI执行Java方法 → 结果回传LLM → 生成最终回复

注:

完整实现

1. 基础函数定义 - @Function注解

java 复制代码
@Component
public class OrderFunctions {

    @Autowired
    private OrderService orderService;

    /**
     * 查询订单状态
     */
    @Bean
    @Description("根据订单号查询订单的当前状态和物流信息")
    public Function<OrderRequest, OrderResponse> queryOrderStatus(
            ApplicationContext context) {
        return new Function<OrderRequest, OrderResponse>(
                "queryOrderStatus",
                (request) -> orderService.getByOrderNo(request.orderNo()),
                "查询订单状态和物流信息",
                context
        );
    }

    /**
     * 取消订单
     */
    @Bean
    @Description("取消指定订单,仅未发货订单可取消")
    public Function<CancelRequest, CancelResponse> cancelOrder(
            ApplicationContext context) {
        return new Function<CancelRequest, CancelResponse>(
                "cancelOrder",
                (request) -> orderService.cancel(request.orderNo(), request.reason()),
                "取消未发货订单",
                context
        );
    }
}

/**
 * 请求记录 - 自动序列化为JSON Schema
 */
public record OrderRequest(
    @JsonPropertyDescription("订单编号,格式为ORD开头加数字") String orderNo,
    @JsonPropertyDescription("可选的用户ID") String userId
) {}

public record OrderResponse(
    String orderNo,
    String status,
    String logisticsInfo,
    LocalDateTime updateTime
) {}

public record CancelRequest(
    @JsonPropertyDescription("订单编号") String orderNo,
    @JsonPropertyDescription("取消原因") String reason
) {}

public record CancelResponse(
    boolean success,
    String message,
    BigDecimal refundAmount
) {}

2. 使用函数的高级方式 - FunctionCallbackResolver

java 复制代码
@Service
public class AssistantService {

    private final ChatClient chatClient;
    private final OrderService orderService;
    private final UserService userService;
    private final ProductService productService;

    /**
     * 注册函数到ChatClient
     */
    public String assistantChat(String userMessage) {
        return chatClient.prompt()
                .user(userMessage)
                .functions("queryOrderStatus", "cancelOrder", "createOrder",
                          "getProductInfo", "getUserInfo")
                .call()
                .content();
    }

    /**
     * 使用FunctionCallback - 更细粒度控制
     */
    public String assistantWithCallback(String userMessage) {
        FunctionCallback weatherCallback = FunctionCallback.builder()
                .function("getWeather", 
                    (WeatherRequest req) -> WeatherService.getWeather(req))
                .description("获取指定城市的实时天气信息")
                .inputType(WeatherRequest.class)
                .build();

        return chatClient.prompt()
                .user(userMessage)
                .functionCallbacks(weatherCallback)
                .call()
                .content();
    }

    /**
     * 动态函数注册 - 根据上下文加载不同函数集
     */
    public String contextualChat(String userMessage, String domain) {
        FunctionCallbackResolver resolver = switch (domain) {
            case "order" -> orderFunctionResolver();
            case "product" -> productFunctionResolver();
            case "user" -> userFunctionResolver();
            default -> FunctionCallbackResolver.empty();
        };

        return chatClient.prompt()
                .user(userMessage)
                .functionCallbacks(resolver)
                .call()
                .content();
    }

    private FunctionCallbackResolver orderFunctionResolver() {
        List<FunctionCallback> callbacks = List.of(
            FunctionCallback.builder()
                .function("queryOrder", this::queryOrder)
                .description("查询订单信息,支持按订单号、用户ID、日期范围查询")
                .inputType(OrderQueryRequest.class)
                .build(),
            FunctionCallback.builder()
                .function("createOrder", this::createOrder)
                .description("创建新订单,需要商品列表和收货地址")
                .inputType(CreateOrderRequest.class)
                .build(),
            FunctionCallback.builder()
                .function("refundOrder", this::refundOrder)
                .description("申请订单退款")
                .inputType(RefundRequest.class)
                .build()
        );
        return FunctionCallbackResolver.from(callbacks);
    }

    private String queryOrder(OrderQueryRequest req) {
        Order order = orderService.findOrders(
            req.orderNo(), req.userId(), req.startDate(), req.endDate()
        );
        return objectMapper.writeValueAsString(order);
    }

    private String createOrder(CreateOrderRequest req) {
        OrderResult result = orderService.create(
            req.userId(), req.items(), req.address()
        );
        return objectMapper.writeValueAsString(result);
    }

    private String refundOrder(RefundRequest req) {
        RefundResult result = orderService.refund(
            req.orderNo(), req.items(), req.reason()
        );
        return objectMapper.writeValueAsString(result);
    }
}

3. 异步函数调用

java 复制代码
@Service
public class AsyncFunctionService {

    private final AsyncTaskExecutor executor;
    private final ChatClient chatClient;

    /**
     * 支持异步执行的函数
     */
    @Bean
    @Description("提交异步数据分析任务,返回任务ID")
    public Function<AnalysisRequest, TaskResponse> submitAnalysisTask(
            ApplicationContext context) {
        return new Function<AnalysisRequest, TaskResponse>(
                "submitAnalysisTask",
                (request) -> {
                    String taskId = UUID.randomUUID().toString();
                    CompletableFuture.runAsync(() -> {
                        performAnalysis(taskId, request);
                    }, executor);
                    return new TaskResponse(taskId, "SUBMITTED", 
                        "任务已提交,使用getAnalysisResult查询结果");
                },
                "提交异步数据分析任务",
                context
        );
    }

    @Bean
    @Description("查询异步分析任务的结果")
    public Function<TaskQueryRequest, AnalysisResponse> getAnalysisResult(
            ApplicationContext context) {
        return new Function<TaskQueryRequest, AnalysisResponse>(
                "getAnalysisResult",
                (request) -> {
                    AnalysisResult result = analysisCache.get(request.taskId());
                    if (result == null) {
                        return new AnalysisResponse(request.taskId(), 
                            "PROCESSING", "任务仍在处理中,请稍后再查询");
                    }
                    return new AnalysisResponse(request.taskId(), 
                        "COMPLETED", objectMapper.writeValueAsString(result));
                },
                "查询异步分析任务的结果",
                context
        );
    }

    /**
     * 长时间运行的任务函数 - 流式反馈
     */
    public Flux<String> chatWithProgress(String message) {
        return Flux.create(sink -> {
            sink.next("开始处理...\n");
            
            chatClient.prompt()
                    .user(message)
                    .functions("submitAnalysisTask", "getAnalysisResult")
                    .advisors(a -> a.param("onFunctionCall", (Consumer<String>) funcName -> 
                        sink.next("正在执行: " + funcName + "\n")))
                    .stream()
                    .content()
                    .doOnNext(chunk -> sink.next(chunk))
                    .doOnComplete(() -> {
                        sink.next("\n处理完成");
                        sink.complete();
                    })
                    .subscribe();
        });
    }

    private void performAnalysis(String taskId, AnalysisRequest request) {
        try {
            // 模拟耗时分析
            Thread.sleep(5000);
            AnalysisResult result = new AnalysisResult();
            result.setTaskId(taskId);
            result.setData(List.of("分析结论1", "分析结论2"));
            analysisCache.put(taskId, result);
        } catch (InterruptedException e) {
            Thread.currentThread().interrupt();
        }
    }
}

4. 函数结果安全性处理

java 复制代码
/**
 * 函数执行器 - 添加超时、重试、降级
 */
@Component
public class SafeFunctionExecutor {

    private final FunctionRegistry functionRegistry;
    private final ScheduledExecutorService scheduler = 
        Executors.newScheduledThreadPool(4);

    /**
     * 带超时的函数执行
     */
    public Object executeWithTimeout(String functionName, String args, 
                                      long timeoutMs) {
        CompletableFuture<Object> future = CompletableFuture.supplyAsync(() ->
            execute(functionName, args)
        );

        try {
            return future.get(timeoutMs, TimeUnit.MILLISECONDS);
        } catch (TimeoutException e) {
            future.cancel(true);
            return Map.of("error", "函数执行超时", "function", functionName);
        } catch (Exception e) {
            return Map.of("error", e.getMessage(), "function", functionName);
        }
    }

    /**
     * 带重试的函数执行
     */
    public Object executeWithRetry(String functionName, String args, 
                                    int maxRetries) {
        int attempts = 0;
        Exception lastException = null;
        
        while (attempts <= maxRetries) {
            try {
                return execute(functionName, args);
            } catch (Exception e) {
                attempts++;
                lastException = e;
                if (attempts <= maxRetries) {
                    try {
                        Thread.sleep((long) Math.pow(2, attempts) * 1000);
                    } catch (InterruptedException ie) {
                        Thread.currentThread().interrupt();
                    }
                }
            }
        }
        
        return Map.of("error", "重试" + maxRetries + "次后仍失败: " 
            + lastException.getMessage());
    }

    /**
     * 带降级的函数执行
     */
    public Object executeWithFallback(String functionName, String args,
                                       String fallbackResponse) {
        try {
            return execute(functionName, args);
        } catch (Exception e) {
            log.warn("函数 {} 执行失败,使用降级响应: {}", functionName, e.getMessage());
            return Map.of(
                "error", "函数调用失败",
                "fallback", fallbackResponse,
                "function", functionName
            );
        }
    }

    private Object execute(String functionName, String args) {
        Method method = functionRegistry.resolveMethod(functionName);
        Object target = functionRegistry.resolveTarget(functionName);
        Object[] parsedArgs = functionRegistry.deserializeArgs(method, args);
        return method.invoke(target, parsedArgs);
    }
}

5. 函数调用链 - 多步骤复杂任务

java 复制代码
@Service
public class MultiStepFunctionService {

    /**
     * 复杂业务流程:下单 → 支付 → 查物流
     * LLM会自动决定调用顺序和参数传递
     */
    public String complexOrderFlow(String userRequest) {
        return chatClient.prompt()
                .system("""
                    你是一个智能订单助手。请按以下流程处理用户请求:
                    1. 先理解用户需求,提取关键信息
                    2. 需要时使用函数查询或执行操作
                    3. 多步骤任务按依赖顺序调用函数
                    4. 遇到错误时告知用户并提供解决方案
                    """)
                .user(userRequest)
                .functions(
                    "searchProducts",    // 搜索商品
                    "addToCart",         // 加入购物车
                    "createOrder",       // 创建订单
                    "getPaymentUrl",     // 获取支付链接
                    "checkPaymentStatus", // 检查支付状态
                    "queryLogistics"     // 查询物流
                )
                .call()
                .content();
    }

    /**
     * 带条件分支的函数集 - LLM根据上下文决定调用路径
     */
    public String conditionalFlow(String userRequest, String intent) {
        return switch (intent) {
            case "purchase" -> chatClient.prompt()
                    .system("帮助用户完成购买流程")
                    .user(userRequest)
                    .functions("searchProducts", "addToCart", "createOrder", 
                              "applyCoupon", "getPaymentUrl")
                    .call()
                    .content();
            
            case "afterSales" -> chatClient.prompt()
                    .system("处理售后问题,包括退款、换货、投诉")
                    .user(userRequest)
                    .functions("queryOrder", "applyRefund", "applyReturn",
                              "escalateComplaint", "trackShipment")
                    .call()
                    .content();
            
            default -> chatClient.prompt()
                    .system("智能客服")
                    .user(userRequest)
                    .functions("searchFAQ", "transferToHuman", "getStoreInfo")
                    .call()
                    .content();
        };
    }
}

6. 自定义函数描述与Schema生成

java 复制代码
/**
 * 高级函数描述 - 通过自定义注解增强Schema
 */
@Target(ElementType.METHOD)
@Retention(RetentionPolicy.RUNTIME)
public @interface AiFunction {
    String name();
    String description();
    boolean required() default true;
    String[] examples() default {};
}

public record ProductSearchRequest(
    @JsonPropertyDescription("搜索关键词,可以是商品名、品牌或类目")
    String keyword,
    
    @JsonPropertyDescription("价格区间下限,单位:元")
    @JsonProperty(required = false)
    BigDecimal minPrice,
    
    @JsonPropertyDescription("价格区间上限,单位:元")
    @JsonProperty(required = false)
    BigDecimal maxPrice,
    
    @JsonPropertyDescription("排序方式:price_asc, price_desc, sales, rating")
    @JsonProperty(required = false)
    String sortBy
) {}

@Component
public class SearchFunctions {

    @AiFunction(
        name = "advancedSearch",
        description = "智能商品搜索函数,支持多维度筛选和排序",
        examples = {
            "{\"keyword\":\"手机\",\"minPrice\":1000,\"maxPrice\":5000,\"sortBy\":\"rating\"}"
        }
    )
    public SearchResult advancedSearch(ProductSearchRequest request) {
        // 执行搜索逻辑
        return productService.search(
            request.keyword(), 
            request.minPrice(), 
            request.maxPrice(),
            request.sortBy()
        );
    }
}

总结

Spring AI Function Calling的关键技术点:

  1. 自动Schema生成: 通过Java反射自动生成JSON Schema描述
  2. 类型安全: record/class定义参数,编译期类型检查
  3. 函数链式调用: LLM自动决定多步骤任务中的函数调用顺序
  4. 安全控制: 超时、重试、降级等容错机制保障稳定性
  5. 动态注册: 根据场景动态加载不同函数集

参考博客: https://blog.csdn.net/badao_liumang_qizhi

相关推荐
鲨鱼辣钊10 小时前
【FastAPI筑基-Day19】APScheduler定时任务全实战|自动执行、动态启停、后台常驻
java·spring·fastapi
学长毕业设计14 小时前
基于SpringBoot的公益基金管理系统(源码+文档+讲解视频)
java·spring boot·后端
东小西15 小时前
【SAA实战】第 3 篇 · 工具调用全攻略:把业务能力交给 Agent 自己调度
java·后端·spring
东小西15 小时前
【SAA实战】第 4 篇 · Agent 短期记忆:saver 让 Agent 跨轮记得住(threadId 隔离)
java·后端·spring
许彰午15 小时前
22-DataCenter报文序列化
java·低代码·架构·状态模式
2601_9620652515 小时前
[MySQL] SQL优化之性能分析
java·sql·mysql
小范同学_15 小时前
JDK1.7 与 JDK1.8 HashMap 底层原理对比 + 数组并发扩容死循环详解
java·开发语言
予昊16 小时前
从零实现“在线五子棋对战“:WebSocket 实时通信 + 段位匹配
java·开发语言·网络·websocket
2601_9622035116 小时前
【SpringAI入门】初识SpringAI
java
weixin_4614085817 小时前
Mybatis-flex小记
java·开发语言·mybatis