agent: restore metadata-only prompt injection
This commit is contained in:
parent
6bfbab2384
commit
d976be1162
2 changed files with 82 additions and 6 deletions
|
|
@ -1649,7 +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)
|
messages = annotateCurrentUserMessageForLLM(messages, ts.userMessage, annotatedUserMessage)
|
||||||
|
|
||||||
cfg := al.GetConfig()
|
cfg := al.GetConfig()
|
||||||
maxMediaSize := cfg.Agents.Defaults.GetMaxMediaSize()
|
maxMediaSize := cfg.Agents.Defaults.GetMaxMediaSize()
|
||||||
|
|
@ -1680,7 +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 = annotateCurrentUserMessageForLLM(messages, ts.userMessage, annotatedUserMessage)
|
||||||
messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize)
|
messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -2053,7 +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)
|
messages = 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())
|
||||||
|
|
@ -3561,9 +3561,9 @@ func annotateCurrentUserMessageForLLM(
|
||||||
messages []providers.Message,
|
messages []providers.Message,
|
||||||
rawUserMessage string,
|
rawUserMessage string,
|
||||||
annotatedUserMessage string,
|
annotatedUserMessage string,
|
||||||
) {
|
) []providers.Message {
|
||||||
if rawUserMessage == annotatedUserMessage {
|
if rawUserMessage == annotatedUserMessage {
|
||||||
return
|
return messages
|
||||||
}
|
}
|
||||||
for i := len(messages) - 1; i >= 0; i-- {
|
for i := len(messages) - 1; i >= 0; i-- {
|
||||||
if messages[i].Role != "user" {
|
if messages[i].Role != "user" {
|
||||||
|
|
@ -3573,8 +3573,10 @@ func annotateCurrentUserMessageForLLM(
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
messages[i].Content = annotatedUserMessage
|
messages[i].Content = annotatedUserMessage
|
||||||
return
|
return messages
|
||||||
}
|
}
|
||||||
|
// Keep session history raw, but ensure metadata-only current turns still reach the LLM.
|
||||||
|
return append(messages, providers.Message{Role: "user", Content: annotatedUserMessage})
|
||||||
}
|
}
|
||||||
|
|
||||||
func formatUserMessageWithThreadMetadata(
|
func formatUserMessageWithThreadMetadata(
|
||||||
|
|
|
||||||
|
|
@ -268,6 +268,80 @@ func TestProcessMessage_AnnotatesThreadMetadataInPromptOnly(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProcessMessage_AnnotatesThreadMetadataOnlyPromptWhenContentEmpty(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
ModelName: "test-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &recordingProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
msg := bus.InboundMessage{
|
||||||
|
Channel: "discord",
|
||||||
|
SenderID: "discord:123",
|
||||||
|
ChatID: "group-1",
|
||||||
|
Content: "",
|
||||||
|
MessageID: "123",
|
||||||
|
Sender: bus.SenderInfo{
|
||||||
|
DisplayName: "Alice",
|
||||||
|
Username: "@alice",
|
||||||
|
},
|
||||||
|
Peer: bus.Peer{
|
||||||
|
Kind: "direct",
|
||||||
|
ID: "discord:123",
|
||||||
|
},
|
||||||
|
Metadata: map[string]string{
|
||||||
|
metadataKeyReplyToMessage: "120",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
response, err := al.processMessage(context.Background(), msg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
if response != "Mock response" {
|
||||||
|
t.Fatalf("processMessage() response = %q, want %q", response, "Mock response")
|
||||||
|
}
|
||||||
|
if len(provider.lastMessages) == 0 {
|
||||||
|
t.Fatal("provider did not receive any messages")
|
||||||
|
}
|
||||||
|
|
||||||
|
wantAnnotated := "[from:Alice (@alice); msgs:#123, reply_to:#120]"
|
||||||
|
lastMessage := provider.lastMessages[len(provider.lastMessages)-1]
|
||||||
|
if lastMessage.Role != "user" || lastMessage.Content != wantAnnotated {
|
||||||
|
t.Fatalf("last provider message = %+v, want user annotation %q", lastMessage, wantAnnotated)
|
||||||
|
}
|
||||||
|
|
||||||
|
route := al.registry.ResolveRoute(routing.RouteInput{
|
||||||
|
Channel: msg.Channel,
|
||||||
|
Peer: extractPeer(msg),
|
||||||
|
})
|
||||||
|
defaultAgent := al.registry.GetDefaultAgent()
|
||||||
|
if defaultAgent == nil {
|
||||||
|
t.Fatal("No default agent found")
|
||||||
|
}
|
||||||
|
history := defaultAgent.Sessions.GetHistory(route.SessionKey)
|
||||||
|
for _, h := range history {
|
||||||
|
if h.Role == "user" && h.Content == wantAnnotated {
|
||||||
|
t.Fatalf("history contains annotated prompt message: %+v", h)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestProcessMessage_UseCommandLoadsRequestedSkill(t *testing.T) {
|
func TestProcessMessage_UseCommandLoadsRequestedSkill(t *testing.T) {
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
skillDir := filepath.Join(tmpDir, "skills", "shell")
|
skillDir := filepath.Join(tmpDir, "skills", "shell")
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue