Spring AI 扩展:ChatMemory 会话记忆 + FunctionCalling 工具调用

2 阅读11分钟

文档简介:基于SpringBoot3 + SpringAI搭建,全程为可直接上线的代码、生产配置、避坑方案,无冗余理论,快速落地AI会话与工具能力。

前置核心依赖(直接覆盖pom.xml)

<dependencies>
    <!-- SpringWeb基础服务 -->
    <dependency>
        <groupId>org.springframework.boot</groupId>
        <artifactId>spring-boot-starter-web</artifactId>
    </dependency>
    <!-- SpringAI大模型核心 -->
    <dependency>
        <groupId>org.springframework.ai</groupId>
        <artifactId>spring-ai-starter-model-openai</artifactId>
    </dependency>
    <!-- Redis会话持久化 -->
    <dependency>
        <groupId>org.springframework.boot</groupId>
        <artifactId>spring-boot-starter-data-redis</artifactId>
    </dependency>
    <!-- Jackson序列化 -->
    <dependency>
        <groupId>com.fasterxml.jackson.core</groupId>
        <artifactId>jackson-databind</artifactId>
    </dependency>
    <!-- 链路追踪 -->
    <dependency>
        <groupId>org.springframework.boot</groupId>
        <artifactId>spring-boot-starter-actuator</artifactId>
    </dependency>
    <dependency>
        <groupId>io.micrometer.tracing</groupId>
        <artifactId>tracing-exporter-zipkin</artifactId>
    </dependency>
    <!-- 参数校验 -->
    <dependency>
        <groupId>org.springframework.boot</groupId>
        <artifactId>spring-boot-starter-validation</artifactId>
    </dependency>
    <dependency>
        <groupId>org.springdoc</groupId>
        <artifactId>springdoc-openapi-starter-webmvc-ui</artifactId>
        <version>2.2.0</version>
    </dependency>
    <!-- 熔断重试容错 -->
    <dependency>
        <groupId>io.github.resilience4j</groupId>
        <artifactId>resilience4j-spring-boot3</artifactId>
        <version>2.2.0</version>
    </dependency>
</dependencies>

一、ChatMemory 会话记忆(生产落地)

1. 核心架构与原理

1.1 核心作用

大模型请求无状态,ChatMemory 统一实现对话上下文的存储、读取、裁剪、过期管控,支撑多轮连贯对话,是AI拟人交互的核心基础。

1.2 核心接口体系

SpringAI会话体系仅两个核心接口,区分测试与生产实现:

  • ChatMemory:顶层接口,定义会话增、查、清核心能力

  • MessageWindow:会话窗口裁剪接口,防止上下文超限

实现类场景区分:

  • InMemoryChatMemory:内存会话,仅本地测试

  • RedisChatMemory:分布式持久化会话,生产唯一方案

1.3 架构流程

graph LR
A[ChatMemory顶层接口] --> B[内存实现-测试]
A --> C[Redis实现-生产]
B & C --> D[MessageWindow窗口裁剪]
D --> E[Advisor自动拦截上下文拼接]

1.4 核心方法

  • add():写入用户/模型对话消息

  • get():读取完整会话历史

  • clear():清空指定会话

  • truncate():自动裁剪老旧对话,控制上下文长度

2. 内存会话(测试专用)

2.1 原理

基于JVM ConcurrentHashMap实现,开箱即用、无需中间件,仅支持单机本地调试。

2.2 测试代码

import org.springframework.ai.chat.memory.InMemoryChatMemory;
import org.springframework.stereotype.Controller;

@Controller
public class MemoryChatDemoController {
    private final InMemoryChatMemory chatMemory = new InMemoryChatMemory();

    public String chat(String sessionId, String message) {
        chatMemory.add(sessionId, message);
        var history = chatMemory.get(sessionId);
        return "对话完成";
    }
}

