package com.xly.agent; import dev.langchain4j.agent.tool.ToolExecutionRequest; import dev.langchain4j.data.message.AiMessage; import dev.langchain4j.data.message.ChatMessage; import dev.langchain4j.data.message.SystemMessage; import dev.langchain4j.data.message.ToolExecutionResultMessage; import dev.langchain4j.data.message.UserMessage; import dev.langchain4j.memory.ChatMemory; import dev.langchain4j.store.memory.chat.ChatMemoryStore; import java.util.ArrayList; import java.util.List; /** * 带**投影**的对话记忆:存储层保留完整消息({@link ChatMemoryStore},硬上限按整轮裁剪), * 读取层({@link #messages()})做 token 预算投影——替换按条数计窗的 MessageWindowChatMemory * (工具消息占条数导致真实轮次只有 5-8 轮,且长短不均)。 * *

投影规则: *

*/ public class ProjectedChatMemory implements ChatMemory { private static final int TOOL_DIGEST_LEN = 120; private final Object id; private final ChatMemoryStore store; private final int charBudget; private final int hardCapMessages; public ProjectedChatMemory(Object id, ChatMemoryStore store, int charBudget, int hardCapMessages) { this.id = id; this.store = store; this.charBudget = charBudget; this.hardCapMessages = hardCapMessages; } @Override public Object id() { return id; } @Override public void add(ChatMessage m) { List full = new ArrayList<>(store.getMessages(id)); if (m instanceof SystemMessage sm) { if (!full.isEmpty() && full.get(0) instanceof SystemMessage cur) { if (cur.text().equals(sm.text())) { return; } full.set(0, sm); } else { full.add(0, sm); } } else { full.add(m); trimToCap(full); } store.updateMessages(id, full); } @Override public List messages() { return project(new ArrayList<>(store.getMessages(id))); } @Override public void clear() { store.deleteMessages(id); } private List project(List full) { if (full.isEmpty()) { return full; } SystemMessage sys = full.get(0) instanceof SystemMessage s ? s : null; List body = full.subList(sys == null ? 0 : 1, full.size()); int lastUser = 0; for (int i = body.size() - 1; i >= 0; i--) { if (body.get(i) instanceof UserMessage) { lastUser = i; break; } } List tail = new ArrayList<>(body.subList(lastUser, body.size())); List head = new ArrayList<>(); int used = 0; int turnEnd = lastUser; for (int i = lastUser - 1; i >= 0 && used < charBudget; i--) { if (!(body.get(i) instanceof UserMessage)) { continue; } List turn = new ArrayList<>(); int size = 0; for (int k = i; k < turnEnd; k++) { ChatMessage c = collapse(body.get(k)); turn.add(c); size += approxLen(c); } if (used + size > charBudget && !head.isEmpty()) { break; } head.addAll(0, turn); used += size; turnEnd = i; } List out = new ArrayList<>(); if (sys != null) { out.add(sys); } out.addAll(head); out.addAll(tail); return out; } /** 历史轮的工具结果压成一行摘要(当前轮不经过此路径,配对结构完整)。 */ private static ChatMessage collapse(ChatMessage m) { if (m instanceof ToolExecutionResultMessage t) { String txt = t.text() == null ? "" : t.text().replace('\n', ' ').trim(); if (txt.length() > TOOL_DIGEST_LEN) { txt = txt.substring(0, TOOL_DIGEST_LEN) + "…"; } return ToolExecutionResultMessage.from(t.id(), t.toolName(), txt); } return m; } private static int approxLen(ChatMessage m) { if (m instanceof UserMessage u && u.hasSingleText()) { return u.singleText().length(); } if (m instanceof AiMessage a) { int n = a.text() == null ? 0 : a.text().length(); if (a.hasToolExecutionRequests()) { for (ToolExecutionRequest r : a.toolExecutionRequests()) { n += (r.arguments() == null ? 0 : r.arguments().length()) + 20; } } return n; } if (m instanceof ToolExecutionResultMessage t) { return t.text() == null ? 0 : t.text().length(); } return 50; } /** 存储硬上限:超限时从最旧的整轮开始删(system 保留)。 */ private void trimToCap(List full) { int start = !full.isEmpty() && full.get(0) instanceof SystemMessage ? 1 : 0; while (full.size() > hardCapMessages) { int next = -1; for (int i = start + 1; i < full.size(); i++) { if (full.get(i) instanceof UserMessage) { next = i; break; } } if (next < 0) { break; } full.subList(start, next).clear(); } } }