Agent三层记忆架构:短期+长期+向量Java实现,再也不踩上下文丢失的坑

0 阅读6分钟

在这里插入图片描述 做AI Agent的同学,几乎都踩过记忆的坑。

原型调试的时候对话少,一切正常;一上线用户聊个十几轮,Agent突然就失忆了,前面说过的业务背景全忘了,用户得反复重复。更严重的,token直接爆掉,接口报错,会话直接中断。

很多人解决办法很粗暴:把历史对话全塞给大模型。结果就是token成本直线飙升,大模型输入超限被截断,关键业务信息丢失,甚至历史无关信息干扰当前回答,出现张冠李戴的幻觉。

生产环境的Agent记忆,从来不是「把对话存起来」这么简单,必须做分层设计。今天我们用Java完整实现三层记忆架构:短期会话记忆、长期向量记忆、实体记忆,代码可以直接复制进项目,也能写进简历。

本文属于Java+AI Agent落地系列,上一篇我们讲了Function Calling生产级实现,没看过可以翻历史文章

一、三层记忆架构,分别解决什么问题

1. 短期记忆(会话记忆) 保存当前会话的最近N轮对话,只存活在本次会话里,作用是维持上下文连贯性。核心是控制轮次上限,避免token无限膨胀。一般保留8-12轮,超出就淘汰最早的对话。

2. 长期记忆(向量记忆) 把重要的业务事实、用户偏好、历史决策,向量化存入Milvus向量库,跨会话持久化。新会话开启时,根据用户当前问题召回相关历史信息,不用把全部历史塞进prompt,既省token又能做到「还记得你上次说过」。

3. 实体记忆 提取用户ID、租户、业务单据号这类关键参数,单独结构化存储,避免大模型幻觉篡改业务主键。

避坑提醒:向量记忆不是召回越多越好,必须做相关性阈值过滤,无关记忆会引入噪声,反而让回答跑偏。

二、Java完整代码实现

1. 短期会话记忆(轮次限制)

package com.aiproject.agent.memory;

import java.util.ArrayDeque;
import java.util.Deque;

/**
 * 短期会话记忆:固定轮次,防止token爆炸
 * 只保存当前会话最近N轮对话,会话结束即销毁
 */
public class ShortTermMemory {

    // 生产环境建议8-12轮,根据模型窗口调整
    private static final int MAX_ROUNDS = 10;

    private final Deque<Message> history = new ArrayDeque<>();

    public void addUserMessage(String content) {
        addMessage("user", content);
    }

    public void addAssistantMessage(String content) {
        addMessage("assistant", content);
    }

    private void addMessage(String role, String content) {
        history.offerLast(new Message(role, content));
        // 超出最大轮次,淘汰最早的对话
        while (history.size() > MAX_ROUNDS) {
            history.pollFirst();
        }
    }

    /**
     * 组装成大模型可用的历史上下文
     */
    public String buildContext() {
        if (history.isEmpty()) {
            return "";
        }
        StringBuilder sb = new StringBuilder("【历史对话】\n");
        for (Message msg : history) {
            sb.append(msg.role()).append(": ").append(msg.content()).append("\n");
        }
        return sb.toString();
    }

    /**
     * 获取当前历史轮次
     */
    public int size() {
        return history.size();
    }

    public void clear() {
        history.clear();
    }

    public record Message(String role, String content) {}
}

2. 长期向量记忆(Milvus完整实现)

package com.aiproject.agent.memory;

import com.alibaba.fastjson2.JSON;
import com.alibaba.fastjson2.JSONObject;
import io.milvus.client.MilvusServiceClient;
import io.milvus.grpc.DataType;
import io.milvus.param.ConnectParam;
import io.milvus.param.collection.CreateCollectionParam;
import io.milvus.param.collection.FieldType;
import io.milvus.param.dml.InsertParam;
import io.milvus.param.dml.SearchParam;
import io.milvus.response.SearchResultsWrapper;
import okhttp3.*;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;

import javax.annotation.PostConstruct;
import java.io.IOException;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.concurrent.TimeUnit;

/**
 * 长期向量记忆:业务事实持久化,跨会话召回
 * 基于Milvus向量数据库 + 通义千问Embedding
 */
@Component
public class LongTermVectorMemory {

    @Value("${milvus.host:localhost}")
    private String milvusHost;

    @Value("${milvus.port:19530}")
    private int milvusPort;

    @Value("${llm.api-key}")
    private String apiKey;

    @Value("${llm.embedding-url:https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding}")
    private String embeddingUrl;

    @Value("${llm.embedding-model:text-embedding-v2}")
    private String embeddingModel;

    // 向量维度:text-embedding-v2是1536维
    private static final int VECTOR_DIM = 1536;
    // 集合名
    private static final String COLLECTION_NAME = "agent_long_term_memory";
    // 相关性阈值,低于此值的记忆不召回
    private static final float SIMILARITY_THRESHOLD = 0.65f;

    private MilvusServiceClient milvusClient;
    private final OkHttpClient httpClient = new OkHttpClient.Builder()
            .connectTimeout(10, TimeUnit.SECONDS)
            .readTimeout(30, TimeUnit.SECONDS)
            .build();

    @PostConstruct
    public void init() {
        milvusClient = new MilvusServiceClient(ConnectParam.newBuilder()
                .withHost(milvusHost)
                .withPort(milvusPort)
                .build());
        // 生产环境启动时确保集合存在
        ensureCollectionExists();
    }

    /**
     * 保存重要业务事实到向量库
     */
    public void saveFact(String tenantId, String userId, String fact) {
        try {
            List<Float> vector = embed(fact);
            List<InsertParam.Field> fields = new ArrayList<>();
            fields.add(new InsertParam.Field("tenant_id", Collections.singletonList(tenantId)));
            fields.add(new InsertParam.Field("user_id", Collections.singletonList(userId)));
            fields.add(new InsertParam.Field("content", Collections.singletonList(fact)));
            fields.add(new InsertParam.Field("vector", Collections.singletonList(vector)));

            milvusClient.insert(InsertParam.newBuilder()
                    .withCollectionName(COLLECTION_NAME)
                    .withFields(fields)
                    .build());
        } catch (Exception e) {
            throw new RuntimeException("保存长期记忆失败", e);
        }
    }

    /**
     * 根据当前query召回TopN相关记忆(带相关性过滤)
     */
    public List<String> recall(String tenantId, String query, int topN) {
        try {
            List<Float> queryVector = embed(query);
            SearchParam searchParam = SearchParam.newBuilder()
                    .withCollectionName(COLLECTION_NAME)
                    .withVectorFieldName("vector")
                    .withVectors(Collections.singletonList(queryVector))
                    .withTopK(topN)
                    // 租户隔离:只召回当前租户的记忆
                    .withExpr("tenant_id == \"" + tenantId + "\"")
                    .build();

            SearchResultsWrapper wrapper = new SearchResultsWrapper(
                    milvusClient.search(searchParam).getData().getResults());

            List<String> results = new ArrayList<>();
            List<?> rowRecords = wrapper.getRowRecords(0);
            for (Object record : rowRecords) {
                SearchResultsWrapper.RowRecord row = (SearchResultsWrapper.RowRecord) record;
                float score = row.getScore();
                // 相关性阈值过滤:低于阈值的记忆不召回,避免噪声
                if (score >= SIMILARITY_THRESHOLD) {
                    String content = (String) row.get("content");
                    results.add(content);
                }
            }
            return results;
        } catch (Exception e) {
            throw new RuntimeException("召回长期记忆失败", e);
        }
    }

    /**
     * 调用通义千问Embedding接口生成向量
     */
    private List<Float> embed(String text) throws IOException {
        JSONObject requestBody = new JSONObject();
        requestBody.put("model", embeddingModel);
        requestBody.put("input", new String[]{text});

        Request request = new Request.Builder()
                .url(embeddingUrl)
                .addHeader("Authorization", "Bearer " + apiKey)
                .addHeader("Content-Type", "application/json")
                .post(RequestBody.create(requestBody.toJSONString(),
                        MediaType.parse("application/json")))
                .build();

        try (Response response = httpClient.newCall(request).execute()) {
            String body = response.body().string();
            JSONObject json = JSON.parseObject(body);
            return json.getJSONArray("output")
                    .getJSONObject(0)
                    .getJSONArray("embedding")
                    .toJavaList(Float.class);
        }
    }

    /**
     * 确保集合存在(生产环境用运维脚本创建,这里仅作演示)
     */
    private void ensureCollectionExists() {
        try {
            FieldType tenantIdField = FieldType.newBuilder()
                    .withName("tenant_id").withDataType(DataType.VarChar).withMaxLength(64).build();
            FieldType userIdField = FieldType.newBuilder()
                    .withName("user_id").withDataType(DataType.VarChar).withMaxLength(64).build();
            FieldType contentField = FieldType.newBuilder()
                    .withName("content").withDataType(DataType.VarChar).withMaxLength(2000).build();
            FieldType vectorField = FieldType.newBuilder()
                    .withName("vector").withDataType(DataType.FloatVector).withDimension(VECTOR_DIM).build();

            milvusClient.createCollection(CreateCollectionParam.newBuilder()
                    .withCollectionName(COLLECTION_NAME)
                    .addFieldType(tenantIdField)
                    .addFieldType(userIdField)
                    .addFieldType(contentField)
                    .addFieldType(vectorField)
                    .build());
        } catch (Exception ignored) {
            // 集合已存在则忽略
        }
    }
}

3. 实体记忆(关键参数结构化存储)

package com.aiproject.agent.memory;

import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;

/**
 * 实体记忆:关键业务参数结构化存储,防止大模型幻觉篡改
 * 生产环境可替换为Redis或数据库持久化
 */
public class EntityMemory {

    // key: sessionId, value: 实体参数Map
    private final Map<String, Map<String, String>> sessionEntities = new ConcurrentHashMap<>();

    /**
     * 保存实体参数
     */
    public void putEntity(String sessionId, String key, String value) {
        sessionEntities.computeIfAbsent(sessionId, k -> new ConcurrentHashMap<>())
                .put(key, value);
    }

    /**
     * 获取实体参数
     */
    public String getEntity(String sessionId, String key) {
        Map<String, String> entities = sessionEntities.get(sessionId);
        return entities == null ? null : entities.get(key);
    }

    /**
     * 组装实体信息注入prompt(确保大模型使用正确的业务参数)
     */
    public String buildEntityContext(String sessionId) {
        Map<String, String> entities = sessionEntities.get(sessionId);
        if (entities == null || entities.isEmpty()) {
            return "";
        }
        StringBuilder sb = new StringBuilder("【业务实体参数】\n");
        entities.forEach((k, v) -> sb.append(k).append(": ").append(v).append("\n"));
        sb.append("以上参数为系统确认值,回答时必须使用,不得自行编造或修改。\n");
        return sb.toString();
    }

    public void clear(String sessionId) {
        sessionEntities.remove(sessionId);
    }
}

三、生产踩过的3个坑

  1. 不要全量召回向量记忆:必须设置相似度阈值(建议0.6-0.7),低于阈值的记忆坚决不用,否则噪声会引发幻觉。
  2. 不要把用户输入全存长期记忆:只存确认过的业务事实和用户偏好,闲聊内容不要落库,否则记忆库会越来越脏,召回质量下降。
  3. 短期记忆轮次不是越多越好:超过12轮边际收益骤降,token成本翻倍,不如把关键信息沉淀到长期记忆。

📦 落地资源推荐

跑向量库和Agent后端,推荐阿里云 轻量应用服务器,2核4G跑Milvus+SpringBoot完全够用,新用户价格友好再加上通过推广渠道的9折优惠券跌破低价,开发调试足够用,有部署需求点击了解www.aliyun.com/daily-act/e…

💡 领取完整源码

关注图片上水印,领取完整SpringBoot工程包含Milvus配置、Embedding工具、记忆过滤逻辑、单元测试。