2.3 生产缺陷

  • 无持久化,服务重启会话丢失

  • 不支持分布式,多实例部署会话断裂

  • 无自动清理,闲置会话堆积引发内存溢出

  • 不支持跨端会话适配

3. Redis 生产级持久化会话

3.1 核心优势

支持分布式会话共享、自动过期、活跃续期、双维度消息裁剪、并发防冲突,适配所有生产场景。

3.2 Redis键设计规范

  • 统一前缀:springai:chat:memory:

  • 完整Key:springai:chat:memory:{sessionId}(绑定用户ID)

  • 存储结构:List有序链表,保证对话时序

  • 过期策略:默认24h过期,活跃自动续期

3.3 完整生产实现

import org.springframework.ai.chat.memory.ChatMemory;
import org.springframework.ai.chat.messages.Message;
import org.springframework.data.redis.core.StringRedisTemplate;
import org.springframework.stereotype.Component;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.type.CollectionType;
import lombok.extern.slf4j.Slf4j;

import java.util.List;
import java.util.concurrent.TimeUnit;

/**
 * 生产级Redis会话记忆
 * 能力:序列化、分布式锁、自动续期、轮数+Token双裁剪、异常兜底
 */
@Slf4j
@Component
public class RedisChatMemory implements ChatMemory {
    private static final String KEY_PREFIX = "springai:chat:memory:";
    private static final String LOCK_PREFIX = "springai:chat:lock:";
    private static final long EXPIRE_SECONDS = 86400;
    private static final int MAX_ROUND_SIZE = 10;
    private static final int MAX_TOKEN_LIMIT = 4096;

    private final StringRedisTemplate redisTemplate;
    private final ObjectMapper objectMapper;

    public RedisChatMemory(StringRedisTemplate redisTemplate, ObjectMapper objectMapper) {
        this.redisTemplate = redisTemplate;
        this.objectMapper = objectMapper;
    }

    @Override
    public void add(String sessionId, Message message) {
        String key = KEY_PREFIX + sessionId;
        String lockKey = LOCK_PREFIX + sessionId;
        // 分布式锁防并发错乱
        Boolean lock = redisTemplate.opsForValue().setIfAbsent(lockKey, "lock", 3, TimeUnit.SECONDS);
        if (!Boolean.TRUE.equals(lock)) {
            log.warn("[会话并发冲突] sessionId:{}", sessionId);
            return;
        }
        try {
            String json = objectMapper.writeValueAsString(message);
            redisTemplate.opsForList().rightPush(key, json);
            redisTemplate.expire(key, EXPIRE_SECONDS, TimeUnit.SECONDS);
            truncateByRound(key);
            truncateByToken(sessionId);
        } catch (Exception e) {
            log.error("[会话存储失败] sessionId:{}", sessionId, e);
        } finally {
            redisTemplate.delete(lockKey);
        }
    }

    @Override
    public List<Message> get(String sessionId) {
        String key = KEY_PREFIX + sessionId;
        try {
            List<String> jsonList = redisTemplate.opsForList().range(key, 0, -1);
            if (jsonList == null || jsonList.isEmpty()) {
                return List.of();
            }
            CollectionType listType = objectMapper.getTypeFactory()
                    .constructCollectionType(List.class, Message.class);
            return objectMapper.readValue(objectMapper.writeValueAsString(jsonList), listType);
        } catch (Exception e) {
            log.error("[读取会话失败] sessionId:{}", sessionId, e);
            return List.of();
        }
    }

    @Override
    public void clear(String sessionId) {
        String key = KEY_PREFIX + sessionId;
        redisTemplate.delete(key);
    }

    // 按对话轮数裁剪
    private void truncateByRound(String key) {
        Long size = redisTemplate.opsForList().size(key);
        if (size != null && size > MAX_ROUND_SIZE) {
            redisTemplate.opsForList().trim(key, size - MAX_ROUND_SIZE, -1);
        }
    }

    // 按Token阈值裁剪,防止上下文超限
    private void truncateByToken(String sessionId) {
        List<Message> messageList = get(sessionId);
        int totalToken = messageList.stream().mapToInt(msg -> msg.getContent().length()).sum();
        if (totalToken > MAX_TOKEN_LIMIT) {
            clear(sessionId);
            log.info("[会话Token超限清空] sessionId:{}", sessionId);
        }
    }
}

3.4 Jackson序列化全局配置

解决SpringAI消息多类型序列化、反序列化异常

import com.fasterxml.jackson.annotation.JsonTypeInfo;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.jsontype.impl.LaissezFaireSubTypeValidator;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;

@Configuration
public class JacksonConfig {
    @Bean
    public ObjectMapper objectMapper() {
        ObjectMapper objectMapper = new ObjectMapper();
        objectMapper.activateDefaultTyping(
                LaissezFaireSubTypeValidator.instance,
                ObjectMapper.DefaultTyping.NON_FINAL,
                JsonTypeInfo.As.PROPERTY
        );
        return objectMapper;
    }
}

3.5 Redis生产连接池配置

解决高并发连接耗尽、超时阻塞问题

spring:
  data:
    redis:
      host: localhost
      port: 6379
      password:
      timeout: 2000ms
      lettuce:
        pool:
          max-active: 32
          max-idle: 16
          min-idle: 8
          max-wait: 1000ms

3.6 生产高级优化策略

  • 服务降级:Redis宕机自动切内存会话,保障服务不中断

  • 数据合规:闲置会话定时清理、数据脱敏返回

  • 多端隔离:sessionId拼接用户ID+设备ID,实现多设备独立会话

4. 全自动会话拦截器配置

4.1 核心作用

替代手动读写会话,自动完成上下文拼接、消息存储、过期续期,全局统一生效。

4.2 全局配置代码

import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor;
import org.springframework.ai.chat.memory.ChatMemory;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;

@Configuration
public class AiMemoryAdvisorConfig {
    private final ChatMemory redisChatMemory;
    private static final int MAX_ROUND_SIZE = 10;

    public AiMemoryAdvisorConfig(ChatMemory redisChatMemory) {
        this.redisChatMemory = redisChatMemory;
    }

    @Bean
    public MessageChatMemoryAdvisor messageChatMemoryAdvisor() {
        return MessageChatMemoryAdvisor.builder(redisChatMemory)
                .maxChatMemorySize(MAX_ROUND_SIZE)
                .build();
    }
}

4.3 自动执行链路

请求拦截 → 读取Redis历史会话 → 拼接完整Prompt → 模型生成结果 → 自动存储对话 → 刷新会话过期时间

二、FunctionCalling 工具调用(生产全套能力)

1. 核心原理

1.1 核心价值

弥补大模型无法读取业务数据、无法执行业务逻辑的短板,由大模型决策调用、SpringAI落地执行,打通AI与本地业务能力。

1.2 调用时序

sequenceDiagram
用户->模型: 业务提问
模型->SpringAI: 返回工具调用指令(方法名+参数)
SpringAI->本地工具: 执行自定义业务方法
本地工具->SpringAI: 返回真实业务数据
SpringAI->模型: 回填工具结果
模型->用户: 生成自然语言最终回答

2. @Tool 注解规范与基础案例

2.1 注解核心参数

  • name:工具唯一标识,模型调用匹配依据

  • description:工具功能描述,决定调用准确率

  • parameters:参数释义,规范模型入参格式

2.2 基础工具实现(含参数与权限校验)

import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.stereotype.Service;

@Slf4j
@Service
public class OrderToolService {
    @Tool(
            name = "queryOrderStatus",
            description = "查询用户订单状态,需传入订单号和登录用户ID",
            parameters = {"orderNo:订单唯一编号", "userId:当前登录用户ID"}
    )
    public String queryOrderStatus(String orderNo, String userId) {
        if (orderNo == null || orderNo.isBlank()) {
            return "参数异常:订单号不能为空";
        }
        if (userId == null || userId.isBlank()) {
            return "权限异常:用户未登录";
        }
        try {
            log.info("[工具调用] 订单查询 userId:{},orderNo:{}", userId, orderNo);
            String result = "订单[" + orderNo + "]状态:已发货,归属用户:" + userId;
            log.info("[工具执行成功] 结果:{}", result);
            return result;
        } catch (Exception e) {
            log.error("[工具执行异常] 订单查询失败", e);
            return "订单查询失败,请稍后重试";
        }
    }
}

3. 工具注册方式

3.1 自动注册

SpringAI自动扫描@Tool注解Bean,零配置生效,适用于固定业务工具。

3.2 手动动态注册

支持工具热更新、动态启停,无需重启服务。

import org.springframework.ai.tool.ToolCallback;
import org.springframework.ai.tool.ToolCallbacks;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import java.util.List;

@Configuration
public class ToolRegisterConfig {
    @Bean
    public List<ToolCallback> toolCallbacks(OrderToolService orderToolService) {
        return ToolCallbacks.from(orderToolService);
    }
}

4. 工具调用基础配置

4.1 工具死循环防护

限制单轮对话最大工具调用次数,避免模型死循环卡死服务。

import org.springframework.ai.chat.client.ChatClient;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;

@Configuration
public class ChatClientLimitConfig {
    @Bean
    public ChatClient chatClient(ChatClient.Builder builder) {
        // 单轮最多3次工具调用
        return builder.defaultToolCallMaxIterations(3).build();
    }
}

5. 生产级高阶能力

5.1 JSON Schema 标准化参数校验

注解驱动自动校验参数,替代手写if判断,统一规范、适配模型自动传参。

5.1.1 参数实体定义
import io.swagger.v3.oas.annotations.media.Schema;
import jakarta.validation.constraints.NotBlank;
import lombok.Data;

@Data
@Schema(description = "订单查询工具请求参数")
public class OrderQueryParam {
    @NotBlank(message = "订单号不能为空")
    @Schema(description = "订单唯一编号", example = "20260804")
    private String orderNo;

    @NotBlank(message = "用户ID不能为空")
    @Schema(description = "当前登录用户ID", example = "10001")
    private String userId;
}

5.1.2 标准化校验工具类
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.ai.tool.annotation.ToolParam;
import org.springframework.stereotype.Service;
import org.springframework.validation.annotation.Validated;

import javax.validation.Valid;

@Slf4j
@Service
@Validated
public class OrderSchemaToolService {
    @Tool(
            name = "queryOrderStatusBySchema",
            description = "根据订单号查询用户订单状态,用于用户咨询订单物流、发货、签收状态场景"
    )
    public String queryOrderStatus(@Valid @ToolParam OrderQueryParam param) {
        String orderNo = param.getOrderNo();
        String userId = param.getUserId();
        log.info("[Schema工具调用] 用户{}查询订单{}", userId, orderNo);
        try {
            String result = "订单[" + orderNo + "]状态:已发货,归属用户:" + userId;
            log.info("[Schema工具执行成功] 结果:{}", result);
            return result;
        } catch (Exception e) {
            log.error("[Schema工具执行异常]", e);
            return "订单查询失败,请稍后重试";
        }
    }
}

5.2 重试+熔断降级容错机制

基于Resilience4j实现瞬时故障自愈、故障隔离、超时拦截,避免工具异常拖垮整体服务。

5.2.1 生产容错配置
# Resilience4j工具容错配置
resilience4j:
  retry:
    instances:
      aiToolRetry:
        max-attempts: 3
        wait-duration: 1000ms
        enable-exponential-backoff: true
  circuitbreaker:
    instances:
      aiToolCircuitBreaker:
        sliding-window-size: 10
        failure-rate-threshold: 50
        wait-duration-in-open-state: 5000ms
        permitted-number-of-calls-in-half-open-state: 2
        register-health-indicator: true
  timelimiter:
    instances:
      aiToolTimeLimiter:
        timeout-duration: 3000ms

5.2.2 容错注解说明
  • @Retry:瞬时异常自动重试,实现故障自愈

  • @CircuitBreaker:持续异常触发熔断,隔离故障

  • @TimeLimiter:强制超时拦截,杜绝线程阻塞

5.2.3 容错工具完整实现
import io.github.resilience4j.circuitbreaker.annotation.CircuitBreaker;
import io.github.resilience4j.retry.annotation.Retry;
import io.github.resilience4j.timelimiter.annotation.TimeLimiter;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.ai.tool.annotation.ToolParam;
import org.springframework.stereotype.Service;
import org.springframework.validation.annotation.Validated;

import javax.validation.Valid;
import java.util.concurrent.CompletableFuture;

@Slf4j
@Service
@Validated
public class OrderSchemaToolService {

    private final ToolCacheManager toolCacheManager;
    private final AiTraceContextUtil traceContextUtil;

    public OrderSchemaToolService(ToolCacheManager toolCacheManager, AiTraceContextUtil traceContextUtil) {
        this.toolCacheManager = toolCacheManager;
        this.traceContextUtil = traceContextUtil;
    }

    @Tool(
            name = "queryOrderStatusBySchema",
            description = "根据订单号查询用户订单状态,用于用户咨询订单物流、发货、签收状态场景"
    )
    @Retry(name = "aiToolRetry", fallbackMethod = "queryOrderFallback")
    @CircuitBreaker(name = "aiToolCircuitBreaker", fallbackMethod = "queryOrderFallback")
    @TimeLimiter(name = "aiToolTimeLimiter")
    public CompletableFuture<String> queryOrderStatus(@Valid @ToolParam OrderQueryParam param) {
        return CompletableFuture.supplyAsync(() -> {
            traceContextUtil.putTraceContext();
            String traceId = traceContextUtil.getCurrentTraceId();
            String paramKey = param.getUserId() + "_" + param.getOrderNo();
            if (toolCacheManager.isRepeatRequest("queryOrderStatusBySchema", paramKey)) {
                log.warn("[工具防抖拦截][traceId:{}] userId:{},orderNo:{}", traceId, param.getUserId(), param.getOrderNo());
                return "请求过于频繁,请稍后再试";
            }
            log.info("[容错工具调用][traceId:{}] userId:{},orderNo:{}", traceId, param.getUserId(), param.getOrderNo());
            String result = "订单[" + param.getOrderNo() + "]状态:已发货,归属用户:" + param.getUserId();
            log.info("[容错工具执行成功][traceId:{}] 结果:{}", traceId, result);
            return result;
        });
    }

    // 熔断、重试、超时统一降级兜底
    public CompletableFuture<String> queryOrderFallback(OrderQueryParam param, Exception e) {
        return CompletableFuture.supplyAsync(() -> {
            traceContextUtil.putTraceContext();
            String traceId = traceContextUtil.getCurrentTraceId();
            log.error("[工具熔断降级][traceId:{}] userId:{},orderNo:{}", traceId, param.getUserId(), param.getOrderNo(), e);
            return "订单查询服务暂时繁忙,请您稍后重试,感谢理解!";
        });
    }
}

5.3 工具防抖缓存

拦截短时间重复请求,避免重复查询接口/数据库,降低资源消耗。

5.3.1 防抖缓存工具类
import lombok.extern.slf4j.Slf4j;
import org.springframework.data.redis.core.StringRedisTemplate;
import org.springframework.stereotype.Component;
import java.util.concurrent.TimeUnit;

@Slf4j
@Component
public class ToolCacheManager {
    private final StringRedisTemplate redisTemplate;
    private static final long TOOL_CACHE_EXPIRE = 5;
    private static final String TOOL_CACHE_PREFIX = "springai:tool:cache:";

    public ToolCacheManager(StringRedisTemplate redisTemplate) {
        this.redisTemplate = redisTemplate;
    }

    public boolean isRepeatRequest(String toolName, String paramKey) {
        String cacheKey = TOOL_CACHE_PREFIX + toolName + ":" + paramKey;
        Boolean exist = redisTemplate.opsForValue().setIfAbsent(cacheKey, "1", TOOL_CACHE_EXPIRE, TimeUnit.SECONDS);
        return !Boolean.TRUE.equals(exist);
    }
}

5.4 全链路TraceId追踪

解决异步、长连接日志链路断裂问题,单一TraceId贯穿全流程,便于线上故障排查。

5.4.1 链路追踪YML配置
# 全链路追踪配置
management:
  tracing:
    sampling:
      probability: 1.0
  endpoints:
    web:
      exposure:
        include: health,info,tracing

logging:
  pattern:
    console: "%d{yyyy-MM-dd HH:mm:ss.SSS} [%thread] [%X{traceId:-N/A}/%X{spanId:-N/A}] %-5level %logger{50} - %msg%n"
  level:
    root: INFO
    org.springframework.ai: INFO

5.4.2 链路上下文工具类
import io.micrometer.tracing.Span;
import io.micrometer.tracing.Tracer;
import lombok.RequiredArgsConstructor;
import org.slf4j.MDC;
import org.springframework.stereotype.Component;

@Component
@RequiredArgsConstructor
public class AiTraceContextUtil {
    private final Tracer tracer;
    private static final String TRACE_ID = "traceId";
    private static final String SPAN_ID = "spanId";

    public String getCurrentTraceId() {
        Span currentSpan = tracer.currentSpan();
        return currentSpan == null ? "NONE" : currentSpan.context().traceId();
    }

    public void putTraceContext() {
        Span currentSpan = tracer.currentSpan();
        if (currentSpan != null) {
            MDC.put(TRACE_ID, currentSpan.context().traceId());
            MDC.put(SPAN_ID, currentSpan.context().spanId());
        }
    }

    public void clearTraceContext() {
        MDC.remove(TRACE_ID);
        MDC.remove(SPAN_ID);
    }
}
5.4.3 全局异常链路改造
import lombok.extern.slf4j.Slf4j;
import org.springframework.web.bind.annotation.ExceptionHandler;
import org.springframework.web.bind.annotation.RestControllerAdvice;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;

@Slf4j
@RestControllerAdvice
public class AiGlobalExceptionHandler {
    private final AiTraceContextUtil traceContextUtil;

    public AiGlobalExceptionHandler(AiTraceContextUtil traceContextUtil) {
        this.traceContextUtil = traceContextUtil;
    }

    @ExceptionHandler(Exception.class)
    public Object handleException(Exception e) {
        traceContextUtil.putTraceContext();
        String traceId = traceContextUtil.getCurrentTraceId();
        log.error("[AI全局异常][traceId:{}]", traceId, e);
        if (e.getStackTrace().toString().contains("SseEmitter")) {
            return "对话服务异常,请稍后重试";
        }
        return "服务繁忙,对话请求失败,请稍后重试";
    }
}

三、SSE流式对话(生产落地)

1. 核心能力

基于SSE长连接实现打字机实时输出,完美兼容会话记忆、工具调用,支持超时回收、异常兜底,是AI对话生产标准方案。

2. 流式对话完整接口

import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;

import java.io.IOException;
import java.util.concurrent.TimeUnit;

@Slf4j
@RestController
public class StreamChatController {
    private final ChatClient chatClient;
    private final AiTraceContextUtil traceContextUtil;

    public StreamChatController(ChatClient.Builder chatClientBuilder, AiTraceContextUtil traceContextUtil) {
        this.chatClient = chatClientBuilder.build();
        this.traceContextUtil = traceContextUtil;
    }

