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 轮,且长短不均)。
*
*
投影规则:
*
* - system 消息永在首位;
* - 当前轮(最后一个 UserMessage 起)原样保留——进行中的 工具调用/结果 配对不可破坏;
* - 历史轮从新到旧按**整轮**(UserMessage 边界)纳入,旧轮的工具结果压成一行摘要,
* 预算(约 {@code charBudget} 字符 ≈ 中文 token 数)用尽即止。
*
*/
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();
}
}
}