agent: keep thread metadata out of session history

This commit is contained in:
Alix-007 2026-03-29 20:52:02 +08:00
parent f9c8b30f8d
commit 6bfbab2384
2 changed files with 43 additions and 17 deletions

View file

@ -1476,15 +1476,6 @@ func (al *AgentLoop) runAgentLoop(
agent *AgentInstance, agent *AgentInstance,
opts processOptions, opts processOptions,
) (string, error) { ) (string, error) {
opts.UserMessage = formatUserMessageWithThreadMetadata(
opts.UserMessage,
opts.SenderDisplayName,
opts.SenderUsername,
opts.SenderID,
opts.MessageID,
opts.ReplyToMessageID,
)
// Record last channel for heartbeat notifications (skip internal channels and cli) // Record last channel for heartbeat notifications (skip internal channels and cli)
if opts.Channel != "" && opts.ChatID != "" && !constants.IsInternalChannel(opts.Channel) { if opts.Channel != "" && opts.ChatID != "" && !constants.IsInternalChannel(opts.Channel) {
channelKey := fmt.Sprintf("%s:%s", opts.Channel, opts.ChatID) channelKey := fmt.Sprintf("%s:%s", opts.Channel, opts.ChatID)
@ -1638,6 +1629,14 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er
summary = ts.agent.Sessions.GetSummary(ts.sessionKey) summary = ts.agent.Sessions.GetSummary(ts.sessionKey)
} }
ts.captureRestorePoint(history, summary) ts.captureRestorePoint(history, summary)
annotatedUserMessage := formatUserMessageWithThreadMetadata(
ts.userMessage,
ts.opts.SenderDisplayName,
ts.opts.SenderUsername,
ts.opts.SenderID,
ts.opts.MessageID,
ts.opts.ReplyToMessageID,
)
messages := ts.agent.ContextBuilder.BuildMessages( messages := ts.agent.ContextBuilder.BuildMessages(
history, history,
@ -1650,6 +1649,7 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er
ts.opts.SenderDisplayName, ts.opts.SenderDisplayName,
activeSkillNames(ts.agent, ts.opts)..., activeSkillNames(ts.agent, ts.opts)...,
) )
annotateCurrentUserMessageForLLM(messages, ts.userMessage, annotatedUserMessage)
cfg := al.GetConfig() cfg := al.GetConfig()
maxMediaSize := cfg.Agents.Defaults.GetMaxMediaSize() maxMediaSize := cfg.Agents.Defaults.GetMaxMediaSize()
@ -1680,6 +1680,7 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er
ts.opts.SenderID, ts.opts.SenderDisplayName, ts.opts.SenderID, ts.opts.SenderDisplayName,
activeSkillNames(ts.agent, ts.opts)..., activeSkillNames(ts.agent, ts.opts)...,
) )
annotateCurrentUserMessageForLLM(messages, ts.userMessage, annotatedUserMessage)
messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize) messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize)
} }
} }
@ -2052,6 +2053,7 @@ turnLoop:
nil, ts.channel, ts.chatID, ts.opts.SenderID, ts.opts.SenderDisplayName, nil, ts.channel, ts.chatID, ts.opts.SenderID, ts.opts.SenderDisplayName,
activeSkillNames(ts.agent, ts.opts)..., activeSkillNames(ts.agent, ts.opts)...,
) )
annotateCurrentUserMessageForLLM(messages, ts.userMessage, annotatedUserMessage)
callMessages = messages callMessages = messages
if gracefulTerminal { if gracefulTerminal {
callMessages = append(append([]providers.Message(nil), messages...), ts.interruptHintMessage()) callMessages = append(append([]providers.Message(nil), messages...), ts.interruptHintMessage())
@ -3555,6 +3557,26 @@ func inboundMetadata(msg bus.InboundMessage, key string) string {
return msg.Metadata[key] return msg.Metadata[key]
} }
func annotateCurrentUserMessageForLLM(
messages []providers.Message,
rawUserMessage string,
annotatedUserMessage string,
) {
if rawUserMessage == annotatedUserMessage {
return
}
for i := len(messages) - 1; i >= 0; i-- {
if messages[i].Role != "user" {
continue
}
if messages[i].Content != rawUserMessage {
continue
}
messages[i].Content = annotatedUserMessage
return
}
}
func formatUserMessageWithThreadMetadata( func formatUserMessageWithThreadMetadata(
content string, content string,
senderDisplayName string, senderDisplayName string,
@ -3574,11 +3596,15 @@ func formatUserMessageWithThreadMetadata(
if from != "" { if from != "" {
metaParts = append(metaParts, fmt.Sprintf("from:%s", from)) metaParts = append(metaParts, fmt.Sprintf("from:%s", from))
} }
msgMetaParts := make([]string, 0, 2)
if messageID != "" { if messageID != "" {
metaParts = append(metaParts, fmt.Sprintf("msg:#%s", messageID)) msgMetaParts = append(msgMetaParts, fmt.Sprintf("#%s", messageID))
} }
if replyToMessageID != "" { if replyToMessageID != "" {
metaParts = append(metaParts, fmt.Sprintf("reply_to:#%s", replyToMessageID)) msgMetaParts = append(msgMetaParts, fmt.Sprintf("reply_to:#%s", replyToMessageID))
}
if len(msgMetaParts) > 0 {
metaParts = append(metaParts, fmt.Sprintf("msgs:%s", strings.Join(msgMetaParts, ", ")))
} }
annotation := fmt.Sprintf("[%s]", strings.Join(metaParts, "; ")) annotation := fmt.Sprintf("[%s]", strings.Join(metaParts, "; "))

View file

@ -178,7 +178,7 @@ func TestFormatUserMessageWithThreadMetadata(t *testing.T) {
t.Run("includes from, message and reply IDs", func(t *testing.T) { t.Run("includes from, message and reply IDs", func(t *testing.T) {
got := formatUserMessageWithThreadMetadata("ping", "Alice", "@alice", "discord:1", "123", "120") got := formatUserMessageWithThreadMetadata("ping", "Alice", "@alice", "discord:1", "123", "120")
want := "[from:Alice (@alice); msg:#123; reply_to:#120]\nping" want := "[from:Alice (@alice); msgs:#123, reply_to:#120]\nping"
if got != want { if got != want {
t.Fatalf("formatUserMessageWithThreadMetadata() = %q, want %q", got, want) t.Fatalf("formatUserMessageWithThreadMetadata() = %q, want %q", got, want)
} }
@ -186,14 +186,14 @@ func TestFormatUserMessageWithThreadMetadata(t *testing.T) {
t.Run("empty content returns annotation only", func(t *testing.T) { t.Run("empty content returns annotation only", func(t *testing.T) {
got := formatUserMessageWithThreadMetadata("", "", "", "discord:1", "123", "") got := formatUserMessageWithThreadMetadata("", "", "", "discord:1", "123", "")
want := "[from:discord:1; msg:#123]" want := "[from:discord:1; msgs:#123]"
if got != want { if got != want {
t.Fatalf("formatUserMessageWithThreadMetadata() = %q, want %q", got, want) t.Fatalf("formatUserMessageWithThreadMetadata() = %q, want %q", got, want)
} }
}) })
} }
func TestProcessMessage_AnnotatesThreadMetadataInPromptAndHistory(t *testing.T) { func TestProcessMessage_AnnotatesThreadMetadataInPromptOnly(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-test-*") tmpDir, err := os.MkdirTemp("", "agent-test-*")
if err != nil { if err != nil {
t.Fatalf("Failed to create temp dir: %v", err) t.Fatalf("Failed to create temp dir: %v", err)
@ -245,7 +245,7 @@ func TestProcessMessage_AnnotatesThreadMetadataInPromptAndHistory(t *testing.T)
t.Fatal("provider did not receive any messages") t.Fatal("provider did not receive any messages")
} }
wantAnnotated := "[from:Alice (@alice); msg:#123; reply_to:#120]\nhello" wantAnnotated := "[from:Alice (@alice); msgs:#123, reply_to:#120]\nhello"
lastMessage := provider.lastMessages[len(provider.lastMessages)-1] lastMessage := provider.lastMessages[len(provider.lastMessages)-1]
if lastMessage.Role != "user" || lastMessage.Content != wantAnnotated { if lastMessage.Role != "user" || lastMessage.Content != wantAnnotated {
t.Fatalf("last provider message = %+v, want user annotation %q", lastMessage, wantAnnotated) t.Fatalf("last provider message = %+v, want user annotation %q", lastMessage, wantAnnotated)
@ -263,8 +263,8 @@ func TestProcessMessage_AnnotatesThreadMetadataInPromptAndHistory(t *testing.T)
if len(history) != 2 { if len(history) != 2 {
t.Fatalf("expected history len=2, got %d", len(history)) t.Fatalf("expected history len=2, got %d", len(history))
} }
if history[0].Role != "user" || history[0].Content != wantAnnotated { if history[0].Role != "user" || history[0].Content != "hello" {
t.Fatalf("history user message = %+v, want %q", history[0], wantAnnotated) t.Fatalf("history user message = %+v, want %q", history[0], "hello")
} }
} }