    @GetMapping("/chat/stream")
    public SseEmitter streamChat(@RequestParam String userId, @RequestParam String msg) {
        traceContextUtil.putTraceContext();
        String traceId = traceContextUtil.getCurrentTraceId();
        log.info("[流式对话开始][traceId:{}] userId:{}", traceId, userId);
        SseEmitter sseEmitter = new SseEmitter(TimeUnit.SECONDS.toMillis(30));

        // 超时回调
        sseEmitter.onTimeout(() -> {
            log.warn("[流式对话超时][traceId:{}] userId:{}", traceId, userId);
            sseEmitter.complete();
        });

        // 异常回调
        sseEmitter.onError((throwable) -> {
            log.error("[流式对话异常中断][traceId:{}] userId:{}", traceId, userId, throwable);
        });

        // 流式核心逻辑
        try {
            chatClient.prompt()
                    .user(msg)
                    .advisors(advisor -> advisor.param("sessionId", userId))
                    .stream()
                    .content()
                    .doOnNext(content -> {
                        try {
                            sseEmitter.send(content);
                        } catch (IOException e) {
                            log.error("[流式推送失败][traceId:{}]", traceId, e);
                        }
                    })
                    .doOnComplete(() -> {
                        log.info("[流式对话完成][traceId:{}] userId:{}", traceId, userId);
                        sseEmitter.complete();
                    })
                    .doOnError(error -> {
                        log.error("[流式对话执行失败][traceId:{}] userId:{}", traceId, userId, error);
                        try {
                            sseEmitter.send("对话异常,请稍后重试");
                            sseEmitter.complete();
                        } catch (IOException e) {
                            log.error("[流式异常兜底失败][traceId:{}]", traceId, e);
                        }
                    })
                    .subscribe();
        } catch (Exception e) {
            log.error("[流式对话初始化失败][traceId:{}] userId:{}", traceId, userId, e);
            try {
                sseEmitter.send("系统繁忙,请稍后重试");
                sseEmitter.complete();
            } catch (IOException ex) {
                log.error("[初始化兜底失败][traceId:{}]", traceId, ex);
            }
        }
        return sseEmitter;
    }
}

3. 前端极简测试页面

<!DOCTYPE html>
<html lang="zh-CN">
<body>
<div id="result" style="white-space: pre-wrap;padding: 20px;"></div>
<script>
const userId = "10001";
const msg = "帮我查询一下我的订单20260804状态";
const eventSource = new EventSource(`/chat/stream?userId=${userId}&msg=${msg}`);
let resultDom = document.getElementById("result");

eventSource.onmessage = function (e) {
    resultDom.innerText += e.data;
};
eventSource.onclose = function () {
    resultDom.innerText += "\n\n【对话结束】";
    eventSource.close();
};
eventSource.onerror = function () {
    resultDom.innerText += "\n\n【对话异常中断】";
    eventSource.close();
};
</script>
</body>
</html>

四、项目整合与启动

1. 启动类

import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;

@SpringBootApplication
public class AiApplication {
    public static void main(String[] args) {
        SpringApplication.run(AiApplication.class, args);
    }
}

2. 普通对话接口

import org.springframework.ai.chat.client.ChatClient;
import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;

@RestController
public class ChatController {
    private final ChatClient chatClient;

    public ChatController(ChatClient.Builder chatClientBuilder) {
        this.chatClient = chatClientBuilder.build();
    }

    @GetMapping("/chat")
    public String chat(@RequestParam String userId, @RequestParam String msg) {
        return chatClient.prompt()
                .user(msg)
                .advisors(advisor -> advisor.param("sessionId", userId))
                .call()
                .content();
    }
}

3. 核心总结

  • ChatMemory:实现会话持久化,保障多轮对话连贯

  • FunctionCalling:打通本地业务数据,突破大模型能力限制

  • 高阶能力:参数校验、防抖、熔断重试、链路追踪、流式输出,全方位满足生产上线标准