001/*
002 *  Copyright (c) 2023-2026, Agents-Flex (fuhai999@gmail.com).
003 *  <p>
004 *  Licensed under the Apache License, Version 2.0 (the "License");
005 *  you may not use this file except in compliance with the License.
006 *  You may obtain a copy of the License at
007 *  <p>
008 *  http://www.apache.org/licenses/LICENSE-2.0
009 *  <p>
010 *  Unless required by applicable law or agreed to in writing, software
011 *  distributed under the License is distributed on an "AS IS" BASIS,
012 *  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
013 *  See the License for the specific language governing permissions and
014 *  limitations under the License.
015 */
016package com.agentsflex.core.prompt;
017
018import com.agentsflex.core.memory.ChatMemory;
019import com.agentsflex.core.memory.DefaultChatMemory;
020import com.agentsflex.core.message.*;
021
022import java.util.ArrayList;
023import java.util.Collection;
024import java.util.List;
025import java.util.function.Function;
026
027public class MemoryPrompt extends Prompt {
028
029    private ChatMemory memory = new DefaultChatMemory();
030
031    private SystemMessage systemMessage;
032
033    private int maxAttachedMessageCount = 100;
034
035    private boolean historyMessageTruncateEnable = false;
036    private int historyMessageTruncateLength = 1000;
037    private Function<String, String> historyMessageTruncateProcessor;
038
039    // 临时消息不回存入 memory,只会当做 “过程消息” 参与大模型交互
040    // 比如用于 Function call 等场景
041    private List<Message> temporaryMessages;
042
043    public SystemMessage getSystemMessage() {
044        return systemMessage;
045    }
046
047    public void setSystemMessage(String content) {
048        this.systemMessage = new SystemMessage(content);
049    }
050
051    public void setSystemMessage(SystemMessage systemMessage) {
052        this.systemMessage = systemMessage;
053    }
054
055    public int getMaxAttachedMessageCount() {
056        return maxAttachedMessageCount;
057    }
058
059    public void setMaxAttachedMessageCount(int maxAttachedMessageCount) {
060        this.maxAttachedMessageCount = maxAttachedMessageCount;
061    }
062
063    public boolean isHistoryMessageTruncateEnable() {
064        return historyMessageTruncateEnable;
065    }
066
067    public void setHistoryMessageTruncateEnable(boolean historyMessageTruncateEnable) {
068        this.historyMessageTruncateEnable = historyMessageTruncateEnable;
069    }
070
071    public int getHistoryMessageTruncateLength() {
072        return historyMessageTruncateLength;
073    }
074
075    public void setHistoryMessageTruncateLength(int historyMessageTruncateLength) {
076        this.historyMessageTruncateLength = historyMessageTruncateLength;
077    }
078
079    public Function<String, String> getHistoryMessageTruncateProcessor() {
080        return historyMessageTruncateProcessor;
081    }
082
083    public void setHistoryMessageTruncateProcessor(Function<String, String> historyMessageTruncateProcessor) {
084        this.historyMessageTruncateProcessor = historyMessageTruncateProcessor;
085    }
086
087    public MemoryPrompt() {
088    }
089
090    public MemoryPrompt(ChatMemory memory) {
091        this.memory = memory;
092    }
093
094    public void addMessage(Message message) {
095        memory.addMessage(message);
096    }
097
098    public void addUserMessage(String content) {
099        this.addMessage(new UserMessage(content));
100    }
101
102    public void addAiMessage(String content) {
103        this.addMessage(new AiMessage(content));
104    }
105
106    public void addMessageTemporary(Message message) {
107        if (temporaryMessages == null) {
108            temporaryMessages = new ArrayList<>();
109        }
110        temporaryMessages.add(message);
111    }
112
113    public void addMessages(Collection<? extends Message> messages) {
114        memory.addMessages(messages);
115    }
116
117    public ChatMemory getMemory() {
118        return memory;
119    }
120
121    public void setMemory(ChatMemory memory) {
122        this.memory = memory;
123    }
124
125    public List<Message> getTemporaryMessages() {
126        return temporaryMessages;
127    }
128
129    public void setTemporaryMessages(List<Message> temporaryMessages) {
130        this.temporaryMessages = temporaryMessages;
131    }
132
133    public void clearTemporaryMessages() {
134        temporaryMessages.clear();
135        temporaryMessages = null;
136    }
137
138    /**
139     * 清空所有消息
140     */
141    public void clear() {
142        memory.clear();
143        if (temporaryMessages != null) {
144            temporaryMessages.clear();
145        }
146    }
147
148    @Override
149    public List<Message> getMessages() {
150        List<Message> messages = memory.getMessages(maxAttachedMessageCount);
151        if (messages == null) {
152            messages = new ArrayList<>();
153        }
154
155        if (historyMessageTruncateEnable) {
156            for (int i = 0; i < messages.size(); i++) {
157                Message msg = messages.get(i);
158                if (msg instanceof AbstractTextMessage) {
159                    AbstractTextMessage<?> textMsg = (AbstractTextMessage<?>) msg;
160                    String content = textMsg.getContent();
161                    if (content == null) continue;
162
163                    // 应用自定义处理器或默认截断
164                    if (historyMessageTruncateProcessor != null) {
165                        content = historyMessageTruncateProcessor.apply(content);
166                    } else if (content.length() > historyMessageTruncateLength) {
167                        content = content.substring(0, historyMessageTruncateLength);
168                    }
169
170                    // 创建新实例,避免修改原始消息
171                    AbstractTextMessage<?> copied = textMsg.copy();
172                    copied.setContent(content);
173                    messages.set(i, copied);
174                }
175            }
176        }
177
178        // 插入系统消息
179        if (systemMessage != null) {
180            if (messages.isEmpty() || !(messages.get(0) instanceof SystemMessage)) {
181                messages.add(0, systemMessage);
182            }
183        }
184
185        //  添加临时消息(如果存在)
186        if (temporaryMessages != null && !temporaryMessages.isEmpty()) {
187            messages.addAll(new ArrayList<>(temporaryMessages));
188
189            // 使用后自动清理
190            temporaryMessages.clear();
191            temporaryMessages = null;
192        }
193
194        return messages;
195    }
196
197
198}