agent: keep thread metadata out of session history
This commit is contained in:
parent
f9c8b30f8d
commit
6bfbab2384
2 changed files with 43 additions and 17 deletions
|
|
@ -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, "; "))
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue